Files
rag/storage/file_fetcher.py
lacerate551 100d1a06eb init: RAG 知识库服务初始提交
- 后端 API(Flask + Gunicorn)
- RAG 引擎(混合检索 + 云端 Reranker + 引用溯源)
- 文档解析(MinerU + 多格式支持)
- Docker 生产部署配置
- 排除前端项目、敏感配置、模型文件
2026-06-04 17:35:27 +08:00

167 lines
5.0 KiB
Python
Raw 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 -*-
"""
企业文件系统集成助手
提供便捷方法从企业文件系统获取文件并进行向量化
"""
import os
import tempfile
import logging
from typing import Optional, Tuple
logger = logging.getLogger(__name__)
def get_file_for_parsing(
file_path: str,
use_file_provider: bool = True
) -> Tuple[str, bool]:
"""
获取用于解析的文件路径
根据配置自动从本地或企业文件系统获取文件。
如果是企业文件系统,会下载到临时目录并返回临时路径。
Args:
file_path: 文件路径(相对路径或绝对路径)
use_file_provider: 是否使用文件提供者(默认 True
Returns:
(文件绝对路径, 是否是临时文件需要清理)
"""
# 如果是绝对路径且文件存在,直接返回
if os.path.isabs(file_path) and os.path.exists(file_path):
return file_path, False
# 如果配置了文件提供者,从企业文件系统获取
if use_file_provider:
try:
from storage import get_file_provider
provider = get_file_provider()
# 检查文件是否存在
if not provider.exists(file_path):
# 尝试本地 documents 目录
from config import DOCUMENTS_PATH
local_path = os.path.join(DOCUMENTS_PATH, file_path)
if os.path.exists(local_path):
return os.path.abspath(local_path), False
raise FileNotFoundError(f"文件不存在: {file_path}")
# 获取文件信息
info = provider.get_file_info(file_path)
logger.info(f"从企业文件系统获取文件: {file_path} ({info.size} bytes)")
# 对于大文件,使用流式下载
if info.size > 100 * 1024 * 1024: # > 100MB
logger.info(f"大文件,使用流式下载: {file_path}")
stream = provider.get_file_stream(file_path)
temp_file = tempfile.NamedTemporaryFile(
delete=False,
suffix=os.path.splitext(file_path)[1]
)
try:
while True:
chunk = stream.read(8192)
if not chunk:
break
temp_file.write(chunk)
finally:
stream.close()
temp_file.close()
return temp_file.name, True
else:
# 小文件直接下载
content = provider.get_file(file_path)
temp_file = tempfile.NamedTemporaryFile(
delete=False,
suffix=os.path.splitext(file_path)[1]
)
temp_file.write(content)
temp_file.close()
return temp_file.name, True
except ImportError:
logger.warning("文件提供者模块未安装,使用本地文件系统")
# 降级到本地文件系统
from config import DOCUMENTS_PATH
local_path = os.path.join(DOCUMENTS_PATH, file_path)
if os.path.exists(local_path):
return os.path.abspath(local_path), False
raise FileNotFoundError(f"文件不存在: {file_path}")
def cleanup_temp_file(file_path: str, is_temp: bool):
"""
清理临时文件
Args:
file_path: 文件路径
is_temp: 是否是临时文件
"""
if is_temp and os.path.exists(file_path):
try:
os.remove(file_path)
logger.debug(f"清理临时文件: {file_path}")
except Exception as e:
logger.warning(f"清理临时文件失败: {e}")
class FileFetcher:
"""
文件获取上下文管理器
使用方式:
with FileFetcher("finance/报销制度.pdf") as f:
# f.path 是可用于解析的文件路径
# 自动处理临时文件清理
parse_document(f.path, ...)
"""
def __init__(self, file_path: str, use_file_provider: bool = True):
self.file_path = file_path
self.use_file_provider = use_file_provider
self.temp_path: Optional[str] = None
self.is_temp = False
def __enter__(self):
self.temp_path, self.is_temp = get_file_for_parsing(
self.file_path,
self.use_file_provider
)
return self
def __exit__(self, exc_type, exc_val, exc_tb):
if self.is_temp and self.temp_path:
cleanup_temp_file(self.temp_path, True)
return False
@property
def path(self) -> str:
"""获取文件路径"""
return self.temp_path
# ==================== 便捷函数 ====================
def fetch_file(file_path: str) -> FileFetcher:
"""
获取文件(上下文管理器)
Args:
file_path: 文件路径
Returns:
FileFetcher 上下文管理器
Example:
with fetch_file("finance/报销制度.pdf") as f:
result = parse_document(f.path)
"""
return FileFetcher(file_path)