1"""Base compressor class for context compression."""
2
3from __future__ import annotations
4
5from abc import ABC, abstractmethod
6
7from lexigram.ai.rag.context_compression.types import CompressionResult
8
9"""Convenience function for context compression."""
10
11from typing import Any
12
13from lexigram.ai.rag.context_compression.types import (
14 CompressionStrategy,
15)
16
17
18async def compress_context(
19 context: str | list[str],
20 strategy: CompressionStrategy = CompressionStrategy.EXTRACTIVE,
21 query: str | None = None,
22 **kwargs: Any,
23) -> CompressionResult:
24 """Convenience function for context compression.
25
26 Args:
27 context: Text or list of texts to compress.
28 strategy: Compression strategy to use.
29 query: Optional query for relevance-based compression.
30 **kwargs: Strategy-specific parameters.
31
32 Returns:
33 CompressionResult.
34
35 Example:
36 >>> result = await compress_context(
37 ... context=long_text,
38 ... strategy=CompressionStrategy.EXTRACTIVE,
39 ... query="What is AI?",
40 ... max_sentences=5
41 ... )
42 """
43 from lexigram.ai.rag.context_compression.strategy_registry import (
44 CompressionStrategyRegistry,
45 )
46
47 registry = CompressionStrategyRegistry.with_defaults()
48 return await registry.compress(strategy, context, query, kwargs)
49
50
51class AbstractCompressor(ABC):
52 """Base class for context compressors."""
53
54 @abstractmethod
55 async def compress(
56 self,
57 context: str | list[str],
58 query: str | None = None,
59 **kwargs,
60 ) -> CompressionResult:
61 """Compress context.
62
63 Args:
64 context: Text or list of texts to compress.
65 query: Optional query for relevance-based compression.
66 **kwargs: Additional compression parameters.
67
68 Returns:
69 CompressionResult with compressed text and statistics.
70 """
71
72 def _estimate_tokens(self, text: str) -> int:
73 """Estimate token count (rough approximation).
74
75 Uses simple heuristic: ~4 characters per token on average.
76 For production, use tiktoken or similar.
77 """
78 return len(text) // 4
79
80 def _normalize_context(self, context: str | list[str]) -> str:
81 """Normalize context to single string."""
82 if isinstance(context, list):
83 return "\n\n".join(str(c) for c in context)
84 return str(context)