1"""Extractive compression strategies."""
2
3from __future__ import annotations
4
5from lexigram.ai.rag.context_compression.base import AbstractCompressor
6from lexigram.ai.rag.context_compression.types import (
7 CompressionResult,
8 CompressionStrategy,
9)
10
11
12class ExtractiveSummaryCompressor(AbstractCompressor):
13 """Extract most relevant sentences from context.
14
15 This compressor selects sentences based on:
16 - Relevance to query (if provided)
17 - Position in document (first/last sentences)
18 - Sentence length and informativeness
19
20 Example:
21 >>> compressor = ExtractiveSummaryCompressor(
22 ... max_sentences=5,
23 ... query_weight=0.7,
24 ... position_weight=0.3
25 ... )
26 >>> result = await compressor.compress(
27 ... context=long_text,
28 ... query="What is machine learning?"
29 ... )
30 """
31
32 def __init__(
33 self,
34 max_sentences: int = 5,
35 query_weight: float = 0.7,
36 position_weight: float = 0.3,
37 min_sentence_length: int = 20,
38 ):
39 """Initialize extractive compressor.
40
41 Args:
42 max_sentences: Maximum sentences to extract.
43 query_weight: Weight for query relevance (0.0 to 1.0).
44 position_weight: Weight for sentence position (0.0 to 1.0).
45 min_sentence_length: Minimum sentence length in characters.
46 """
47 self.max_sentences = max_sentences
48 self.query_weight = query_weight
49 self.position_weight = position_weight
50 self.min_sentence_length = min_sentence_length
51
52 async def compress(
53 self,
54 context: str | list[str],
55 query: str | None = None,
56 **kwargs,
57 ) -> CompressionResult:
58 """Compress by extracting most relevant sentences."""
59 from datetime import UTC, datetime
60
61 original_text = self._normalize_context(context)
62 original_tokens = self._estimate_tokens(original_text)
63
64 # Split into sentences
65 sentences = self._split_sentences(original_text)
66
67 # Filter short sentences
68 sentences = list(
69 filter(lambda s: len(s) >= self.min_sentence_length, sentences),
70 )
71
72 if len(sentences) <= self.max_sentences:
73 # Already short enough
74 compressed_text = original_text
75 else:
76 # Score and select top sentences
77 scored_sentences = self._score_sentences(sentences, query)
78 top_sentences = sorted(
79 scored_sentences,
80 key=lambda x: x[1],
81 reverse=True,
82 )[: self.max_sentences]
83
84 # Sort by original position to maintain flow
85 top_sentences = sorted(top_sentences, key=lambda x: x[2])
86
87 # Join selected sentences
88 compressed_text = " ".join(s[0] for s in top_sentences)
89
90 compressed_tokens = self._estimate_tokens(compressed_text)
91 compression_ratio = (
92 compressed_tokens / original_tokens if original_tokens > 0 else 1.0
93 )
94
95 return CompressionResult(
96 original_text=original_text,
97 compressed_text=compressed_text,
98 original_tokens=original_tokens,
99 compressed_tokens=compressed_tokens,
100 compression_ratio=compression_ratio,
101 strategy=CompressionStrategy.EXTRACTIVE,
102 metadata={
103 "total_sentences": len(sentences),
104 "selected_sentences": min(len(sentences), self.max_sentences),
105 "query_used": query is not None,
106 "timestamp": datetime.now(UTC).isoformat(),
107 },
108 )
109
110 def _split_sentences(self, text: str) -> list[str]:
111 """Split text into sentences (simple approach)."""
112 import re
113
114 # Split on period, exclamation, question mark followed by space/newline
115 sentences = re.split(r"[.!?]+\s+", text)
116 return list(map(str.strip, filter(str.strip, sentences)))
117
118 def _score_sentences(
119 self,
120 sentences: list[str],
121 query: str | None,
122 ) -> list[tuple[str, float, int]]:
123 """Score sentences based on relevance and position.
124
125 Returns:
126 List of (sentence, score, original_index) tuples.
127 """
128 scored = []
129 total = len(sentences)
130
131 for idx, sentence in enumerate(sentences):
132 score = 0.0
133
134 # Position score (first and last sentences are important)
135 if idx == 0:
136 position_score = 1.0
137 elif idx == total - 1:
138 position_score = 0.8
139 elif idx < total * 0.2: # First 20%
140 position_score = 0.7
141 elif idx > total * 0.8: # Last 20%
142 position_score = 0.6
143 else:
144 position_score = 0.3
145
146 score += position_score * self.position_weight
147
148 # Query relevance score
149 if query:
150 query_score = self._compute_relevance(sentence, query)
151 score += query_score * self.query_weight
152
153 # Length bonus (longer sentences often more informative)
154 length_score = min(len(sentence) / 200, 1.0) # Normalize to 200 chars
155 score += length_score * 0.1
156
157 scored.append((sentence, score, idx))
158
159 return scored
160
161 def _compute_relevance(self, sentence: str, query: str) -> float:
162 """Compute relevance score between sentence and query.
163
164 Simple approach using word overlap.
165 For production, use embeddings or cross-encoder.
166 """
167 # Lowercase and split into words
168 sentence_words = set(sentence.lower().split())
169 query_words = set(query.lower().split())
170
171 # Remove common stopwords (simple list)
172 stopwords = {
173 "the",
174 "a",
175 "an",
176 "is",
177 "are",
178 "was",
179 "were",
180 "in",
181 "on",
182 "at",
183 "to",
184 "for",
185 "of",
186 "and",
187 "or",
188 "but",
189 }
190 sentence_words -= stopwords
191 query_words -= stopwords
192
193 if not query_words:
194 return 0.0
195
196 # Calculate overlap
197 overlap = len(sentence_words & query_words)
198 return overlap / len(query_words)