81 lines
3.2 KiB
Python
81 lines
3.2 KiB
Python
"""Tests for spreadsheet layout detection and row splitting."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from rag_cut.parsers.pdf.tables import detect_spreadsheet_layout, is_qa_style_table
|
|
from rag_cut.pipeline import chunk_document
|
|
from rag_cut.splitters.by_row import split_by_row
|
|
from rag_cut.models import Block, BlockType, SplitConfig, SplitMode
|
|
|
|
|
|
BACKMAN_HEADER = [
|
|
[
|
|
"Session会话(必填):用于标识1个对话",
|
|
"query 用户输入(必填):消息内容",
|
|
"用户ID(必填)",
|
|
"使用的大语言模型(必填)",
|
|
"要求AI回复的语言(必填)",
|
|
"reference_output 标准答案(可选)",
|
|
],
|
|
["session", "query", "userid", "model", "lang", "reference_output"],
|
|
["1", "账户余额是多少?", "SUPPORT", "gpt-4o-mini", "zh", "SELECT 1"],
|
|
["2", "今日有哪些账户透支?", "SUPPORT", "gpt-4o-mini", "zh", "SELECT 2"],
|
|
]
|
|
|
|
|
|
class SpreadsheetLayoutTest(unittest.TestCase):
|
|
def test_detect_template_description_and_header_rows(self) -> None:
|
|
layout = detect_spreadsheet_layout(BACKMAN_HEADER)
|
|
self.assertEqual(layout["preamble_rows"], 1)
|
|
self.assertEqual(layout["header_row_start"], 2)
|
|
self.assertEqual(layout["header_row_end"], 2)
|
|
self.assertEqual(layout["data_start_row"], 3)
|
|
self.assertTrue(is_qa_style_table(BACKMAN_HEADER, layout))
|
|
|
|
def test_split_by_row_uses_real_header_and_one_row_per_chunk(self) -> None:
|
|
layout = detect_spreadsheet_layout(BACKMAN_HEADER)
|
|
block = Block(
|
|
type=BlockType.TABLE,
|
|
markdown="",
|
|
meta={"rows": BACKMAN_HEADER, "table_title": "testset", **layout},
|
|
)
|
|
groups = split_by_row(
|
|
[block],
|
|
SplitConfig(
|
|
mode=SplitMode.BY_ROW,
|
|
header_row_start=layout["header_row_start"],
|
|
header_row_end=layout["header_row_end"],
|
|
start_row=layout["data_start_row"],
|
|
rows_per_chunk=1,
|
|
),
|
|
)
|
|
self.assertEqual(len(groups), 2)
|
|
first_md = groups[0][0].markdown
|
|
self.assertIn("| session | query | userid | model | lang | reference_output |", first_md)
|
|
self.assertIn("账户余额是多少?", first_md)
|
|
self.assertNotIn("Session会话(必填)", first_md)
|
|
|
|
|
|
class BackmanFixtureTest(unittest.TestCase):
|
|
def test_backman_upload_chunks_one_row_each(self) -> None:
|
|
from pathlib import Path
|
|
|
|
uploads = list((Path(__file__).resolve().parents[2] / "storage" / "uploads").rglob("BackmanAI*.xlsx"))
|
|
if not uploads:
|
|
self.skipTest("BackmanAI fixture not uploaded")
|
|
result = chunk_document(uploads[0])
|
|
self.assertEqual(result.split_config["rows_per_chunk"], 1)
|
|
self.assertEqual(result.split_config["header_row_start"], 2)
|
|
self.assertEqual(result.split_config["start_row"], 3)
|
|
self.assertGreaterEqual(result.chunk_count, 40)
|
|
first = result.chunks[0].content
|
|
self.assertIn("| session | query |", first)
|
|
self.assertNotIn("Session会话(必填)", first)
|
|
self.assertLess(result.chunks[0].char_count, 1200)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|