part3知识库RAG方案设计


前置准备:知识库得需求、场景、价值分析

知识库的这个背景就像是我们在一个团队中,文档信息不统一,我们就用一个库,把需要的所有知识全部放进去,然后进行一个调用之类的

知识库的核心流程

第一步:文档拆分
知识库就是把我们上传的文档,按照章节、段落、甚至自定义的规则进行划分

每个拆出来的片段,都要包含完整的语义信息

什么意思?举个例子,假设文档里有一段话:”当API返回401错误码时,通常是因为Token过期。解
决方案是重新调用/auth/refresh接口获取新Token。
如果拆分的时候,把”解决方案是重新调用/auth/refresh接口”和前面的”API返回401错误码”拆到了两个不同的片段里,那后续AI检索的时候,就可能只找到问题描述,却找不到解决方案。这就是”语义信息不完整”带来的问题。

好的拆分,是整个知识库质量的基石

第二步:文本向量化
调用一个叫Embedding模型的东西,把每个文本片段转换为一组高维向量

比如一个人,每个人都有很多特征:身高、体重、年龄等等,我们用一组数字来表示就是[175,70,25],那这组数字就是这个人的“向量表示”

同样的道理。Embedding模型会根据文本的语义,给每个片段生成一组数字(通常是几百到上千维的)。语义越相近的版本,它们的向量在数学空间中就越靠近

第三步:向量库储存
向量生成好了,总得有个地方存起来?这就是向量数据库的职责
知识库会把生成的向量,连同它的原信息(比如:这段内容来自哪份文档、属于第几章、作者是谁、最后更新时间是什么),一起批量存入向量数据库中
为什么要存元信息?
1.方便溯源:当AI给你一个回答的时候,你可以看到“这个答案来自。。。”,这样你就能判断这个答案靠不靠谱,而不是盲注信任AI
2.支持过滤:比如你只想搜索最近一个月更新的文档,或者只搜索某个特定作者写的内容,有了元信息就能轻松做到

然后就是知识库的一个核心价值了,当然有很多,比如找东西就没有那么复杂了等等

架构设计RAG全流程解析

为什么需要RAG

就是做一个检索功能,不然一下内容太多了的话,上下文窗口有限,可能不支持,成本高等等问题

RAG核心流程

提问前有个数据准备:分片-Embedding (嵌入向量)-存储
提问后回答生成:召回-重排-生成

提问前做的事情,本质上就是把你的文档变成一个AI能理解的知识库。提问后做的事情,就是从知识库里面找答案,然后让AI组织语言回答你

第一阶段

这个阶段的目标只有一个:把你的原始文档,变成一个可以被快速检索的知识库
整个过程分三步:分片、向量化、存储
分片:提前把文档切成一段一段的小片段,每个片段聚焦一个具体的知识点
1.按固定数切
2.按照段落切
3.按章节/标题切
4.一页一个片段

向量化:
通过把文本向量化,来计算向量之间的距离,来判断两端文字的意思是否相近。这就是RAG能够“检索相关内容”的数学基础

存储:
把片段和向量存进数据库
这份数据库就是向量数据库
向量数据库存的不只是向量,而是向量+原始文本

1
2
3
4
{
"content": "产品A支持7天无理由退货,拆封后的电子产品除外。",
"vector": [0.12, 0.34, -0.56, 0.78, ...]
}

第二阶段,提问后的回答生成

召回——从知识库里“广撒网”
召回阶段做的事情,和提问前的索引过程是对称的
先把用户的问题也通过Embedding 模型转成向量
拿这个问题,去向量数据库里找最相似的片段

重排:
这步简单来说就是把问题以大模型的方式弄懂,或者理解为能够经过检索后回答的更加好嘛
原理还是有点复杂的

生成:
大模型拿到这些信息后,就能基于真实的知识内容来组织语言,生成回答,而不是凭自己的想象来回答

能够减少幻觉、成本更低、速度更快、准确率更高

RAG进阶

1
2
3
4
5
6
7
8
9
Agent 判断是否需要知识库
→ 校验 query、topK、过滤器和知识库权限
→ 并行召回
├─ query embedding → Milvus 语义召回 Top 20
└─ scoped chunks → BM25L 词项召回 Top 20
→ owner / tenant / 文档 / metadata 二次过滤
→ RRF 融合并去重,保留 Top 20
→ qwen3-vl-rerank 真实精排
→ 返回最多 5 条结果及对应引用

感觉好复杂

一句话概述:这次主要升级的不是RAG的生成,而是检索中间层-召回更全面、排序更准确、权限更严格、结果更可解释

源码分析-RAG代码实战

我先把源码下下来

放到

1
D:\typora\document\大模型\大模型项目

先放到这里来

流程梳理

我们的目标是将文件向量化后存储到数据库中,具体步骤如下

1.读取文件

2.切分文件

3.索引(向量化和储存)

读取文件

我们直接传入文件路径,读取文件到内存中

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
"""向量索引服务模块"""

from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Optional

from loguru import logger

from app.services.document_splitter_service import document_splitter_service
from app.services.vector_store_manager import vector_store_manager


class IndexingResult:
"""索引结果类"""

def __init__(self):
self.success = False
self.directory_path = ""
self.total_files = 0
self.success_count = 0
self.fail_count = 0
self.start_time: Optional[datetime] = None
self.end_time: Optional[datetime] = None
self.error_message = ""
self.failed_files: Dict[str, str] = {}

def increment_success_count(self):
"""增加成功计数"""
self.success_count += 1

def increment_fail_count(self):
"""增加失败计数"""
self.fail_count += 1

def add_failed_file(self, file_path: str, error: str):
"""添加失败文件"""
self.failed_files[file_path] = error

def get_duration_ms(self) -> int:
"""获取耗时(毫秒)"""
if self.start_time and self.end_time:
return int((self.end_time - self.start_time).total_seconds() * 1000)
return 0

def to_dict(self) -> Dict[str, Any]:
"""转换为字典"""
return {
"success": self.success,
"directory_path": self.directory_path,
"total_files": self.total_files,
"success_count": self.success_count,
"fail_count": self.fail_count,
"duration_ms": self.get_duration_ms(),
"error_message": self.error_message,
"failed_files": self.failed_files,
}


class VectorIndexService:
"""向量索引服务 - 负责读取文件、生成向量、存储到 Milvus"""

def __init__(self):
"""初始化向量索引服务"""
self.upload_path = "./uploads"
logger.info("向量索引服务初始化完成")

def index_directory(self, directory_path: Optional[str] = None) -> IndexingResult:
"""
索引指定目录下的所有文件

Args:
directory_path: 目录路径(可选,默认使用配置的上传目录)

Returns:
IndexingResult: 索引结果
"""
result = IndexingResult()
result.start_time = datetime.now()

try:
# 使用指定目录或默认上传目录
target_path = directory_path if directory_path else self.upload_path
dir_path = Path(target_path).resolve()

if not dir_path.exists() or not dir_path.is_dir():
raise ValueError(f"目录不存在或不是有效目录: {target_path}")

result.directory_path = str(dir_path)

# 获取所有支持的文件
files = list(dir_path.glob("*.txt")) + list(dir_path.glob("*.md"))

if not files:
logger.warning(f"目录中没有找到支持的文件: {target_path}")
result.total_files = 0
result.success = True
result.end_time = datetime.now()
return result

result.total_files = len(files)
logger.info(f"开始索引目录: {target_path}, 找到 {len(files)} 个文件")

# 遍历并索引每个文件
for file_path in files:
try:
self.index_single_file(str(file_path))
result.increment_success_count()
logger.info(f"✓ 文件索引成功: {file_path.name}")
except Exception as e:
result.increment_fail_count()
result.add_failed_file(str(file_path), str(e))
logger.error(f"✗ 文件索引失败: {file_path.name}, 错误: {e}")

result.success = result.fail_count == 0
result.end_time = datetime.now()

logger.info(
f"目录索引完成: 总数={result.total_files}, "
f"成功={result.success_count}, 失败={result.fail_count}"
)

return result

except Exception as e:
logger.error(f"索引目录失败: {e}")
result.success = False
result.error_message = str(e)
result.end_time = datetime.now()
return result

def index_single_file(self, file_path: str):
"""
索引单个文件 (使用新的 LangChain 分割器)

Args:
file_path: 文件路径

Raises:
ValueError: 文件不存在时抛出
RuntimeError: 索引失败时抛出
"""
path = Path(file_path).resolve()

if not path.exists() or not path.is_file():
raise ValueError(f"文件不存在: {file_path}")

logger.info(f"开始索引文件: {path}")

try:
# 1. 读取文件内容
content = path.read_text(encoding="utf-8")
logger.info(f"读取文件: {path}, 内容长度: {len(content)} 字符")

# 2. 删除该文件的旧数据(如果存在)
normalized_path = path.as_posix()
vector_store_manager.delete_by_source(normalized_path)

# 3. 使用新的文档分割器
documents = document_splitter_service.split_document(content, normalized_path)
logger.info(f"文档分割完成: {file_path} -> {len(documents)} 个分片")

# 4. 添加文档到向量存储
if documents:
vector_store_manager.add_documents(documents)
logger.info(f"文件索引完成: {file_path}, 共 {len(documents)} 个分片")
else:
logger.warning(f"文件内容为空或无法分割: {file_path}")

except Exception as e:
logger.error(f"索引文件失败: {file_path}, 错误: {e}")
raise RuntimeError(f"索引文件失败: {e}") from e


# 全局单例
vector_index_service = VectorIndexService()

然后index_single_file是索引单个文件的入口方法

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
def index_single_file(self, file_path: str):
"""
索引单个文件 (使用新的 LangChain 分割器)

Args:
file_path: 文件路径

Raises:
ValueError: 文件不存在时抛出
RuntimeError: 索引失败时抛出
"""
path = Path(file_path).resolve()

if not path.exists() or not path.is_file():
raise ValueError(f"文件不存在: {file_path}")

logger.info(f"开始索引文件: {path}")

try:
# 1. 读取文件内容
content = path.read_text(encoding="utf-8")
logger.info(f"读取文件: {path}, 内容长度: {len(content)} 字符")

# 2. 删除该文件的旧数据(如果存在)
normalized_path = path.as_posix()
vector_store_manager.delete_by_source(normalized_path)

# 3. 使用新的文档分割器
documents = document_splitter_service.split_document(content, normalized_path)
logger.info(f"文档分割完成: {file_path} -> {len(documents)} 个分片")

# 4. 添加文档到向量存储
if documents:
vector_store_manager.add_documents(documents)
logger.info(f"文件索引完成: {file_path}, 共 {len(documents)} 个分片")
else:
logger.warning(f"文件内容为空或无法分割: {file_path}")

except Exception as e:
logger.error(f"索引文件失败: {file_path}, 错误: {e}")
raise RuntimeError(f"索引文件失败: {e}") from e

文件分块

分档分块使用langChain提供的分割器,分为三个阶段

第一阶段:按markdown标题(#、##)切分,将文档分割成多个章节

第二阶段:对每个章节使用RecursivecharacterTextSplitter进行二次分割,超过chunk_size*2的章节会被拆分

第三阶段:合并过小的分片(<300字符),避免过度碎片化。同时通过chunk_overlap保持分片间的上下文语义连贯

DocumentSplitterService初始化时会配置好这两个分割器

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
class DocumentSplitterService:
def __init__(self):
self.chunk_size = config.chunk_max_size # 默认 800
self.chunk_overlap = config.chunk_overlap # 默认 100

# 第一阶段: Markdown 标题分割器 (按 # 和 ## 切分)
self.markdown_splitter = MarkdownHeaderTextSplitter(
headers_to_split_on=[
("#", "h1"),
("##", "h2"),
],
strip_headers=False, # 保留标题在内容中
)

# 第二阶段: 递归字符分割器 (用于二次分割)
self.text_splitter = RecursiveCharacterTextSplitter(
chunk_size=self.chunk_size * 2, # 加倍 chunk_size, 减少分片数
chunk_overlap=self.chunk_overlap,
length_function=len,
is_separator_regex=False,
)

Markdown文档完整的三阶段分割逻辑

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
def split_markdown(self, content: str, file_path: str = "") -> List[Document]:
"""分割 Markdown 文档(两阶段分割 + 合并小片段)"""
# 第一阶段: 按标题分割
md_docs = self.markdown_splitter.split_text(content)

# 第二阶段: 按大小进一步分割
docs_after_split = self.text_splitter.split_documents(md_docs)

# 第三阶段: 合并太小的分片 (< 300 字符)
final_docs = self._merge_small_chunks(docs_after_split, min_size=300)

# 添加文件路径元数据
for doc in final_docs:
doc.metadata["_source"] = file_path
doc.metadata["_extension"] = ".md"
doc.metadata["_file_name"] = Path(file_path).name

logger.info(f"Markdown 分割完成: {file_path} -> {len(final_docs)} 个分片")
return final_docs

合并小分片的逻辑(_merge_small_chunks):遍历所有分片,若当前分片小于min_size且合并后不超限,则将追加到上一个分片

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
def _merge_small_chunks(self, documents: List[Document], min_size: int = 300) -> List[Document]:
merged_docs = []
current_doc = None

for doc in documents:
doc_size = len(doc.page_content)

if current_doc is None:
current_doc = doc
elif doc_size < min_size and len(current_doc.page_content) < self.chunk_size:
# 当前分片太小且合并后不会太大,则合并
current_doc.page_content += "\n\n" + doc.page_content
else:
# 保存当前文档,开始新文档
merged_docs.append(current_doc)
current_doc = doc

if current_doc is not None:
merged_docs.append(current_doc)

return merged_docs

文件索引(向量化和存储到数据库)

Embedding生成

DashScopeEmbeddings 实现了 LangChain 标准的 Embeddings 接口,通过阿里云 DashScope 的 OpenAI 兼容模式调用 text-embedding-v4 模型,生成 1024 维向量

1
2
3
4
5
6
7
8
9
class DashScopeEmbeddings(Embeddings):
def __init__(self, api_key: str, model: str = "text-embedding-v4", dimensions: int):
self.client = OpenAI(
api_key=api_key,
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1"
)
self.model = model
self.dimensions = dimensions

批量向量化文档(embed_documents)和单挑查询向量化(embed_query)分别对应入库和检索场景

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
def embed_documents(self, texts: List[str]) -> List[List[float]]:
"""批量嵌入文档列表,返回向量列表"""
response = self.client.embeddings.create(
model=self.model,
input=texts,
dimensions=self.dimensions,
encoding_format="float"
)
return [item.embedding for item in response.data]

def embed_query(self, text: str) -> List[float]:
"""嵌入单个查询文本,返回单条向量"""
response = self.client.embeddings.create(
model=self.model,
input=text,
dimensions=self.dimensions,
encoding_format="float"
)
return response.data[0].embedding

向量存储到Milvus

VectorStoreManager封装了Langchain_milvus.Molvus,将LangChain Document对象直接批量写入Milvus,字段映射关系如下

LangChain 字段 Milvus Collection 字段 说明
page_content content 文本内容
向量(自动计算) vector 1024 维 float 向量
id(UUID) id 主键
metadata metadata JSON 元数据(含 _source_file_name 等)

初始化时连接Milvus

1
2
3
4
5
6
7
8
9
10
11
12
self.vector_store = Milvus(
embedding_function=vector_embedding_service, # 自动调用 embed_documents
collection_name="biz",
connection_args={"host": config.milvus_host, "port": config.milvus_port},
auto_id=False, # 使用自定义 UUID
drop_old=False,
text_field="content",
vector_field="vector",
primary_field="id",
metadata_field="metadata",
)

批量入库时,LangChain会自动调用embed_documents完成向量化,无需手动循环处理每个分片

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
def add_documents(self, documents: List[Document]) -> List[str]:
"""批量添加文档到向量存储(自动批量向量化)"""
import time, uuid

start_time = time.time()

# 为每个文档生成唯一 UUID (auto_id=False 时必须手动提供)
ids = [str(uuid.uuid4()) for _ in documents]
# LangChain Milvus 的 add_documents 自动调用 embedding_function 批量向量化并写入
result_ids = self.vector_store.add_documents(documents, ids=ids)

elapsed = time.time() - start_time
logger.info(
f"批量添加 {len(documents)} 个文档完成,"
f"耗时: {elapsed:.2f}秒,平均: {elapsed/len(documents):.2f}秒/个"
)
return result_ids

在重新索引同一个文件前,会先按照_source路径删除旧数据

1
2
3
4
5
6
7
8
9
10
def delete_by_source(self, file_path: str) -> int:
"""删除指定文件的所有文档"""
collection = milvus_manager.get_collection()
# metadata 是 JSON 字段, 使用 JSON 路径查询语法
expr = f'metadata["_source"] == "{file_path}"'
result = collection.delete(expr)
deleted_count = result.delete_count if hasattr(result, "delete_count") else 0
logger.info(f"删除文件旧数据: {file_path}, 删除数量: {deleted_count}")
return deleted_count

短暂小结

上面就是agent的一个上半部分,也就是我们的这个知识索引部分

接下来就是要实现这个agent检索好这个数据库然后完成回答问题的这个部分了

召回

我们之前已经实现将文档向量化存储到了Milvus,所以召回时也从这个数据库去查询

1.将查询文本向量化

2.相似度查询

sear_similar_documents是底层召回的完整实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
def search_similar_documents(self, query: str, top_k: int = 3) -> List[SearchResult]:
"""
搜索相似文档

Args:
query: 查询文本
top_k: 返回最相似的k个结果

Returns:
List[SearchResult]: 搜索结果列表
"""
logger.info(f"开始搜索相似文档,查询: {query}, topK: {top_k}")
# 1. 将查询文本向量化
query_vector = vector_embedding_service.embed_query(query)
logger.debug(f"查询向量生成成功,维度: {len(query_vector)}")
# 2. 获取 collection
collection: Collection = milvus_manager.get_collection()

# 3. 构建搜索参数
search_params = {
"metric_type": "L2", # 欧氏距离, 与入库时的索引类型保持一致
"params": {"nprobe": 10},
}

# 4. 执行搜索
results = collection.search(
data=[query_vector],
anns_field="vector",
param=search_params,
limit=top_k,
output_fields=["id", "content", "metadata"],
)

# 5. 解析搜索结果
search_results = []
for hits in results:
for hit in hits:
result = SearchResult(
id=hit.entity.get("id"),
content=hit.entity.get("content"),
score=hit.distance, # L2 距离,越小越相似
metadata=hit.entity.get("metadata", {}),
)
search_results.append(result)

logger.info(f"搜索完成,找到 {len(search_results)} 个相似文档")
return search_results

查询文本向量化

首先是对用户问题向量化,调用DashScopeEmbeddings.embed_query,通过DashScope OpenAI兼容接口获取1024维向量

1
query_vector = vector_embedding_service.embed_query(query)

embed_query的实现

1
2
3
4
5
6
7
8
9
10
def embed_query(self, text: str) -> List[float]:
"""嵌入单个查询文本, 返回单条向量"""
response = self.client.embeddings.create(
model=self.model, # text-embedding-v4
input=text,
dimensions=self.dimensions, # 1024
encoding_format="float"
)
return response.data[0].embedding

构建搜索参数并执行向量检索

然后使用 PyMilvus 的 collection.search 进行相似度查询,获取距离最近的向量数据

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# 构建搜索参数
search_params = {
"metric_type": "L2", # 欧氏距离
"params": {"nprobe": 10}, # 搜索时探测的 cluster 数量, 越大越精确但越慢
}

# 执行搜索
results = collection.search(
data=[query_vector], # 查询向量 (批量, 这里只有一条)
anns_field="vector", # 向量字段名, 与入库时一致
param=search_params,
limit=top_k, # 返回最相似的前 K 条
output_fields=["id", "content", "metadata"], # 需要返回的字段
)

搜索结果封装到 SearchResult 对象中,score 为 L2 欧氏距离,越小表示越相似:

1
2
3
4
5
6
7
class SearchResult:
def __init__(self, id: str, content: str, score: float, metadata: Dict[str, Any]):
self.id = id
self.content = content
self.score = score # L2 距离, 越小越相似
self.metadata = metadata

源码分析:API接口与Agent的整合

文件上传接口的定义

这里就是很简单的一个原理,就不多说了

文件上传接口的核心实现

除了特有的这种对文件进行向量化索引存储,其他应该就是常规的文件上传的一个操作了

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
"""文件上传接口模块"""

from pathlib import Path

from fastapi import APIRouter, File, HTTPException, UploadFile
from fastapi.responses import JSONResponse

from app.services.vector_index_service import vector_index_service
from loguru import logger

router = APIRouter()

# 文件上传后存储的路径
UPLOAD_DIR = Path("./uploads")
# 支持的文件类型
ALLOWED_EXTENSIONS = ["txt", "md"]
# 单个文件支持最大大小
MAX_FILE_SIZE = 10 * 1024 * 1024 # 10MB


@router.post("/upload")
async def upload_file(file: UploadFile = File(...)):
"""
上传文件并自动创建向量索引

Args:
file: 上传的文件

Returns:
JSONResponse: 上传结果
"""
try:
# 1. 验证文件
if not file.filename:
raise HTTPException(status_code=400, detail="文件名不能为空")

# 2. 规范化文件名(去除空格,处理 Windows 上传的文件)
safe_filename = _sanitize_filename(file.filename)

# 3. 验证文件扩展名
file_extension = _get_file_extension(safe_filename)
if file_extension not in ALLOWED_EXTENSIONS:
raise HTTPException(
status_code=400,
detail=f"不支持的文件格式,仅支持: {', '.join(ALLOWED_EXTENSIONS)}",
)

# 4. 创建上传目录
UPLOAD_DIR.mkdir(parents=True, exist_ok=True)

# 5. 保存文件
file_path = UPLOAD_DIR / safe_filename

# 如果文件已存在,先删除旧文件(实现覆盖更新)
if file_path.exists():
logger.info(f"文件已存在,将覆盖: {file_path}")
file_path.unlink()

# 读取并保存文件内容
content = await file.read()

# 验证文件大小
if len(content) > MAX_FILE_SIZE:
raise HTTPException(status_code=400, detail=f"文件大小超过限制(最大 {MAX_FILE_SIZE} 字节)")

file_path.write_bytes(content)

logger.info(f"文件上传成功: {file_path}")

# 5. 自动创建向量索引
try:
logger.info(f"开始为上传文件创建向量索引: {file_path}")
vector_index_service.index_single_file(str(file_path))
logger.info(f"向量索引创建成功: {file_path}")
except Exception as e:
logger.error(f"向量索引创建失败: {file_path}, 错误: {e}")
# 注意:即使索引失败,文件上传仍然成功,只是记录错误日志

# 6. 返回响应
return JSONResponse(
status_code=200,
content={
"code": 200,
"message": "success",
"data": {
"filename": safe_filename,
"file_path": str(file_path),
"size": len(content),
},
},
)

except HTTPException:
raise
except Exception as e:
logger.error(f"文件上传失败: {e}")
raise HTTPException(status_code=500, detail=f"文件上传失败: {e}")


@router.post("/index_directory")
async def index_directory(directory_path: str = None):
"""
索引指定目录下的所有文件

Args:
directory_path: 目录路径(可选,默认使用 uploads 目录)

Returns:
JSONResponse: 索引结果
"""
try:
logger.info(f"开始索引目录: {directory_path or 'uploads'}")

# 执行索引
result = vector_index_service.index_directory(directory_path)

return JSONResponse(
status_code=200,
content={
"code": 200,
"message": "success" if result.success else "partial_success",
"data": result.to_dict(),
},
)

except Exception as e:
logger.error(f"索引目录失败: {e}")
raise HTTPException(status_code=500, detail=f"索引目录失败: {e}")


def _get_file_extension(filename: str) -> str:
"""
获取文件扩展名

Args:
filename: 文件名

Returns:
str: 扩展名(小写,不含点)
"""
parts = filename.rsplit(".", 1)
if len(parts) == 2:
return parts[1].lower()
return ""


def _sanitize_filename(filename: str) -> str:
"""
规范化文件名,去除空格和特殊字符

Args:
filename: 原始文件名

Returns:
str: 规范化后的文件名
"""
# 去除空格
sanitized = filename.replace(" ", "_")
# 去除其他可能导致问题的字符
for char in ['\\', '/', ':', '*', '?', '"', '<', '>', '|']:
sanitized = sanitized.replace(char, "_")
return sanitized


文章作者: wuk0Ng
版权声明: 本博客所有文章除特別声明外,均采用 CC BY 4.0 许可协议。转载请注明来源 wuk0Ng !
评论
  目录