103 lines
3.1 KiB
Python
103 lines
3.1 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
预处理:
|
||
获取飞书知识库节点列表,并就文档节点与预处理任务比照,若:
|
||
1)该文档不存在则下载该文档MD文件,然后切块、向量化、存储到数据库中
|
||
2)该文档存在但更新时间戳不一致则先再数据库中删除该文档所有块向量记录,再重新切块、向量化、存储到数据库中
|
||
"""
|
||
|
||
import asyncio
|
||
from pathlib import Path
|
||
from pydantic_ai import Embedder
|
||
from pydantic_ai.embeddings.sentence_transformers import (
|
||
SentenceTransformerEmbeddingModel,
|
||
SentenceTransformersEmbeddingSettings,
|
||
)
|
||
from typing import cast, Dict, List, Any
|
||
|
||
from json import loads
|
||
|
||
from lark_oapi.api.wiki.service import WikiService
|
||
|
||
from lark_oapi import Client, JSON
|
||
from lark_oapi.api.wiki.v2 import (
|
||
ListSpaceNodeRequest,
|
||
ListSpaceNodeResponse,
|
||
ListSpaceNodeResponseBody,
|
||
Node,
|
||
)
|
||
|
||
|
||
# 实例化飞书服务端
|
||
server = (
|
||
Client.builder()
|
||
.app_id("cli_a1587980be78500c")
|
||
.app_secret("vZXGZomwfmyaHXoG8s810d1YYGLsIqCA")
|
||
.build()
|
||
)
|
||
# 实例化知识库
|
||
if not (wiki := server.wiki):
|
||
raise Exception("server.wiki is None")
|
||
|
||
# 在线文档唯一标识字典
|
||
online_document_ids: Dict[str, int] = {}
|
||
# 分页标记
|
||
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":
|
||
# 文档唯一标识
|
||
document_id: str = cast(str, item.obj_token)
|
||
# 文档更新时间戳
|
||
document_updated_at: int = cast(int, item.obj_edit_time)
|
||
online_document_ids[document_id] = document_updated_at
|
||
# 更新分页标记
|
||
page_token = cast(str, data.page_token)
|
||
# 若是否还有更多项为否则跳出循环
|
||
if not cast(bool, data.has_more):
|
||
break
|
||
|
||
print(online_document_ids)
|
||
|
||
exit()
|
||
|
||
|
||
async def main():
|
||
# 自动定位:脚本目录/models/模型文件夹
|
||
base_path = Path(__file__).parent
|
||
model_path = str(base_path / "models" / "paraphrase-multilingual-MiniLM-L12-v2")
|
||
|
||
model = SentenceTransformerEmbeddingModel(
|
||
model_path,
|
||
settings=SentenceTransformersEmbeddingSettings(
|
||
sentence_transformers_device="cpu",
|
||
sentence_transformers_normalize_embeddings=True,
|
||
),
|
||
)
|
||
embedder = Embedder(model)
|
||
|
||
query_text = "刘弼仁"
|
||
counts = await embedder.count_tokens(query_text)
|
||
print(counts)
|
||
max_tokens = await embedder.max_input_tokens()
|
||
print(f"Max tokens: {max_tokens}")
|
||
res = await embedder.embed_query(query_text)
|
||
print(f"向量维度:{len(res.embeddings[0])}")
|
||
|
||
|
||
asyncio.run(main())
|