File size: 2,442 Bytes
8e874f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Tokenization utilities for QAFD-RAG.

Provides tiktoken-based encoding/decoding and token-aware list truncation.
"""

from typing import Callable, List, TypeVar

import tiktoken

# Global encoder instance (lazy initialized)
_ENCODER = None


def _get_encoder(model_name: str = "gpt-4o-mini"):
    """Get or initialize the tiktoken encoder."""
    global _ENCODER
    if _ENCODER is None:
        _ENCODER = tiktoken.encoding_for_model(model_name)
    return _ENCODER


def encode_string_by_tiktoken(content: str, model_name: str = "gpt-4o-mini") -> List[int]:
    """
    Encode a string into tokens using tiktoken.

    Parameters:
    -----------
    content : str
        Text content to encode
    model_name : str, optional
        Model name for tokenizer selection (default: gpt-4o-mini)

    Returns:
    --------
    List[int]
        List of token IDs
    """
    encoder = _get_encoder(model_name)
    return encoder.encode(content)


def decode_tokens_by_tiktoken(tokens: List[int], model_name: str = "gpt-4o-mini") -> str:
    """
    Decode tokens back to a string using tiktoken.

    Parameters:
    -----------
    tokens : List[int]
        List of token IDs
    model_name : str, optional
        Model name for tokenizer selection (default: gpt-4o-mini)

    Returns:
    --------
    str
        Decoded text content
    """
    encoder = _get_encoder(model_name)
    return encoder.decode(tokens)


T = TypeVar('T')


def truncate_list_by_token_size(
    list_data: List[T],
    key: Callable[[T], str],
    max_token_size: int
) -> List[T]:
    """
    Truncate a list based on cumulative token count.

    Iterates through the list and includes items until the total token
    count exceeds max_token_size.

    Parameters:
    -----------
    list_data : List[T]
        List of items to truncate
    key : Callable[[T], str]
        Function to extract text content from each item
    max_token_size : int
        Maximum total tokens allowed

    Returns:
    --------
    List[T]
        Truncated list that fits within token limit
    """
    if max_token_size <= 0:
        return []

    tokens = 0
    for i, data in enumerate(list_data):
        tokens += len(encode_string_by_tiktoken(key(data)))
        if tokens > max_token_size:
            return list_data[:i]
    return list_data


__all__ = [
    "encode_string_by_tiktoken",
    "decode_tokens_by_tiktoken",
    "truncate_list_by_token_size",
]