Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
200 changes: 200 additions & 0 deletions tests/chunking_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,76 @@ def test_empty(self):
with self.assertRaises(StopIteration):
next(sentence_iter)

def test_rejects_negative_start_position(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with self.assertRaisesRegex(IndexError, "can not be negative"):
chunking.SentenceIterator(tokenized_text, curr_token_pos=-1)

def test_rejects_start_position_past_document(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with self.assertRaisesRegex(IndexError, "past the length"):
chunking.SentenceIterator(
tokenized_text, curr_token_pos=len(tokenized_text.tokens) + 1
)


class TokenIntervalUtilTest(absltest.TestCase):

def test_create_token_interval_rejects_invalid_ranges(self):
with self.assertRaisesRegex(ValueError, "must be positive"):
chunking.create_token_interval(-1, 2)
with self.assertRaisesRegex(ValueError, "must be < end index"):
chunking.create_token_interval(2, 2)
with self.assertRaisesRegex(ValueError, "must be < end index"):
chunking.create_token_interval(3, 1)

def test_get_token_interval_text_rejects_empty_interval(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with self.assertRaisesRegex(ValueError, "must be < end index"):
chunking.get_token_interval_text(
tokenized_text, tokenizer.TokenInterval(start_index=1, end_index=1)
)

def test_get_token_interval_text_surfaces_tokenizer_interval_errors(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with self.assertRaises(tokenizer.InvalidTokenIntervalError):
chunking.get_token_interval_text(
tokenized_text, tokenizer.TokenInterval(start_index=0, end_index=99)
)

def test_get_token_interval_text_rejects_unexpected_empty_text(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with mock.patch.object(
chunking.tokenizer_lib, "tokens_text", return_value=""
):
with self.assertRaises(chunking.TokenUtilError):
chunking.get_token_interval_text(
tokenized_text,
tokenizer.TokenInterval(start_index=0, end_index=1),
)

def test_get_char_interval_rejects_empty_interval(self):
tokenized_text = tokenizer.tokenize("One sentence.")

with self.assertRaisesRegex(ValueError, "must be < end index"):
chunking.get_char_interval(
tokenized_text, tokenizer.TokenInterval(start_index=1, end_index=1)
)

def test_get_char_interval_uses_token_boundaries(self):
tokenized_text = tokenizer.tokenize("Alpha beta gamma.")

char_interval = chunking.get_char_interval(
tokenized_text, tokenizer.TokenInterval(start_index=1, end_index=3)
)

self.assertEqual(data.CharInterval(start_pos=6, end_pos=16), char_interval)


class ChunkIteratorTest(absltest.TestCase):

Expand Down Expand Up @@ -341,6 +411,58 @@ def test_tokenizer_propagation(self):
self.assertEqual(text_chunk.document_text, mock_tokenized_text)
self.assertEqual(text_chunk.chunk_text, text)

def test_constructor_requires_text_or_document(self):
with self.assertRaisesRegex(ValueError, "Either text or document"):
chunking.ChunkIterator(
text=None,
max_char_buffer=100,
tokenizer_impl=tokenizer.RegexTokenizer(),
)

def test_constructor_uses_document_text_when_text_is_none(self):
document = data.Document(text="Document fallback text.")
chunk_iter = chunking.ChunkIterator(
text=None,
max_char_buffer=100,
document=document,
tokenizer_impl=tokenizer.RegexTokenizer(),
)

self.assertEqual(next(chunk_iter).chunk_text, "Document fallback text.")

def test_constructor_retokenizes_empty_tokenized_text(self):
tokenized_text = tokenizer.TokenizedText(text="Retokenize this.", tokens=[])
tokenizer_impl = mock.Mock(spec=tokenizer.Tokenizer)
tokenizer_impl.tokenize.return_value = tokenizer.tokenize(
"Retokenize this."
)

chunk_iter = chunking.ChunkIterator(
text=tokenized_text,
max_char_buffer=100,
tokenizer_impl=tokenizer_impl,
)

self.assertEqual(next(chunk_iter).chunk_text, "Retokenize this.")
tokenizer_impl.tokenize.assert_called_once_with("Retokenize this.")

def test_broken_sentence_reset_prevents_following_sentence_merge(self):
text = "Alpha beta gamma delta. Zeta."
tokenized_text = tokenizer.tokenize(text)
chunk_iter = chunking.ChunkIterator(
tokenized_text,
max_char_buffer=18,
tokenizer_impl=tokenizer.RegexTokenizer(),
)

first_chunk = next(chunk_iter)
second_chunk = next(chunk_iter)
third_chunk = next(chunk_iter)

self.assertEqual(first_chunk.chunk_text, "Alpha beta gamma")
self.assertEqual(second_chunk.chunk_text, "delta.")
self.assertEqual(third_chunk.chunk_text, "Zeta.")


class BatchingTest(parameterized.TestCase):

Expand Down Expand Up @@ -443,6 +565,19 @@ def test_make_batches_of_textchunk(
"Batched chunks should match expected structure",
)

def test_make_batches_respects_short_batch_length(self):
chunk_iter = chunking.ChunkIterator(
self._SAMPLE_DOCUMENT.tokenized_text,
max_char_buffer=10,
tokenizer_impl=tokenizer.RegexTokenizer(),
)

batches = list(chunking.make_batches_of_textchunk(chunk_iter, 2))

self.assertLen(batches, 5)
self.assertLen(batches[0], 2)
self.assertLen(batches[-1], 1)


class TextChunkTest(absltest.TestCase):

Expand All @@ -465,6 +600,71 @@ def test_string_output(self):
text_chunk = next(chunk_iter)
self.assertEqual(str(text_chunk), expected)

def test_chunk_text_requires_document(self):
text_chunk = chunking.TextChunk(
token_interval=tokenizer.TokenInterval(start_index=0, end_index=1)
)

with self.assertRaisesRegex(ValueError, "document_text must be set"):
_ = text_chunk.chunk_text

def test_char_interval_requires_document(self):
text_chunk = chunking.TextChunk(
token_interval=tokenizer.TokenInterval(start_index=0, end_index=1)
)

with self.assertRaisesRegex(ValueError, "document_text must be set"):
_ = text_chunk.char_interval

def test_sanitized_chunk_text_collapses_whitespace(self):
document = data.Document(text="Alpha\n\n beta\tgamma.")
text_chunk = chunking.TextChunk(
token_interval=tokenizer.TokenInterval(
start_index=0, end_index=len(document.tokenized_text.tokens)
),
document=document,
)

self.assertEqual(text_chunk.sanitized_chunk_text, "Alpha beta gamma.")

def test_chunk_text_is_cached(self):
document = data.Document(text="Cache this text.")
text_chunk = chunking.TextChunk(
token_interval=tokenizer.TokenInterval(
start_index=0, end_index=len(document.tokenized_text.tokens)
),
document=document,
)

with mock.patch.object(
chunking,
"get_token_interval_text",
wraps=chunking.get_token_interval_text,
) as mock_get_text:
self.assertEqual(text_chunk.chunk_text, "Cache this text.")
self.assertEqual(text_chunk.chunk_text, "Cache this text.")

mock_get_text.assert_called_once()

def test_char_interval_is_cached(self):
document = data.Document(text="Cache this interval.")
text_chunk = chunking.TextChunk(
token_interval=tokenizer.TokenInterval(start_index=0, end_index=2),
document=document,
)

with mock.patch.object(
chunking, "get_char_interval", wraps=chunking.get_char_interval
) as mock_get_interval:
self.assertEqual(
text_chunk.char_interval, data.CharInterval(start_pos=0, end_pos=10)
)
self.assertEqual(
text_chunk.char_interval, data.CharInterval(start_pos=0, end_pos=10)
)

mock_get_interval.assert_called_once()


class TextAdditionalContextTest(absltest.TestCase):

Expand Down
Loading