Repository navigation
Expand file tree
/
Copy pathabstractive_compressor.py
More file actions
203 lines (158 loc) · 8.57 KB
/
Copy pathabstractive_compressor.py
File metadata and controls
203 lines (158 loc) · 8.57 KB
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
import asyncio
import tiktoken
from dataclasses import dataclass, field
from typing import Optional
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
from langchain_core.prompts import ChatPromptTemplate
COMPRESSION_SYSTEM = """You are a context compression engine for LLM conversations.
Your task: compress a block of conversation turns into a dense, lossless summary.
CRITICAL RULES — never violate these:
1. Preserve ALL named entities: variable names, function names, class names, file paths, URLs.
2. Preserve ALL numerical values: counts, prices, IDs, versions, scores, timestamps.
3. Preserve ALL error messages and exception types verbatim (in backticks).
4. Preserve ALL code snippets that were agreed upon or defined — use inline code.
5. Preserve ALL decisions: "we decided to use X", "user rejected Y", "the bug was Z".
6. Compress prose aggressively. Remove pleasantries, repetition, and verbose explanations.
7. Use structured format: bullet points, key:value pairs, inline code for identifiers.
8. Never invent information. Never generalize specific facts into vague statements.
9. Target compression ratio: 8:1 (800 tokens -> ~100 tokens).
Output format:
[COMPRESSED HISTORY - turns {start} to {end}]
• Context: <1-2 sentences of what was being worked on>
• Decisions: <bullet list of key decisions made>
• Code defined: <inline code for any functions/classes/variables established>
• Errors seen: <exact error messages if any>
• Current state: <where things stood at the end of this block>
"""
COMPRESSION_HUMAN = """Compress these conversation turns:
{conversation_block}
Previous compressed history (incorporate this context):
{previous_summary}
"""
@dataclass
class ConversationMessage:
role: str
content: str
@dataclass
class ContextWindow:
"""Manages a live conversation context with a token budget."""
budget_tokens: int
verbatim_turns: int = 8
compressor_model: str = "gpt-4o-mini"
full_model: str = "gpt-4o"
messages: list[ConversationMessage] = field(default_factory=list)
compressed_summary: str = ""
_total_turns: int = 0
def __post_init__(self):
self.enc = tiktoken.encoding_for_model(self.full_model)
self.llm = ChatOpenAI(model=self.compressor_model, temperature=0.0)
def count_tokens(self, text: str) -> int:
return len(self.enc.encode(text))
def total_tokens(self) -> int:
msg_tokens = sum(self.count_tokens(m.content) for m in self.messages)
summary_tokens = self.count_tokens(self.compressed_summary)
return msg_tokens + summary_tokens
def add_message(self, role: str, content: str):
"""Add a new message to the context."""
self.messages.append(ConversationMessage(role=role, content=content))
self._total_turns += 1
async def compress_oldest_block(self, block_size: int = 6) -> str:
"""
Compress the oldest `block_size` turns into a summary.
Incorporates any previously compressed summary for continuity.
Args:
block_size: Number of turns to compress in one pass.
Returns:
The new compressed summary string.
"""
if len(self.messages) <= block_size:
return self.compressed_summary
to_compress = self.messages[:block_size]
self.messages = self.messages[block_size:]
turn_start = self._total_turns - len(self.messages) - block_size + 1
turn_end = turn_start + block_size - 1
conversation_block = "\n\n".join(
f"[{m.role.upper()}]: {m.content}" for m in to_compress
)
prompt = ChatPromptTemplate.from_messages([
("system", COMPRESSION_SYSTEM.format(start=turn_start, end=turn_end)),
("human", COMPRESSION_HUMAN),
])
chain = prompt | self.llm
response = await chain.ainvoke({
"conversation_block": conversation_block,
"previous_summary": self.compressed_summary or "None — this is the first compression.",
})
new_summary = response.content.strip()
original_tokens = sum(self.count_tokens(m.content) for m in to_compress)
compressed_tokens = self.count_tokens(new_summary)
ratio = original_tokens / compressed_tokens if compressed_tokens else 0
print(
f"[Compressor] Turns {turn_start}-{turn_end}: "
f"{original_tokens} -> {compressed_tokens} tokens ({ratio:.1f}x compression)"
)
self.compressed_summary = new_summary
return new_summary
async def ensure_fits(self) -> bool:
"""
Compress until the context fits within the token budget.
Keeps the most recent `verbatim_turns` messages uncompressed.
Returns:
True if compression was performed, False if already within budget.
"""
if self.total_tokens() <= self.budget_tokens:
return False
print(f"[ContextWindow] Budget exceeded: {self.total_tokens()} / {self.budget_tokens} tokens. Compressing...")
while self.total_tokens() > self.budget_tokens and len(self.messages) > self.verbatim_turns:
compressible = len(self.messages) - self.verbatim_turns
block = min(6, compressible)
if block < 2:
break
await self.compress_oldest_block(block_size=block)
print(f"[ContextWindow] After compression: {self.total_tokens()} tokens")
return True
def build_prompt_messages(self) -> list:
"""
Build the final list of LangChain messages to send to the LLM.
Injects compressed history as a system context block.
"""
result = []
if self.compressed_summary:
result.append(SystemMessage(content=(
"The following is a compressed summary of earlier conversation history. "
"All facts, names, and code in this summary are accurate and should be treated "
"as if you had the full conversation.\n\n"
+ self.compressed_summary
)))
for msg in self.messages:
if msg.role == "user":
result.append(HumanMessage(content=msg.content))
elif msg.role == "assistant":
result.append(AIMessage(content=msg.content))
else:
result.append(SystemMessage(content=f"[{msg.role.upper()}]: {msg.content}"))
return result
async def main():
window = ContextWindow(budget_tokens=2000, verbatim_turns=4)
conversation = [
("user", "I need to build a FastAPI service that processes webhook payloads from Stripe."),
("assistant", "Let's start with the endpoint structure. We'll use a POST /webhook route with signature verification using the Stripe-Signature header and your webhook secret stored in STRIPE_WEBHOOK_SECRET env var."),
("user", "What about the payment_intent.succeeded event?"),
("assistant", "For payment_intent.succeeded, extract event.data.object which is a PaymentIntent. Key fields: id (pi_xxx), amount (in cents), currency, customer, metadata. Write handle_payment_succeeded(payment_intent: dict) to update order status."),
("user", "I'm getting a 400 error: No signatures found matching the expected signature for payload."),
("assistant", "That error means the raw request body is being consumed before Stripe can verify it. In FastAPI you must read the body with `await request.body()` not `await request.json()`. Pass the raw bytes directly to `stripe.Webhook.construct_event(payload, sig_header, secret)`."),
("user", "Fixed! Now I need to handle refund.created too."),
("assistant", "For refund.created: the event object is a Refund with fields id (re_xxx), payment_intent (pi_xxx), amount, status, reason. Write handle_refund_created(refund: dict) that looks up the order by payment_intent ID. Remember refunds can be partial."),
("user", "Should I store the raw webhook payload in the DB?"),
("assistant", "Yes, store it in a webhook_events table: id, stripe_event_id (UNIQUE), event_type, payload (JSONB), processed_at, status. Check stripe_event_id before processing to avoid duplicate handling on Stripe retries."),
]
for role, content in conversation:
window.add_message(role, content)
await window.ensure_fits()
print(f"\nFinal context: {window.total_tokens()} tokens")
print(f"Compressed summary:\n{window.compressed_summary}")
prompt_messages = window.build_prompt_messages()
print(f"\nMessages in final prompt: {len(prompt_messages)}")
if __name__ == "__main__":
asyncio.run(main())