Skip to content

Commit 294afae

Browse files
committed
feat(segment): 新增文档片段管理功能
新增文档片段的创建、更新、删除、启用/禁用等功能,包括相关handler、service、schema的添加和修改。同时优化了分页查询逻辑和向量数据库的批量处理性能。
1 parent bdc209f commit 294afae

12 files changed

Lines changed: 2145 additions & 632 deletions

File tree

internal/handler/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from .upload_file_handler import UploadFileHandler
55
from .dataset_handler import DatasetHandler
66
from .document_handler import DocumentHandler
7+
from .segment_handler import SegmentHandler
78

89
__all__ = [
910
"AppHandler",
@@ -12,4 +13,5 @@
1213
"UploadFileHandler",
1314
"DatasetHandler",
1415
"DocumentHandler",
16+
"SegmentHandler",
1517
]

internal/handler/dataset_handler.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,14 @@ def hit(self, dataset_id: UUID):
3434
"""根据传递的知识库id+检索参数执行召回测试"""
3535
pass
3636

37+
q = "alembic==1.15.2"
38+
retriever = self.vector_database_service.vector_store.as_retriever(
39+
search_type="mmr",
40+
search_kwargs={"k": 10},
41+
)
42+
documents = retriever.invoke(q)
43+
return success_json({"documents": [doc.page_content for doc in documents]})
44+
3745
def embeddings_query(self):
3846
# upload_file = self.db.session.query(UploadFile).get(
3947
# "94acfe76-bbad-4d4c-9751-bfcd36bca124"
Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
from dataclasses import dataclass
2+
from uuid import UUID
3+
from injector import inject
4+
5+
from internal.schema.segment_schema import (
6+
GetSegmentsWithPageReq,
7+
GetSegmentsWithPageResp,
8+
GetSegmentResp,
9+
UpdateSegmentEnabledReq,
10+
)
11+
from internal.service.segment_service import SegmentService
12+
from pkg.response.response import validate_error_json, success_json, success_message
13+
from pkg.paginator import PageModel
14+
15+
16+
@inject
17+
@dataclass
18+
class SegmentHandler:
19+
20+
segment_service: SegmentService
21+
22+
def get_segments_with_page(self, dataset_id: UUID, document_id: UUID):
23+
24+
req = GetSegmentsWithPageReq()
25+
26+
if not req.validate():
27+
return validate_error_json(req.errors)
28+
29+
segments, paginator = self.segment_service.get_segments_with_page(
30+
dataset_id, document_id, req
31+
)
32+
33+
resp = GetSegmentsWithPageResp(many=True)
34+
35+
return success_json(PageModel(list=resp.dump(segments), paginator=paginator))
36+
37+
def create_segment(self, dataset_id: UUID, document_id: UUID):
38+
pass
39+
40+
def get_segment(self, dataset_id: UUID, document_id: UUID, segment_id: UUID):
41+
"""获取指定的文档片段信息详情"""
42+
segment = self.segment_service.get_segment(dataset_id, document_id, segment_id)
43+
resp = GetSegmentResp()
44+
return success_json(resp.dump(segment))
45+
46+
def update_segment_enabled(
47+
self, dataset_id: UUID, document_id: UUID, segment_id: UUID
48+
):
49+
"""根据传递的信息更新指定的文档片段启用状态"""
50+
# 1.提取请求并校验
51+
req = UpdateSegmentEnabledReq()
52+
if not req.validate():
53+
return validate_error_json(req.errors)
54+
55+
# 2.调用服务更新文档片段的启用状态
56+
self.segment_service.update_segment_enabled(
57+
dataset_id, document_id, segment_id, req.enabled.data
58+
)
59+
60+
return success_message("修改片段状态成功")

internal/model/dataset.py

Lines changed: 15 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
from internal.extension.database_extension import db
1616
from .app import AppDatasetJoin
17-
from internal.model.upload_file import UploadFile
17+
from .upload_file import UploadFile
1818

1919

2020
class Dataset(db.Model):
@@ -122,25 +122,27 @@ class Document(db.Model):
122122
)
123123

124124
@property
125-
def upload_file(self):
126-
"""只读属性,获取文档关联的上传文件"""
125+
def upload_file(self) -> "UploadFile":
127126
return (
128127
db.session.query(UploadFile)
129-
.filter(UploadFile.id == self.upload_file_id)
128+
.filter(
129+
UploadFile.id == self.upload_file_id,
130+
)
130131
.one_or_none()
131132
)
132133

133134
@property
134-
def process_rule(self):
135-
"""只读属性,获取文档关联的处理规则"""
135+
def process_rule(self) -> "ProcessRule":
136136
return (
137137
db.session.query(ProcessRule)
138-
.filter(ProcessRule.id == self.process_rule_id)
138+
.filter(
139+
ProcessRule.id == self.process_rule_id,
140+
)
139141
.one_or_none()
140142
)
141143

142144
@property
143-
def segment_count(self):
145+
def segment_count(self) -> int:
144146
return (
145147
db.session.query(func.count(Segment.id))
146148
.filter(
@@ -150,7 +152,7 @@ def segment_count(self):
150152
)
151153

152154
@property
153-
def hit_count(self):
155+
def hit_count(self) -> int:
154156
return (
155157
db.session.query(func.coalesce(func.sum(Segment.hit_count), 0))
156158
.filter(
@@ -200,6 +202,10 @@ class Segment(db.Model):
200202
DateTime, nullable=False, server_default=text("CURRENT_TIMESTAMP(0)")
201203
)
202204

205+
@property
206+
def document(self) -> "Document":
207+
return db.session.query(Document).get(self.document_id)
208+
203209

204210
class KeywordTable(db.Model):
205211
"""关键词表模型"""

internal/router/router.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
from internal.handler import AppHandler, BuiltinToolHandler, ApiToolHandler
77
from internal.handler.dataset_handler import DatasetHandler
88
from internal.handler.document_handler import DocumentHandler
9+
from internal.handler.segment_handler import SegmentHandler
910
from internal.handler.upload_file_handler import UploadFileHandler
1011

1112

@@ -20,6 +21,7 @@ class Router:
2021
upload_file_handler: UploadFileHandler
2122
dataset_handler: DatasetHandler
2223
document_handler: DocumentHandler
24+
segment_handler: SegmentHandler
2325

2426
def register_router(self, app: Flask):
2527
"""注册路由"""
@@ -182,5 +184,25 @@ def register_router(self, app: Flask):
182184
view_func=self.document_handler.delete_document,
183185
)
184186

187+
# 文档片段模块
188+
bp.add_url_rule(
189+
"/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments",
190+
view_func=self.segment_handler.get_segments_with_page,
191+
)
192+
bp.add_url_rule(
193+
"/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments",
194+
methods=["POST"],
195+
view_func=self.segment_handler.create_segment,
196+
)
197+
bp.add_url_rule(
198+
"/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments/<uuid:segment_id>",
199+
view_func=self.segment_handler.get_segment,
200+
)
201+
bp.add_url_rule(
202+
"/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments/<uuid:segment_id>/enabled",
203+
methods=["POST"],
204+
view_func=self.segment_handler.update_segment_enabled,
205+
)
206+
185207
# 在应用上去注册蓝图
186208
app.register_blueprint(bp)

internal/schema/__init__.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,10 @@
99
GetDatasetsWithPageReq,
1010
GetDatasetsWithPageResp,
1111
)
12+
from .segment_schema import (
13+
GetSegmentsWithPageReq,
14+
GetSegmentsWithPageResp,
15+
)
1216

1317
__all__ = [
1418
"CompletionReq",
@@ -23,4 +27,6 @@
2327
"GetDatasetResp",
2428
"GetDatasetsWithPageReq",
2529
"GetDatasetsWithPageResp",
30+
"GetSegmentsWithPageReq",
31+
"GetSegmentsWithPageResp",
2632
]

internal/schema/segment_schema.py

Lines changed: 162 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,162 @@
1+
from flask_wtf import FlaskForm
2+
from marshmallow import Schema, fields, pre_dump
3+
from wtforms import StringField, BooleanField
4+
from wtforms.validators import Optional, ValidationError, DataRequired
5+
6+
from internal.lib.helper import datetime_to_timestamp
7+
from internal.model import Segment
8+
from pkg.paginator import PaginatorReq
9+
from .schema import ListField
10+
11+
12+
class GetSegmentsWithPageReq(PaginatorReq):
13+
"""获取文档片段列表请求"""
14+
15+
search_word = StringField("search_word", default="", validators=[Optional()])
16+
17+
18+
class GetSegmentsWithPageResp(Schema):
19+
"""获取文档片段列表响应结构"""
20+
21+
id = fields.UUID(dump_default="")
22+
document_id = fields.UUID(dump_default="")
23+
dataset_id = fields.UUID(dump_default="")
24+
position = fields.Integer(dump_default=0)
25+
content = fields.String(dump_default="")
26+
keywords = fields.List(fields.String, dump_default=[])
27+
character_count = fields.Integer(dump_default=0)
28+
token_count = fields.Integer(dump_default=0)
29+
hit_count = fields.Integer(dump_default=0)
30+
enabled = fields.Boolean(dump_default=False)
31+
disabled_at = fields.Integer(dump_default=0)
32+
status = fields.String(dump_default="")
33+
error = fields.String(dump_default="")
34+
updated_at = fields.Integer(dump_default=0)
35+
created_at = fields.Integer(dump_default=0)
36+
37+
@pre_dump
38+
def process_data(self, data: Segment, **kwargs):
39+
return {
40+
"id": data.id,
41+
"document_id": data.document_id,
42+
"dataset_id": data.dataset_id,
43+
"position": data.position,
44+
"content": data.content,
45+
"keywords": data.keywords,
46+
"character_count": data.character_count,
47+
"token_count": data.token_count,
48+
"hit_count": data.hit_count,
49+
"enabled": data.enabled,
50+
"disabled_at": datetime_to_timestamp(data.disabled_at),
51+
"status": data.status,
52+
"error": data.error,
53+
"updated_at": datetime_to_timestamp(data.updated_at),
54+
"created_at": datetime_to_timestamp(data.created_at),
55+
}
56+
57+
58+
class GetSegmentResp(Schema):
59+
"""获取文档详情响应结构"""
60+
61+
id = fields.UUID(dump_default="")
62+
document_id = fields.UUID(dump_default="")
63+
dataset_id = fields.UUID(dump_default="")
64+
position = fields.Integer(dump_default=0)
65+
content = fields.String(dump_default="")
66+
keywords = fields.List(fields.String, dump_default=[])
67+
character_count = fields.Integer(dump_default=0)
68+
token_count = fields.Integer(dump_default=0)
69+
hit_count = fields.Integer(dump_default=0)
70+
hash = fields.String(dump_default="")
71+
enabled = fields.Boolean(dump_default=False)
72+
disabled_at = fields.Integer(dump_default=0)
73+
status = fields.String(dump_default="")
74+
error = fields.String(dump_default="")
75+
updated_at = fields.Integer(dump_default=0)
76+
created_at = fields.Integer(dump_default=0)
77+
78+
@pre_dump
79+
def process_data(self, data: Segment, **kwargs):
80+
return {
81+
"id": data.id,
82+
"document_id": data.document_id,
83+
"dataset_id": data.dataset_id,
84+
"position": data.position,
85+
"content": data.content,
86+
"keywords": data.keywords,
87+
"character_count": data.character_count,
88+
"token_count": data.token_count,
89+
"hit_count": data.hit_count,
90+
"hash": data.hash,
91+
"enabled": data.enabled,
92+
"disabled_at": datetime_to_timestamp(data.disabled_at),
93+
"status": data.status,
94+
"error": data.error,
95+
"updated_at": datetime_to_timestamp(data.updated_at),
96+
"created_at": datetime_to_timestamp(data.created_at),
97+
}
98+
99+
100+
class UpdateSegmentEnabledReq(FlaskForm):
101+
"""更新文档片段启用状态请求"""
102+
103+
enabled = BooleanField("enabled")
104+
105+
def validate_enabled(self, field: BooleanField) -> None:
106+
"""校验文档启用状态enabled"""
107+
if not isinstance(field.data, bool):
108+
raise ValidationError("enabled状态不能为空且必须为布尔值")
109+
110+
111+
class CreateSegmentReq(FlaskForm):
112+
"""创建文档片段请求结构"""
113+
114+
content = StringField("content", validators=[DataRequired("片段内容不能为空")])
115+
keywords = ListField("keywords")
116+
117+
def validate_keywords(self, field: ListField):
118+
"""校验关键词列表,涵盖长度不能为空,默认为值为空列表"""
119+
# 1.校验数据类型+非空
120+
if field.data is None:
121+
field.data = []
122+
if not isinstance(field.data, list):
123+
raise ValidationError("关键词列表格式必须是数组")
124+
125+
# 2.校验数据的长度,最长不能超过10个关键词
126+
if len(field.data) > 10:
127+
raise ValidationError("关键词长度范围数量在1-10")
128+
129+
# 3.循环校验关键词信息,关键词必须是字符串
130+
for keyword in field.data:
131+
if not isinstance(keyword, str):
132+
raise ValidationError("关键词必须是字符串")
133+
134+
# 4.删除重复数据并更新
135+
field.data = list(dict.fromkeys(field.data))
136+
137+
138+
class UpdateSegmentReq(FlaskForm):
139+
"""更新文档片段请求"""
140+
141+
content = StringField("content", validators=[DataRequired("片段内容不能为空")])
142+
keywords = ListField("keywords")
143+
144+
def validate_keywords(self, field: ListField):
145+
"""校验关键词列表,涵盖长度不能为空,默认为值为空列表"""
146+
# 1.校验数据类型+非空
147+
if field.data is None:
148+
field.data = []
149+
if not isinstance(field.data, list):
150+
raise ValidationError("关键词列表格式必须是数组")
151+
152+
# 2.校验数据的长度,最长不能超过10个关键词
153+
if len(field.data) > 10:
154+
raise ValidationError("关键词长度范围数量在1-10")
155+
156+
# 3.循环校验关键词信息,关键词必须是字符串
157+
for keyword in field.data:
158+
if not isinstance(keyword, str):
159+
raise ValidationError("关键词必须是字符串")
160+
161+
# 4.删除重复数据并更新
162+
field.data = list(dict.fromkeys(field.data))

internal/service/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
from .indexing_service import IndexingService
1313
from .keyword_table_service import KeywordTableService
1414
from .process_rule_service import ProcessRuleService
15+
from .segment_service import SegmentService
1516

1617
__all__ = [
1718
"BuiltinToolService",
@@ -28,4 +29,5 @@
2829
"IndexingService",
2930
"KeywordTableService",
3031
"ProcessRuleService",
32+
"SegmentService",
3133
]

0 commit comments

Comments
 (0)