Python/agent/application/preprocessing.py

404 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- coding: utf-8 -*-
"""
预处理:
获取飞书知识库节点列表,并就文档节点与预处理任务比照,若:
1该文档不存在则下载该文档MD文件然后分块、向量化、存储到数据库中
2该文档存在但更新时间戳不一致则先再数据库中删除该文档所有块向量记录再重新分块、向量化、存储到数据库中
"""
from pathlib import Path
from typing import Dict, List, Tuple, cast
from langchain_text_splitters import (
MarkdownHeaderTextSplitter,
RecursiveCharacterTextSplitter,
)
from lark_oapi import Client
from lark_oapi.api.docs.v1 import (
GetContentRequest,
GetContentResponse,
GetContentResponseBody,
)
from lark_oapi.api.wiki.v2 import (
ListSpaceNodeRequest,
ListSpaceNodeResponse,
ListSpaceNodeResponseBody,
Node,
)
from pgvector.psycopg import register_vector
from psycopg import connect
from pydantic import BaseModel, Field
from pydantic_ai import Embedder
from pydantic_ai.embeddings.sentence_transformers import (
SentenceTransformerEmbeddingModel,
SentenceTransformersEmbeddingSettings,
)
# 建立数据库连接
db_connection = connect(
host="liubiren.cloud",
port=5201,
user="root",
password="PG198752",
dbname="database",
)
register_vector(db_connection)
# 实例化飞书服务端
feishu_server_client = (
Client.builder()
.app_id("cli_a1587980be78500c")
.app_secret("vZXGZomwfmyaHXoG8s810d1YYGLsIqCA")
.build()
)
# 实例化词嵌器
embedder = Embedder(
SentenceTransformerEmbeddingModel(
str(Path(__file__).parent / "models" / "paraphrase-multilingual-MiniLM-L12-v2"),
settings=SentenceTransformersEmbeddingSettings(
sentence_transformers_device="cpu",
sentence_transformers_normalize_embeddings=True,
),
),
)
# 最大分块大小
max_chunk_size: int = cast(int, embedder.max_input_tokens_sync())
# 分块重叠
chunk_overlap = max(20, min(int(round(max_chunk_size * 0.15)), 60))
# 实例化 Markdown 标题分割器
header_splitter = MarkdownHeaderTextSplitter(
headers_to_split_on=[
("#", "h1"),
("##", "h2"),
("###", "h3"),
("####", "h4"),
("#####", "h5"),
],
strip_headers=True, # 不保留标题
)
class Chunk(BaseModel):
"""
分块数据模型
"""
document_id: str = Field(..., description="文档唯一标识")
idx: int = Field(..., description="分块索引")
heading_chunk_h1: str = Field(..., description="标题块标题1")
heading_chunk_h2: str = Field(default="", description="标题块标题2")
heading_chunk_h3: str = Field(default="", description="标题块标题3")
heading_chunk_h4: str = Field(default="", description="标题块标题4")
heading_chunk_h5: str = Field(default="", description="标题块标题5")
heading_chunk_heading: str = Field(..., description="标题块标题")
heading_chunk_content: str = Field(..., description="标题块内容")
content_chunk_content: str = Field(..., description="内容块内容")
content_chunk_size: int = Field(..., description="内容块大小")
embedding_text: str = Field(..., description="词嵌文本")
def get_online_documents() -> Dict[str, int]:
"""
获取在线文档字典
:return: 在线文档字典
"""
# 在线文档字典
online_documents: Dict[str, int] = {}
# 实例化知识库
if not (wiki := feishu_server_client.wiki):
raise Exception("client.wiki is None")
# 分页标记
page_token = ""
while True:
# 构造获取知识库空间子节点列表请求实例
request: ListSpaceNodeRequest = (
ListSpaceNodeRequest.builder()
.space_id("7615153684095257820") # 默认为产品设计知识库
.page_size(50) # 分页大小
.page_token(page_token)
.build()
)
# 请求
response: ListSpaceNodeResponse = wiki.v2.space_node.list(request)
if not response.success():
raise Exception("response not success")
# 响应数据
data = cast(ListSpaceNodeResponseBody, response.data)
for item in cast(List[Node], data.items):
if cast(str, item.obj_type) == "docx":
# 文档唯一标识
id: str = cast(str, item.obj_token)
# 文档更新时间戳
updated_at: int = int(
cast(int, item.obj_edit_time)
) # Lark说明类型为 int但实际为 str
online_documents[id] = updated_at
# 更新分页标记
page_token = cast(str, data.page_token)
# 若是否还有更多项为否则跳出循环
if not cast(bool, data.has_more):
break
return online_documents
def get_offline_documents() -> Dict[str, int]:
"""
获取离线文档字典
:return: 离线文档字典
"""
# 离线文档字典
offline_documents: Dict[str, int] = {}
with db_connection.cursor() as cursor:
cursor.execute(
"""
SELECT id, updated_at
FROM documents
"""
)
for id, updated_at in cursor.fetchall():
offline_documents[id] = updated_at
return offline_documents
def download_document(id: str) -> str:
"""
下载文档
:param id: 文档唯一标识
:return: 文档内容Markdown格式字符串
"""
# 实例化云文档
if not (docs := feishu_server_client.docs):
raise Exception("client.docs is None")
# 构造获取云文档内容请求实例
request: GetContentRequest = (
GetContentRequest.builder()
.doc_token(id)
.doc_type("docx")
.content_type("markdown")
.build()
)
# 请求
response: GetContentResponse = docs.v1.content.get(request)
if not response.success():
print(response.code, response.msg)
raise Exception("response not success")
# 响应数据
data = cast(GetContentResponseBody, response.data)
return cast(str, data.content)
def sync_documents() -> None:
"""
同步文档
:return: None
"""
# 获取在线文档字典
online_documents = get_online_documents()
# 获取离线文档字典
offline_documents = get_offline_documents()
# 待删除文档唯一标识集合
pending_delete_documents = set(offline_documents.keys() - online_documents.keys())
# 待下载文档唯一标识字典
pending_download_documents: Dict[str, int] = {}
for document_id, document_updated_at in online_documents.items():
# 若在线文档唯一标识于离线文档字典不存在则将该文档唯一标识添加至待下载文档唯一标识字典
if document_id not in offline_documents:
pending_download_documents[document_id] = document_updated_at
continue
# 若在线文档唯一标识于离线文档字典存在且更新时间戳不一致则将该文档唯一标识添加至待删除/下载文档唯一标识集合
if (
document_id in offline_documents
and document_updated_at != offline_documents[document_id]
):
pending_delete_documents.add(document_id)
pending_download_documents[document_id] = document_updated_at
continue
# 批量删除待删除文档
if pending_delete_documents:
with db_connection.cursor() as cursor:
cursor.execute(
"""
DELETE FROM documents
WHERE id = ANY(%s);
""",
(list(pending_delete_documents),),
)
db_connection.commit()
for document_id, document_updated_at in pending_download_documents.items():
# 下载文档
document_content = download_document(document_id)
print(document_content)
exit()
chunks = []
# 分块索引
idx = 0
embedding_texts = []
for heading_chunk in header_splitter.split_text(document_content):
# 若标题块内容为空则跳过
if not (heading_chunk_content := heading_chunk.page_content.strip()):
continue
# 构建标题块标题
heading_chunk_heading = f"# {(heading_chunk_h1:=heading_chunk.metadata["h1"])}\n" # 必定包含 h1 标题
if heading_chunk_h2 := heading_chunk.metadata.get("h2", ""):
heading_chunk_heading += f"## {heading_chunk_h2}\n"
if heading_chunk_h3 := heading_chunk.metadata.get("h3", ""):
heading_chunk_heading += f"### {heading_chunk_h3}\n"
if heading_chunk_h4 := heading_chunk.metadata.get("h4", ""):
heading_chunk_heading += f"#### {heading_chunk_h4}\n"
if heading_chunk_h5 := heading_chunk.metadata.get("h5", ""):
heading_chunk_heading += f"##### {heading_chunk_h5}\n"
# 最大标题块内容大小
max_heading_chunk_content_size = (
max_chunk_size - embedder.count_tokens_sync(heading_chunk_heading)
)
# 若标题块内容大小超过最大标题块内容大小则再次分割
if (
embedder.count_tokens_sync(heading_chunk_content)
> max_heading_chunk_content_size
):
# 实例化递归分割器
recursive_splitter = RecursiveCharacterTextSplitter(
chunk_size=max_heading_chunk_content_size,
chunk_overlap=chunk_overlap, # 分块重叠:最大分块大小的 15%,且最小为 20、最大为 60
length_function=embedder.count_tokens_sync,
separators=["\n\n", "\n", "", "", "", ".", "?", "!", ""],
)
content_chunk_contents = recursive_splitter.split_text(
heading_chunk_content
)
else:
content_chunk_contents = [heading_chunk_content]
for content_chunk_content in content_chunk_contents:
# 若内容块为空则跳过
if not content_chunk_content.strip():
continue
chunks.append(
Chunk(
idx=idx,
document_id=document_id,
heading_chunk_h1=heading_chunk_h1,
heading_chunk_h2=heading_chunk_h2,
heading_chunk_h3=heading_chunk_h3,
heading_chunk_h4=heading_chunk_h4,
heading_chunk_h5=heading_chunk_h5,
heading_chunk_heading=heading_chunk_heading,
heading_chunk_content=heading_chunk_content,
content_chunk_content=content_chunk_content,
content_chunk_size=embedder.count_tokens_sync(
content_chunk_content
),
embedding_text=(
embedding_text := heading_chunk_heading
+ content_chunk_content
),
)
)
idx += 1
embedding_texts.append(embedding_text)
with db_connection.cursor() as cursor:
# 写入文档表
cursor.execute(
"""
INSERT INTO documents (id, updated_at)
VALUES (%s, %s)
""",
(document_id, document_updated_at),
)
# 写入分块表
cursor.executemany(
"""
INSERT INTO chunks (
document_id, idx,
heading_chunk_h1, heading_chunk_h2, heading_chunk_h3, heading_chunk_h4, heading_chunk_h5, heading_chunk_heading, heading_chunk_content,
content_chunk_content, content_chunk_size,
embedding_text, embedding_vector
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s);
""",
[
(
chunk.document_id,
chunk.idx,
chunk.heading_chunk_h1,
chunk.heading_chunk_h2,
chunk.heading_chunk_h3,
chunk.heading_chunk_h4,
chunk.heading_chunk_h5,
chunk.heading_chunk_heading,
chunk.heading_chunk_content,
chunk.content_chunk_content,
chunk.content_chunk_size,
chunk.embedding_text,
embedding_vector,
)
for chunk, embedding_vector in zip(
chunks,
embedder.embed_documents_sync(
embedding_texts
).embeddings, # 生成词嵌向量
)
], # 构建记录
)
db_connection.commit()
def query(
question: str,
) -> List[Tuple[float, str, str]]:
"""
查询
:param question: 问题
:return: 相似度最高的前 k 个分块(每个分块包含相似度、文档 ID、标题、内容
"""
# 生成词嵌向量
embedding_vector = embedder.embed_query_sync(query=question).embeddings[0]
with db_connection.cursor() as cursor:
cursor.execute(
"""
SELECT
1 - (embedding_vector <=> %s::vector) AS similarity_score,
heading_chunk_heading,
heading_chunk_content
FROM chunks
WHERE 1 - (embedding_vector <=> %s::vector) >= %s
ORDER BY embedding_vector <=> %s::vector
LIMIT %s;
""",
(
embedding_vector,
embedding_vector,
0.45,
embedding_vector,
5,
),
)
rows = cursor.fetchall()
return [
(similarity_score, heading_chunk_heading, heading_chunk_content)
for similarity_score, heading_chunk_heading, heading_chunk_content in rows
]
if __name__ == "__main__":
# 同步文档
sync_documents()
# print(query("需求评审有哪些流程"))