Skip to main content

langchain_rust/text_splitter/
token_splitter.rs

1use async_trait::async_trait;
2use text_splitter::ChunkConfig;
3use tiktoken_rs::tokenizer::Tokenizer;
4
5use super::{SplitterOptions, TextSplitter, TextSplitterError};
6
7#[derive(Debug, Clone)]
8pub struct TokenSplitter {
9    splitter_options: SplitterOptions,
10}
11
12impl Default for TokenSplitter {
13    fn default() -> Self {
14        TokenSplitter::new(SplitterOptions::default())
15    }
16}
17
18impl TokenSplitter {
19    pub fn new(options: SplitterOptions) -> TokenSplitter {
20        TokenSplitter {
21            splitter_options: options,
22        }
23    }
24
25    #[deprecated = "Use `SplitterOptions::get_tokenizer_from_str` instead"]
26    pub fn get_tokenizer_from_str(&self, s: &str) -> Option<Tokenizer> {
27        match s.to_lowercase().as_str() {
28            "cl100k_base" => Some(Tokenizer::Cl100kBase),
29            "p50k_base" => Some(Tokenizer::P50kBase),
30            "r50k_base" => Some(Tokenizer::R50kBase),
31            "p50k_edit" => Some(Tokenizer::P50kEdit),
32            "gpt2" => Some(Tokenizer::Gpt2),
33            _ => None,
34        }
35    }
36}
37
38#[async_trait]
39impl TextSplitter for TokenSplitter {
40    async fn split_text(&self, text: &str) -> Result<Vec<String>, TextSplitterError> {
41        let chunk_config = ChunkConfig::try_from(&self.splitter_options)?;
42        Ok(text_splitter::TextSplitter::new(chunk_config)
43            .chunks(text)
44            .map(|x| x.to_string())
45            .collect())
46    }
47}