chopratejas commited on
Commit
138b32e
·
2 Parent(s): 3d38c44dba81a5

Merge pull request #68 from KunalLohtia/feat/langgraph-compress-tool-messages

Browse files
OSS_PR_STRATEGY.md CHANGED
@@ -15,7 +15,7 @@ Contribute to popular LangChain ecosystem repos to demonstrate Headroom's value
15
  - **What**: Core LangGraph framework
16
  - **PR**: Add `compress_tool_messages` pre-model hook example
17
  - **Issues it addresses**: #3717 (ToolMessage overflow), #11405 (agent token limit), #2140 (127K tokens from plugin)
18
- - **Status**: TODO
19
 
20
  ### Priority 3: `langchain-ai/deepagents` (~17K stars)
21
  - **What**: LangChain's coding agent (like Claude Code but OSS)
 
15
  - **What**: Core LangGraph framework
16
  - **PR**: Add `compress_tool_messages` pre-model hook example
17
  - **Issues it addresses**: #3717 (ToolMessage overflow), #11405 (agent token limit), #2140 (127K tokens from plugin)
18
+ - **Status**: DONE — `compress_tool_messages()` and `create_compress_tool_messages_node()` in `headroom/integrations/langchain/langgraph.py`
19
 
20
  ### Priority 3: `langchain-ai/deepagents` (~17K stars)
21
  - **What**: LangChain's coding agent (like Claude Code but OSS)
docs/langchain.md CHANGED
@@ -340,6 +340,59 @@ print(f"Tokens saved: {llm.get_metrics()['tokens_saved']}")
340
 
341
  ---
342
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
343
  ### Example 2: RAG Pipeline with Document Filtering
344
 
345
  ```python
 
340
 
341
  ---
342
 
343
+ ### Example 1b: LangGraph Custom Graph with compress_tool_messages Node
344
+
345
+ If you're building a custom LangGraph `StateGraph` (instead of using `create_react_agent`),
346
+ you can insert a compression node between tools and the agent. This compresses all
347
+ `ToolMessage` content in the graph state before the LLM sees it.
348
+
349
+ ```python
350
+ from langchain_openai import ChatOpenAI
351
+ from langchain_core.messages import HumanMessage
352
+ from langgraph.graph import StateGraph, MessagesState, START, END
353
+ from headroom.integrations.langchain import create_compress_tool_messages_node
354
+
355
+ # Define your agent and tools nodes
356
+ def agent_node(state: MessagesState):
357
+ llm = ChatOpenAI(model="gpt-4o")
358
+ response = llm.invoke(state["messages"])
359
+ return {"messages": [response]}
360
+
361
+ def tools_node(state: MessagesState):
362
+ # Your tool execution logic here
363
+ ...
364
+
365
+ # Build the graph with a compression step
366
+ graph = StateGraph(MessagesState)
367
+ graph.add_node("agent", agent_node)
368
+ graph.add_node("tools", tools_node)
369
+ graph.add_node("compress", create_compress_tool_messages_node(
370
+ min_tokens_to_compress=100, # Only compress outputs > ~100 tokens
371
+ ))
372
+
373
+ # Wire: tools -> compress -> agent (instead of tools -> agent directly)
374
+ graph.add_edge(START, "agent")
375
+ graph.add_edge("tools", "compress")
376
+ graph.add_edge("compress", "agent")
377
+ # ... add conditional edges from agent to tools/END as needed
378
+
379
+ app = graph.compile()
380
+ result = app.invoke({"messages": [HumanMessage(content="Find sales data")]})
381
+ ```
382
+
383
+ You can also use `compress_tool_messages` directly as a standalone function:
384
+
385
+ ```python
386
+ from headroom.integrations.langchain import compress_tool_messages
387
+
388
+ # Compress ToolMessages in any list of LangChain messages
389
+ result = compress_tool_messages(messages, min_tokens_to_compress=100)
390
+ compressed_messages = result.messages
391
+ print(f"Saved {result.total_tokens_saved} tokens across {result.messages_compressed} messages")
392
+ ```
393
+
394
+ ---
395
+
396
  ### Example 2: RAG Pipeline with Document Filtering
397
 
398
  ```python
headroom/integrations/langchain/__init__.py CHANGED
@@ -7,6 +7,8 @@ This package provides seamless integration with LangChain, including:
7
  - HeadroomToolWrapper: Tool output compression for agents
8
  - StreamingMetricsTracker: Token counting during streaming
9
  - HeadroomLangSmithCallbackHandler: LangSmith trace enrichment
 
 
10
 
11
  Example:
12
  from langchain_openai import ChatOpenAI
@@ -21,7 +23,6 @@ Example:
21
  Install: pip install headroom[langchain]
22
  """
23
 
24
- # Core chat model wrapper
25
  # Agent tool wrapping
26
  from .agents import (
27
  HeadroomToolWrapper,
@@ -31,6 +32,8 @@ from .agents import (
31
  reset_tool_metrics,
32
  wrap_tools_with_headroom,
33
  )
 
 
34
  from .chat_model import (
35
  HeadroomCallbackHandler,
36
  HeadroomChatModel,
@@ -40,6 +43,15 @@ from .chat_model import (
40
  optimize_messages,
41
  )
42
 
 
 
 
 
 
 
 
 
 
43
  # LangSmith integration
44
  from .langsmith import (
45
  HeadroomLangSmithCallbackHandler,
@@ -93,6 +105,12 @@ __all__ = [
93
  "wrap_tools_with_headroom",
94
  "get_tool_metrics",
95
  "reset_tool_metrics",
 
 
 
 
 
 
96
  # LangSmith
97
  "HeadroomLangSmithCallbackHandler",
98
  "is_langsmith_available",
 
7
  - HeadroomToolWrapper: Tool output compression for agents
8
  - StreamingMetricsTracker: Token counting during streaming
9
  - HeadroomLangSmithCallbackHandler: LangSmith trace enrichment
10
+ - compress_tool_messages: LangGraph pre-model hook for ToolMessage compression
11
+ - create_compress_tool_messages_node: LangGraph node factory
12
 
13
  Example:
14
  from langchain_openai import ChatOpenAI
 
23
  Install: pip install headroom[langchain]
24
  """
25
 
 
26
  # Agent tool wrapping
27
  from .agents import (
28
  HeadroomToolWrapper,
 
32
  reset_tool_metrics,
33
  wrap_tools_with_headroom,
34
  )
35
+
36
+ # Core chat model wrapper
37
  from .chat_model import (
38
  HeadroomCallbackHandler,
39
  HeadroomChatModel,
 
43
  optimize_messages,
44
  )
45
 
46
+ # LangGraph integration
47
+ from .langgraph import (
48
+ CompressToolMessagesConfig,
49
+ CompressToolMessagesResult,
50
+ ToolMessageCompressionMetrics,
51
+ compress_tool_messages,
52
+ create_compress_tool_messages_node,
53
+ )
54
+
55
  # LangSmith integration
56
  from .langsmith import (
57
  HeadroomLangSmithCallbackHandler,
 
105
  "wrap_tools_with_headroom",
106
  "get_tool_metrics",
107
  "reset_tool_metrics",
108
+ # LangGraph
109
+ "compress_tool_messages",
110
+ "create_compress_tool_messages_node",
111
+ "CompressToolMessagesConfig",
112
+ "CompressToolMessagesResult",
113
+ "ToolMessageCompressionMetrics",
114
  # LangSmith
115
  "HeadroomLangSmithCallbackHandler",
116
  "is_langsmith_available",
headroom/integrations/langchain/langgraph.py ADDED
@@ -0,0 +1,399 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LangGraph integration for Headroom tool message compression.
2
+
3
+ This module provides a compress_tool_messages utility and a LangGraph-compatible
4
+ node factory for compressing ToolMessage content before it reaches the LLM,
5
+ solving context bloat from large tool outputs (JSON arrays, DB results, logs).
6
+
7
+ Addresses:
8
+ - LangGraph Issue #3717 (ToolMessage overflow)
9
+ - LangChain Issue #11405 (agent token limit)
10
+ - LangChain Issue #2140 (127K tokens from plugin)
11
+
12
+ Example:
13
+ from langgraph.graph import StateGraph, MessagesState
14
+ from headroom.integrations.langchain.langgraph import (
15
+ compress_tool_messages,
16
+ create_compress_tool_messages_node,
17
+ )
18
+
19
+ # Option 1: Use as a LangGraph node
20
+ graph = StateGraph(MessagesState)
21
+ graph.add_node("agent", agent_node)
22
+ graph.add_node("tools", tool_node)
23
+ graph.add_node("compress", create_compress_tool_messages_node())
24
+ graph.add_edge("tools", "compress")
25
+ graph.add_edge("compress", "agent")
26
+
27
+ # Option 2: Use as a standalone function
28
+ compressed = compress_tool_messages(messages)
29
+ """
30
+
31
+ from __future__ import annotations
32
+
33
+ import logging
34
+ import threading
35
+ from dataclasses import dataclass, field
36
+ from datetime import datetime, timezone
37
+ from typing import Any
38
+ from uuid import uuid4
39
+
40
+ # LangChain imports - optional dependencies
41
+ try:
42
+ from langchain_core.messages import BaseMessage, ToolMessage
43
+
44
+ LANGCHAIN_AVAILABLE = True
45
+ except ImportError:
46
+ LANGCHAIN_AVAILABLE = False
47
+ BaseMessage = object # type: ignore[misc,assignment]
48
+ ToolMessage = object # type: ignore[misc,assignment]
49
+
50
+ from headroom.transforms.smart_crusher import SmartCrusher, SmartCrusherConfig
51
+
52
+ logger = logging.getLogger(__name__)
53
+
54
+
55
+ def _check_langchain_available() -> None:
56
+ """Raise ImportError if LangChain is not installed."""
57
+ if not LANGCHAIN_AVAILABLE:
58
+ raise ImportError(
59
+ "LangChain is required for this integration. "
60
+ "Install with: pip install headroom[langchain] "
61
+ "or: pip install langchain-core"
62
+ )
63
+
64
+
65
+ def _estimate_tokens(text: str) -> int:
66
+ """Estimate token count using ~4 characters per token heuristic."""
67
+ if not text:
68
+ return 0
69
+ return len(text) // 4
70
+
71
+
72
+ @dataclass
73
+ class ToolMessageCompressionMetrics:
74
+ """Metrics from compressing a single ToolMessage."""
75
+
76
+ request_id: str
77
+ timestamp: datetime
78
+ tool_call_id: str
79
+ tokens_before: int
80
+ tokens_after: int
81
+ tokens_saved: int
82
+ savings_percent: float
83
+ was_compressed: bool
84
+ skip_reason: str | None = None
85
+
86
+
87
+ @dataclass
88
+ class CompressToolMessagesConfig:
89
+ """Configuration for compress_tool_messages.
90
+
91
+ Attributes:
92
+ min_tokens_to_compress: Minimum estimated token count in a ToolMessage
93
+ before compression is applied. Default 100.
94
+ preserve_errors: If True, skip compression on ToolMessages whose content
95
+ contains error indicators. Default True.
96
+ error_indicators: Strings that indicate a ToolMessage contains an error.
97
+ """
98
+
99
+ min_tokens_to_compress: int = 100
100
+ preserve_errors: bool = True
101
+ error_indicators: tuple[str, ...] = ('"error"', '"ERROR"', "Error:", "Traceback")
102
+
103
+
104
+ @dataclass
105
+ class CompressToolMessagesResult:
106
+ """Result from compress_tool_messages including metrics."""
107
+
108
+ messages: list[Any] # list[BaseMessage] but Any for when langchain not installed
109
+ metrics: list[ToolMessageCompressionMetrics] = field(default_factory=list)
110
+
111
+ @property
112
+ def total_tokens_saved(self) -> int:
113
+ """Total tokens saved across all compressed messages."""
114
+ return sum(m.tokens_saved for m in self.metrics if m.was_compressed)
115
+
116
+ @property
117
+ def messages_compressed(self) -> int:
118
+ """Number of messages that were actually compressed."""
119
+ return sum(1 for m in self.metrics if m.was_compressed)
120
+
121
+
122
+ class _CrusherSingleton:
123
+ """Thread-safe lazy singleton for SmartCrusher."""
124
+
125
+ def __init__(self, min_tokens: int) -> None:
126
+ self._crusher: SmartCrusher | None = None
127
+ self._min_tokens = min_tokens
128
+ self._lock = threading.Lock()
129
+
130
+ def get(self) -> SmartCrusher:
131
+ if self._crusher is None:
132
+ with self._lock:
133
+ if self._crusher is None:
134
+ config = SmartCrusherConfig(
135
+ min_tokens_to_crush=self._min_tokens,
136
+ )
137
+ self._crusher = SmartCrusher(config=config)
138
+ return self._crusher
139
+
140
+
141
+ # Module-level singleton, lazily initialized on first call
142
+ _crusher_singleton: _CrusherSingleton | None = None
143
+ _crusher_lock = threading.Lock()
144
+
145
+
146
+ def _get_crusher(min_tokens: int) -> SmartCrusher:
147
+ """Get or create the module-level SmartCrusher singleton."""
148
+ global _crusher_singleton
149
+ if _crusher_singleton is None:
150
+ with _crusher_lock:
151
+ if _crusher_singleton is None:
152
+ _crusher_singleton = _CrusherSingleton(min_tokens)
153
+ return _crusher_singleton.get()
154
+
155
+
156
+ def _should_skip(
157
+ content: str,
158
+ config: CompressToolMessagesConfig,
159
+ ) -> str | None:
160
+ """Check if a ToolMessage should skip compression.
161
+
162
+ Returns skip reason string, or None if it should be compressed.
163
+ """
164
+ if not content:
165
+ return "empty_content"
166
+
167
+ tokens = _estimate_tokens(content)
168
+ if tokens < config.min_tokens_to_compress:
169
+ return f"below_threshold:{tokens}<{config.min_tokens_to_compress}"
170
+
171
+ if config.preserve_errors:
172
+ for indicator in config.error_indicators:
173
+ if indicator in content:
174
+ return "error_content_preserved"
175
+
176
+ return None
177
+
178
+
179
+ def compress_tool_messages(
180
+ messages: list[BaseMessage], # type: ignore[type-arg]
181
+ *,
182
+ min_tokens_to_compress: int = 100,
183
+ preserve_errors: bool = True,
184
+ config: CompressToolMessagesConfig | None = None,
185
+ ) -> CompressToolMessagesResult:
186
+ """Compress ToolMessage content in a list of LangChain messages.
187
+
188
+ Iterates through messages, finds ToolMessages with large content,
189
+ and compresses them using SmartCrusher. Non-tool messages are
190
+ returned unchanged. tool_call_id is always preserved.
191
+
192
+ Args:
193
+ messages: List of LangChain BaseMessage objects.
194
+ min_tokens_to_compress: Minimum estimated tokens to trigger compression.
195
+ preserve_errors: If True, skip ToolMessages containing error indicators.
196
+ config: Full configuration object (overrides other kwargs if provided).
197
+
198
+ Returns:
199
+ CompressToolMessagesResult with compressed messages and metrics.
200
+
201
+ Example:
202
+ from langchain_core.messages import HumanMessage, AIMessage, ToolMessage
203
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
204
+
205
+ messages = [
206
+ HumanMessage(content="Get sales data"),
207
+ AIMessage(content="", tool_calls=[{"id": "call_1", "name": "db", "args": {}}]),
208
+ ToolMessage(content='[{"row": 1}, {"row": 2}, ...]', tool_call_id="call_1"),
209
+ ]
210
+
211
+ result = compress_tool_messages(messages)
212
+ print(f"Saved {result.total_tokens_saved} tokens")
213
+ compressed_messages = result.messages
214
+ """
215
+ _check_langchain_available()
216
+
217
+ if config is None:
218
+ config = CompressToolMessagesConfig(
219
+ min_tokens_to_compress=min_tokens_to_compress,
220
+ preserve_errors=preserve_errors,
221
+ )
222
+
223
+ crusher = _get_crusher(config.min_tokens_to_compress)
224
+ result_messages: list[BaseMessage] = []
225
+ metrics: list[ToolMessageCompressionMetrics] = []
226
+
227
+ for msg in messages:
228
+ if not isinstance(msg, ToolMessage):
229
+ result_messages.append(msg)
230
+ continue
231
+
232
+ content = msg.content if isinstance(msg.content, str) else str(msg.content)
233
+ request_id = str(uuid4())
234
+
235
+ # Check if we should skip
236
+ skip_reason = _should_skip(content, config)
237
+ if skip_reason:
238
+ result_messages.append(msg)
239
+ tokens = _estimate_tokens(content)
240
+ metrics.append(
241
+ ToolMessageCompressionMetrics(
242
+ request_id=request_id,
243
+ timestamp=datetime.now(timezone.utc),
244
+ tool_call_id=getattr(msg, "tool_call_id", "unknown"),
245
+ tokens_before=tokens,
246
+ tokens_after=tokens,
247
+ tokens_saved=0,
248
+ savings_percent=0.0,
249
+ was_compressed=False,
250
+ skip_reason=skip_reason,
251
+ )
252
+ )
253
+ logger.debug(
254
+ "Skipping ToolMessage %s compression: %s",
255
+ getattr(msg, "tool_call_id", "unknown"),
256
+ skip_reason,
257
+ )
258
+ continue
259
+
260
+ # Compress
261
+ tokens_before = _estimate_tokens(content)
262
+ try:
263
+ crush_result = crusher.crush(content=content, query="")
264
+ compressed_text = crush_result.compressed
265
+ was_modified = crush_result.was_modified
266
+ except Exception as e:
267
+ logger.warning(
268
+ "Compression failed for ToolMessage %s: %s. Keeping original.",
269
+ getattr(msg, "tool_call_id", "unknown"),
270
+ str(e),
271
+ )
272
+ result_messages.append(msg)
273
+ metrics.append(
274
+ ToolMessageCompressionMetrics(
275
+ request_id=request_id,
276
+ timestamp=datetime.now(timezone.utc),
277
+ tool_call_id=getattr(msg, "tool_call_id", "unknown"),
278
+ tokens_before=tokens_before,
279
+ tokens_after=tokens_before,
280
+ tokens_saved=0,
281
+ savings_percent=0.0,
282
+ was_compressed=False,
283
+ skip_reason=f"compression_error:{type(e).__name__}",
284
+ )
285
+ )
286
+ continue
287
+
288
+ tokens_after = _estimate_tokens(compressed_text)
289
+
290
+ if was_modified and tokens_after < tokens_before:
291
+ # Create new ToolMessage with compressed content, preserving tool_call_id
292
+ compressed_msg = ToolMessage(
293
+ content=compressed_text,
294
+ tool_call_id=msg.tool_call_id,
295
+ )
296
+ result_messages.append(compressed_msg)
297
+ tokens_saved = tokens_before - tokens_after
298
+
299
+ metrics.append(
300
+ ToolMessageCompressionMetrics(
301
+ request_id=request_id,
302
+ timestamp=datetime.now(timezone.utc),
303
+ tool_call_id=msg.tool_call_id,
304
+ tokens_before=tokens_before,
305
+ tokens_after=tokens_after,
306
+ tokens_saved=tokens_saved,
307
+ savings_percent=(tokens_saved / tokens_before * 100)
308
+ if tokens_before > 0
309
+ else 0.0,
310
+ was_compressed=True,
311
+ )
312
+ )
313
+
314
+ logger.info(
315
+ "Compressed ToolMessage %s: %d -> %d tokens (%.1f%% saved)",
316
+ msg.tool_call_id,
317
+ tokens_before,
318
+ tokens_after,
319
+ (tokens_saved / tokens_before * 100) if tokens_before > 0 else 0,
320
+ )
321
+ else:
322
+ # Compression didn't help, keep original
323
+ result_messages.append(msg)
324
+ metrics.append(
325
+ ToolMessageCompressionMetrics(
326
+ request_id=request_id,
327
+ timestamp=datetime.now(timezone.utc),
328
+ tool_call_id=msg.tool_call_id,
329
+ tokens_before=tokens_before,
330
+ tokens_after=tokens_before,
331
+ tokens_saved=0,
332
+ savings_percent=0.0,
333
+ was_compressed=False,
334
+ skip_reason="no_reduction",
335
+ )
336
+ )
337
+
338
+ return CompressToolMessagesResult(messages=result_messages, metrics=metrics)
339
+
340
+
341
+ def create_compress_tool_messages_node(
342
+ *,
343
+ min_tokens_to_compress: int = 100,
344
+ preserve_errors: bool = True,
345
+ config: CompressToolMessagesConfig | None = None,
346
+ ) -> Any:
347
+ """Create a LangGraph node that compresses ToolMessages in graph state.
348
+
349
+ Returns a function compatible with LangGraph's StateGraph that reads
350
+ messages from state, compresses ToolMessages, and returns updated state.
351
+
352
+ Args:
353
+ min_tokens_to_compress: Minimum estimated tokens to trigger compression.
354
+ preserve_errors: If True, skip ToolMessages containing error indicators.
355
+ config: Full configuration object (overrides other kwargs if provided).
356
+
357
+ Returns:
358
+ A callable suitable for use as a LangGraph node.
359
+
360
+ Example:
361
+ from langgraph.graph import StateGraph, MessagesState
362
+
363
+ graph = StateGraph(MessagesState)
364
+ graph.add_node("agent", agent_node)
365
+ graph.add_node("tools", tool_node)
366
+ graph.add_node("compress", create_compress_tool_messages_node(
367
+ min_tokens_to_compress=200,
368
+ ))
369
+
370
+ # Wire: tools -> compress -> agent
371
+ graph.add_edge("tools", "compress")
372
+ graph.add_edge("compress", "agent")
373
+ """
374
+ _check_langchain_available()
375
+
376
+ if config is None:
377
+ config = CompressToolMessagesConfig(
378
+ min_tokens_to_compress=min_tokens_to_compress,
379
+ preserve_errors=preserve_errors,
380
+ )
381
+
382
+ def compress_node(state: dict[str, Any]) -> dict[str, Any]:
383
+ """LangGraph node that compresses ToolMessages in state.
384
+
385
+ Args:
386
+ state: LangGraph state dict containing a "messages" key.
387
+
388
+ Returns:
389
+ Updated state dict with compressed messages.
390
+ """
391
+ messages = state.get("messages", [])
392
+ if not messages:
393
+ return state
394
+
395
+ result = compress_tool_messages(messages, config=config)
396
+
397
+ return {"messages": result.messages}
398
+
399
+ return compress_node
tests/test_integrations/langchain/test_langgraph.py ADDED
@@ -0,0 +1,342 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for LangGraph tool message compression integration.
2
+
3
+ Tests cover:
4
+ 1. compress_tool_messages - Compresses large ToolMessages in a message list
5
+ 2. create_compress_tool_messages_node - LangGraph node factory
6
+ 3. CompressToolMessagesConfig - Configuration options
7
+ 4. CompressToolMessagesResult - Result with metrics
8
+ 5. ToolMessageCompressionMetrics - Per-message metrics
9
+ """
10
+
11
+ import json
12
+ from unittest.mock import MagicMock, patch
13
+
14
+ import pytest
15
+
16
+ # Check if LangChain is available
17
+ try:
18
+ from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, ToolMessage
19
+
20
+ LANGCHAIN_AVAILABLE = True
21
+ except ImportError:
22
+ LANGCHAIN_AVAILABLE = False
23
+
24
+ # Skip all tests if LangChain not installed
25
+ pytestmark = pytest.mark.skipif(not LANGCHAIN_AVAILABLE, reason="LangChain not installed")
26
+
27
+
28
+ def _make_large_tool_output(num_items: int = 200) -> str:
29
+ """Generate a large JSON array string that will trigger compression."""
30
+ items = [{"id": i, "name": f"item_{i}", "value": i * 1.5, "status": "ok"} for i in range(num_items)]
31
+ return json.dumps(items)
32
+
33
+
34
+ def _make_messages_with_tool_output(tool_content: str, tool_call_id: str = "call_1") -> list:
35
+ """Create a typical message sequence with a tool call and result."""
36
+ return [
37
+ HumanMessage(content="Get the data"),
38
+ AIMessage(content="", tool_calls=[{"id": tool_call_id, "name": "search", "args": {}}]),
39
+ ToolMessage(content=tool_content, tool_call_id=tool_call_id),
40
+ ]
41
+
42
+
43
+ class TestCompressToolMessages:
44
+ """Tests for the compress_tool_messages function."""
45
+
46
+ def test_compresses_large_tool_message(self):
47
+ """Large ToolMessage content should be compressed."""
48
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
49
+
50
+ large_output = _make_large_tool_output(200)
51
+ messages = _make_messages_with_tool_output(large_output)
52
+
53
+ result = compress_tool_messages(messages)
54
+
55
+ # Should have same number of messages
56
+ assert len(result.messages) == 3
57
+ # ToolMessage should be smaller
58
+ compressed_content = result.messages[2].content
59
+ assert len(compressed_content) < len(large_output)
60
+
61
+ def test_preserves_small_tool_messages(self):
62
+ """Small ToolMessages should not be compressed."""
63
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
64
+
65
+ small_output = '{"result": "ok"}'
66
+ messages = _make_messages_with_tool_output(small_output)
67
+
68
+ result = compress_tool_messages(messages)
69
+
70
+ # Content should be unchanged
71
+ assert result.messages[2].content == small_output
72
+ assert result.messages_compressed == 0
73
+
74
+ def test_preserves_non_tool_messages(self):
75
+ """HumanMessage and AIMessage should pass through unchanged."""
76
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
77
+
78
+ large_output = _make_large_tool_output(200)
79
+ messages = _make_messages_with_tool_output(large_output)
80
+
81
+ result = compress_tool_messages(messages)
82
+
83
+ assert isinstance(result.messages[0], HumanMessage)
84
+ assert result.messages[0].content == "Get the data"
85
+ assert isinstance(result.messages[1], AIMessage)
86
+ tool_call = result.messages[1].tool_calls[0]
87
+ assert tool_call["id"] == "call_1"
88
+ assert tool_call["name"] == "search"
89
+ assert tool_call["args"] == {}
90
+
91
+ def test_preserves_tool_call_id(self):
92
+ """Compressed ToolMessages must keep their tool_call_id."""
93
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
94
+
95
+ large_output = _make_large_tool_output(200)
96
+ messages = _make_messages_with_tool_output(large_output, tool_call_id="call_abc123")
97
+
98
+ result = compress_tool_messages(messages)
99
+
100
+ tool_msg = result.messages[2]
101
+ assert isinstance(tool_msg, ToolMessage)
102
+ assert tool_msg.tool_call_id == "call_abc123"
103
+
104
+ def test_preserves_error_content_by_default(self):
105
+ """ToolMessages with error indicators should be skipped by default."""
106
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
107
+
108
+ # Large content but contains error indicator
109
+ error_output = json.dumps({
110
+ "error": "Database connection failed",
111
+ "details": "x" * 2000,
112
+ })
113
+ messages = _make_messages_with_tool_output(error_output)
114
+
115
+ result = compress_tool_messages(messages)
116
+
117
+ # Should be unchanged — error preserved
118
+ assert result.messages[2].content == error_output
119
+ assert result.metrics[0].skip_reason == "error_content_preserved"
120
+
121
+ def test_compresses_error_content_when_disabled(self):
122
+ """Error content should be compressed when preserve_errors=False."""
123
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
124
+
125
+ error_output = json.dumps({
126
+ "error": "fail",
127
+ "data": [{"id": i} for i in range(200)],
128
+ })
129
+ messages = _make_messages_with_tool_output(error_output)
130
+
131
+ result = compress_tool_messages(messages, preserve_errors=False)
132
+
133
+ # Should have attempted compression (no error_content_preserved skip)
134
+ assert result.metrics[0].skip_reason != "error_content_preserved"
135
+
136
+ def test_handles_empty_messages(self):
137
+ """Empty message list should return empty result."""
138
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
139
+
140
+ result = compress_tool_messages([])
141
+
142
+ assert result.messages == []
143
+ assert result.metrics == []
144
+ assert result.total_tokens_saved == 0
145
+
146
+ def test_handles_no_tool_messages(self):
147
+ """Message list with no ToolMessages should pass through."""
148
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
149
+
150
+ messages = [
151
+ HumanMessage(content="Hello"),
152
+ AIMessage(content="Hi there!"),
153
+ ]
154
+
155
+ result = compress_tool_messages(messages)
156
+
157
+ assert len(result.messages) == 2
158
+ assert result.messages[0].content == "Hello"
159
+ assert result.messages[1].content == "Hi there!"
160
+ assert result.metrics == []
161
+
162
+ def test_multiple_tool_messages(self):
163
+ """Should compress multiple ToolMessages independently."""
164
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
165
+
166
+ large_output_1 = _make_large_tool_output(200)
167
+ large_output_2 = _make_large_tool_output(150)
168
+
169
+ messages = [
170
+ HumanMessage(content="Get all data"),
171
+ AIMessage(
172
+ content="",
173
+ tool_calls=[
174
+ {"id": "call_1", "name": "search", "args": {}},
175
+ {"id": "call_2", "name": "database", "args": {}},
176
+ ],
177
+ ),
178
+ ToolMessage(content=large_output_1, tool_call_id="call_1"),
179
+ ToolMessage(content=large_output_2, tool_call_id="call_2"),
180
+ ]
181
+
182
+ result = compress_tool_messages(messages)
183
+
184
+ assert len(result.messages) == 4
185
+ # Both tool messages should have their correct tool_call_ids
186
+ assert result.messages[2].tool_call_id == "call_1"
187
+ assert result.messages[3].tool_call_id == "call_2"
188
+
189
+ def test_min_tokens_to_compress_config(self):
190
+ """Custom min_tokens_to_compress should be respected."""
191
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
192
+
193
+ # Content that's ~100 tokens (400 chars) — below a 200 token threshold
194
+ medium_output = json.dumps({"data": "x" * 400})
195
+ messages = _make_messages_with_tool_output(medium_output)
196
+
197
+ result = compress_tool_messages(messages, min_tokens_to_compress=200)
198
+
199
+ # Should be skipped due to being below threshold
200
+ assert result.metrics[0].was_compressed is False
201
+ assert "below_threshold" in (result.metrics[0].skip_reason or "")
202
+
203
+
204
+ class TestCompressToolMessagesResult:
205
+ """Tests for CompressToolMessagesResult properties."""
206
+
207
+ def test_total_tokens_saved(self):
208
+ """total_tokens_saved should sum across compressed metrics."""
209
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
210
+
211
+ large_output = _make_large_tool_output(200)
212
+ messages = _make_messages_with_tool_output(large_output)
213
+
214
+ result = compress_tool_messages(messages)
215
+
216
+ assert result.total_tokens_saved >= 0
217
+ # If compression happened, tokens_saved should be positive
218
+ if result.messages_compressed > 0:
219
+ assert result.total_tokens_saved > 0
220
+
221
+ def test_messages_compressed_count(self):
222
+ """messages_compressed should count actually compressed messages."""
223
+ from headroom.integrations.langchain.langgraph import compress_tool_messages
224
+
225
+ messages = [
226
+ HumanMessage(content="test"),
227
+ ToolMessage(content='{"small": true}', tool_call_id="call_1"),
228
+ ]
229
+
230
+ result = compress_tool_messages(messages)
231
+
232
+ assert result.messages_compressed == 0
233
+
234
+
235
+ class TestCompressToolMessagesConfig:
236
+ """Tests for CompressToolMessagesConfig."""
237
+
238
+ def test_config_object(self):
239
+ """Config object should override kwargs."""
240
+ from headroom.integrations.langchain.langgraph import (
241
+ CompressToolMessagesConfig,
242
+ compress_tool_messages,
243
+ )
244
+
245
+ config = CompressToolMessagesConfig(
246
+ min_tokens_to_compress=500,
247
+ preserve_errors=False,
248
+ )
249
+
250
+ medium_output = json.dumps({"data": "x" * 800})
251
+ messages = _make_messages_with_tool_output(medium_output)
252
+
253
+ result = compress_tool_messages(messages, config=config)
254
+
255
+ # ~200 tokens, below the 500 threshold
256
+ assert result.metrics[0].was_compressed is False
257
+
258
+ def test_default_config(self):
259
+ """Default config should have sensible defaults."""
260
+ from headroom.integrations.langchain.langgraph import CompressToolMessagesConfig
261
+
262
+ config = CompressToolMessagesConfig()
263
+ assert config.min_tokens_to_compress == 100
264
+ assert config.preserve_errors is True
265
+
266
+
267
+ class TestCreateCompressToolMessagesNode:
268
+ """Tests for the LangGraph node factory."""
269
+
270
+ def test_returns_callable(self):
271
+ """Factory should return a callable node function."""
272
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
273
+
274
+ node = create_compress_tool_messages_node()
275
+ assert callable(node)
276
+
277
+ def test_node_reads_messages_from_state(self):
278
+ """Node should read messages from state dict and return updated state."""
279
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
280
+
281
+ large_output = _make_large_tool_output(200)
282
+ state = {
283
+ "messages": _make_messages_with_tool_output(large_output),
284
+ }
285
+
286
+ node = create_compress_tool_messages_node()
287
+ result_state = node(state)
288
+
289
+ assert "messages" in result_state
290
+ assert len(result_state["messages"]) == 3
291
+ # ToolMessage should be compressed
292
+ assert len(result_state["messages"][2].content) < len(large_output)
293
+
294
+ def test_node_preserves_tool_call_id(self):
295
+ """Node should preserve tool_call_id on compressed messages."""
296
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
297
+
298
+ large_output = _make_large_tool_output(200)
299
+ state = {
300
+ "messages": [
301
+ HumanMessage(content="test"),
302
+ AIMessage(content="", tool_calls=[{"id": "call_xyz", "name": "db", "args": {}}]),
303
+ ToolMessage(content=large_output, tool_call_id="call_xyz"),
304
+ ],
305
+ }
306
+
307
+ node = create_compress_tool_messages_node()
308
+ result_state = node(state)
309
+
310
+ assert result_state["messages"][2].tool_call_id == "call_xyz"
311
+
312
+ def test_node_handles_empty_state(self):
313
+ """Node should handle empty messages gracefully."""
314
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
315
+
316
+ node = create_compress_tool_messages_node()
317
+ result_state = node({"messages": []})
318
+
319
+ assert result_state == {"messages": []}
320
+
321
+ def test_node_handles_missing_messages_key(self):
322
+ """Node should handle state without messages key."""
323
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
324
+
325
+ node = create_compress_tool_messages_node()
326
+ result_state = node({})
327
+
328
+ assert "messages" not in result_state or result_state.get("messages") == []
329
+
330
+ def test_node_with_custom_config(self):
331
+ """Node should respect custom configuration."""
332
+ from headroom.integrations.langchain.langgraph import create_compress_tool_messages_node
333
+
334
+ node = create_compress_tool_messages_node(min_tokens_to_compress=10000)
335
+
336
+ large_output = _make_large_tool_output(200)
337
+ state = {"messages": _make_messages_with_tool_output(large_output)}
338
+
339
+ result_state = node(state)
340
+
341
+ # With very high threshold, nothing should be compressed
342
+ assert result_state["messages"][2].content == large_output