feat: 新增webdav存储

This commit is contained in:
Lan
2025-02-09 20:25:23 +08:00
parent 2e44fe29a4
commit 1cbe8d5f2a
20 changed files with 281 additions and 82 deletions
+254 -60
View File
@@ -3,7 +3,7 @@
# @File : storage.py
# @Software: PyCharm
from typing import Optional
from urllib.parse import quote, unquote
import aiohttp
import asyncio
from pathlib import Path
@@ -22,11 +22,13 @@ from fastapi.responses import FileResponse
class FileStorageInterface:
_instance: Optional['FileStorageInterface'] = None
_instance: Optional["FileStorageInterface"] = None
def __new__(cls, *args, **kwargs):
if cls._instance is None:
cls._instance = super(FileStorageInterface, cls).__new__(cls, *args, **kwargs)
cls._instance = super(FileStorageInterface, cls).__new__(
cls, *args, **kwargs
)
return cls._instance
async def save_file(self, file: UploadFile, save_path: str):
@@ -66,7 +68,7 @@ class SystemFileStorage(FileStorageInterface):
self.root_path = data_root
def _save(self, file, save_path):
with open(save_path, 'wb') as f:
with open(save_path, "wb") as f:
chunk = file.read(self.chunk_size)
while chunk:
f.write(chunk)
@@ -89,7 +91,7 @@ class SystemFileStorage(FileStorageInterface):
async def get_file_response(self, file_code: FileCodes):
file_path = self.root_path / await file_code.get_file_path()
if not file_path.exists():
return APIResponse(code=404, detail='文件已过期删除')
return APIResponse(code=404, detail="文件已过期删除")
return FileResponse(file_path, filename=file_code.prefix + file_code.suffix)
@@ -101,30 +103,62 @@ class S3FileStorage(FileStorageInterface):
self.s3_hostname = settings.s3_hostname
self.region_name = settings.s3_region_name
self.signature_version = settings.s3_signature_version
self.endpoint_url = settings.s3_endpoint_url or f'https://{self.s3_hostname}'
self.endpoint_url = settings.s3_endpoint_url or f"https://{self.s3_hostname}"
self.aws_session_token = settings.aws_session_token
self.proxy = settings.s3_proxy
self.session = aioboto3.Session(aws_access_key_id=self.access_key_id, aws_secret_access_key=self.secret_access_key)
self.session = aioboto3.Session(
aws_access_key_id=self.access_key_id,
aws_secret_access_key=self.secret_access_key,
)
if not settings.s3_endpoint_url:
self.endpoint_url = f'https://{self.s3_hostname}'
self.endpoint_url = f"https://{self.s3_hostname}"
else:
# 如果提供了 s3_endpoint_url,则优先使用它
self.endpoint_url = settings.s3_endpoint_url
async def save_file(self, file: UploadFile, save_path: str):
async with self.session.client("s3", endpoint_url=self.endpoint_url, aws_session_token=self.aws_session_token, region_name=self.region_name,
config=Config(signature_version=self.signature_version)) as s3:
await s3.put_object(Bucket=self.bucket_name, Key=save_path, Body=await file.read(), ContentType=file.content_type)
async with self.session.client(
"s3",
endpoint_url=self.endpoint_url,
aws_session_token=self.aws_session_token,
region_name=self.region_name,
config=Config(signature_version=self.signature_version),
) as s3:
await s3.put_object(
Bucket=self.bucket_name,
Key=save_path,
Body=await file.read(),
ContentType=file.content_type,
)
async def delete_file(self, file_code: FileCodes):
async with self.session.client("s3", endpoint_url=self.endpoint_url, region_name=self.region_name, config=Config(signature_version=self.signature_version)) as s3:
await s3.delete_object(Bucket=self.bucket_name, Key=await file_code.get_file_path())
async with self.session.client(
"s3",
endpoint_url=self.endpoint_url,
region_name=self.region_name,
config=Config(signature_version=self.signature_version),
) as s3:
await s3.delete_object(
Bucket=self.bucket_name, Key=await file_code.get_file_path()
)
async def get_file_response(self, file_code: FileCodes):
try:
filename = file_code.prefix + file_code.suffix
async with self.session.client("s3", endpoint_url=self.endpoint_url, region_name=self.region_name, config=Config(signature_version=self.signature_version)) as s3:
link = await s3.generate_presigned_url('get_object', Params={'Bucket': self.bucket_name, 'Key': await file_code.get_file_path()}, ExpiresIn=3600)
async with self.session.client(
"s3",
endpoint_url=self.endpoint_url,
region_name=self.region_name,
config=Config(signature_version=self.signature_version),
) as s3:
link = await s3.generate_presigned_url(
"get_object",
Params={
"Bucket": self.bucket_name,
"Key": await file_code.get_file_path(),
},
ExpiresIn=3600,
)
tmp = io.BytesIO()
async with aiohttp.ClientSession() as session:
async with session.get(link) as resp:
@@ -132,18 +166,36 @@ class S3FileStorage(FileStorageInterface):
tmp.seek(0)
content = tmp.read()
tmp.close()
return Response(content, media_type="application/octet-stream", headers={"Content-Disposition": f'attachment; filename="{filename.encode("utf-8").decode("latin-1")}"'})
return Response(
content,
media_type="application/octet-stream",
headers={
"Content-Disposition": f'attachment; filename="{filename.encode("utf-8").decode("latin-1")}"'
},
)
except Exception:
raise HTTPException(status_code=503, detail='服务代理下载异常,请稍后再试')
raise HTTPException(status_code=503, detail="服务代理下载异常,请稍后再试")
async def get_file_url(self, file_code: FileCodes):
if file_code.prefix == '文本分享':
if file_code.prefix == "文本分享":
return file_code.text
if self.proxy:
return await get_file_url(file_code.code)
else:
async with self.session.client("s3", endpoint_url=self.endpoint_url, region_name=self.region_name, config=Config(signature_version=self.signature_version)) as s3:
result = await s3.generate_presigned_url('get_object', Params={'Bucket': self.bucket_name, 'Key': await file_code.get_file_path()}, ExpiresIn=3600)
async with self.session.client(
"s3",
endpoint_url=self.endpoint_url,
region_name=self.region_name,
config=Config(signature_version=self.signature_version),
) as s3:
result = await s3.generate_presigned_url(
"get_object",
Params={
"Bucket": self.bucket_name,
"Key": await file_code.get_file_path(),
},
ExpiresIn=3600,
)
return result
@@ -152,9 +204,11 @@ class OneDriveFileStorage(FileStorageInterface):
try:
import msal
from office365.graph_client import GraphClient
from office365.runtime.client_request_exception import ClientRequestException
from office365.runtime.client_request_exception import (
ClientRequestException,
)
except ImportError:
raise ImportError('请先安装`msal`和`Office365-REST-Python-Client`')
raise ImportError("请先安装`msal`和`Office365-REST-Python-Client`")
self.msal = msal
self.domain = settings.onedrive_domain
self.client_id = settings.onedrive_client_id
@@ -165,36 +219,45 @@ class OneDriveFileStorage(FileStorageInterface):
try:
client = GraphClient(self.acquire_token_pwd)
self.root_path = client.me.drive.root.get_by_path(settings.onedrive_root_path).get().execute_query()
self.root_path = (
client.me.drive.root.get_by_path(settings.onedrive_root_path)
.get()
.execute_query()
)
except ClientRequestException as e:
if e.code == 'itemNotFound':
if e.code == "itemNotFound":
client.me.drive.root.create_folder(settings.onedrive_root_path)
self.root_path = client.me.drive.root.get_by_path(settings.onedrive_root_path).get().execute_query()
self.root_path = (
client.me.drive.root.get_by_path(settings.onedrive_root_path)
.get()
.execute_query()
)
else:
raise e
except Exception as e:
raise Exception('OneDrive验证失败,请检查配置是否正确\n' + str(e))
raise Exception("OneDrive验证失败,请检查配置是否正确\n" + str(e))
def acquire_token_pwd(self):
authority_url = f'https://login.microsoftonline.com/{self.domain}'
authority_url = f"https://login.microsoftonline.com/{self.domain}"
app = self.msal.PublicClientApplication(
authority=authority_url,
client_id=self.client_id
authority=authority_url, client_id=self.client_id
)
result = app.acquire_token_by_username_password(
username=self.username,
password=self.password,
scopes=["https://graph.microsoft.com/.default"],
)
result = app.acquire_token_by_username_password(username=self.username,
password=self.password,
scopes=['https://graph.microsoft.com/.default'])
return result
def _get_path_str(self, path):
if isinstance(path, str):
path = path.replace('\\', '/').replace('//', '/').split('/')
path = path.replace("\\", "/").replace("//", "/").split("/")
elif isinstance(path, Path):
path = str(path).replace('\\', '/').replace('//', '/').split('/')
path = str(path).replace("\\", "/").replace("//", "/").split("/")
else:
raise TypeError('path must be str or Path')
path[-1] = path[-1].split('.')[0]
return '/'.join(path)
raise TypeError("path must be str or Path")
path[-1] = path[-1].split(".")[0]
return "/".join(path)
def _save(self, file, save_path):
content = file.file.read()
@@ -210,7 +273,7 @@ class OneDriveFileStorage(FileStorageInterface):
try:
self.root_path.get_by_path(path).delete_object().execute_query()
except self._ClientRequestException as e:
if e.code == 'itemNotFound':
if e.code == "itemNotFound":
pass
else:
raise e
@@ -219,23 +282,29 @@ class OneDriveFileStorage(FileStorageInterface):
await asyncio.to_thread(self._delete, await file_code.get_file_path())
def _convert_link_to_download_link(self, link):
p1 = re.search(r'https:\/\/(.+)\.sharepoint\.com', link).group(1)
p2 = re.search(r'personal\/(.+)\/', link).group(1)
p3 = re.search(rf'{p2}\/(.+)', link).group(1)
return f'https://{p1}.sharepoint.com/personal/{p2}/_layouts/52/download.aspx?share={p3}'
p1 = re.search(r"https:\/\/(.+)\.sharepoint\.com", link).group(1)
p2 = re.search(r"personal\/(.+)\/", link).group(1)
p3 = re.search(rf"{p2}\/(.+)", link).group(1)
return f"https://{p1}.sharepoint.com/personal/{p2}/_layouts/52/download.aspx?share={p3}"
def _get_file_url(self, save_path, name):
path = self._get_path_str(save_path)
remote_file = self.root_path.get_by_path(path + '/' + name)
expiration_datetime = datetime.datetime.now(tz=datetime.timezone.utc) + datetime.timedelta(hours=1)
remote_file = self.root_path.get_by_path(path + "/" + name)
expiration_datetime = datetime.datetime.now(
tz=datetime.timezone.utc
) + datetime.timedelta(hours=1)
expiration_datetime = expiration_datetime.strftime("%Y-%m-%dT%H:%M:%SZ")
premission = remote_file.create_link("view", "anonymous", expiration_datetime=expiration_datetime).execute_query()
premission = remote_file.create_link(
"view", "anonymous", expiration_datetime=expiration_datetime
).execute_query()
return self._convert_link_to_download_link(premission.link.webUrl)
async def get_file_response(self, file_code: FileCodes):
try:
filename = file_code.prefix + file_code.suffix
link = await asyncio.to_thread(self._get_file_url, await file_code.get_file_path(), filename)
link = await asyncio.to_thread(
self._get_file_url, await file_code.get_file_path(), filename
)
tmp = io.BytesIO()
async with aiohttp.ClientSession() as session:
async with session.get(link) as resp:
@@ -243,15 +312,25 @@ class OneDriveFileStorage(FileStorageInterface):
tmp.seek(0)
content = tmp.read()
tmp.close()
return Response(content, media_type="application/octet-stream", headers={"Content-Disposition": f'attachment; filename="{filename.encode("utf-8").decode("latin-1")}"'})
return Response(
content,
media_type="application/octet-stream",
headers={
"Content-Disposition": f'attachment; filename="{filename.encode("utf-8").decode("latin-1")}"'
},
)
except Exception:
raise HTTPException(status_code=503, detail='服务代理下载异常,请稍后再试')
raise HTTPException(status_code=503, detail="服务代理下载异常,请稍后再试")
async def get_file_url(self, file_code: FileCodes):
if self.proxy:
return await get_file_url(file_code.code)
else:
return await asyncio.to_thread(self._get_file_url, await file_code.get_file_path(), f'{file_code.prefix}{file_code.suffix}')
return await asyncio.to_thread(
self._get_file_url,
await file_code.get_file_path(),
f"{file_code.prefix}{file_code.suffix}",
)
class OpenDALFileStorage(FileStorageInterface):
@@ -263,10 +342,12 @@ class OpenDALFileStorage(FileStorageInterface):
self.service = settings.opendal_scheme
service_settings = {}
for key, value in settings.items():
if key.startswith('opendal_' + self.service):
setting_name = key.split('_', 2)[2]
if key.startswith("opendal_" + self.service):
setting_name = key.split("_", 2)[2]
service_settings[setting_name] = value
self.operator = opendal.AsyncOperator(settings.opendal_scheme, **service_settings)
self.operator = opendal.AsyncOperator(
settings.opendal_scheme, **service_settings
)
async def save_file(self, file: UploadFile, save_path: str):
await self.operator.write(save_path, file.file.read())
@@ -281,18 +362,131 @@ class OpenDALFileStorage(FileStorageInterface):
try:
filename = file_code.prefix + file_code.suffix
content = await self.operator.read(await file_code.get_file_path())
headers = {
"Content-Disposition": f'attachment; filename="{filename}"'
}
return Response(content, headers=headers, media_type="application/octet-stream")
headers = {"Content-Disposition": f'attachment; filename="{filename}"'}
return Response(
content, headers=headers, media_type="application/octet-stream"
)
except Exception as e:
print(e, file=sys.stderr)
raise HTTPException(status_code=404, detail="文件已过期删除")
class WebDAVFileStorage(FileStorageInterface):
_instance: Optional["WebDAVFileStorage"] = None
def __init__(self):
if not hasattr(self, "_initialized"):
self.base_url = settings.webdav_url.rstrip("/") + "/"
self.auth = aiohttp.BasicAuth(
login=settings.webdav_username, password=settings.webdav_password
)
self._initialized = True
def _build_url(self, path: str) -> str:
encoded_path = quote(str(path).lstrip("/"))
return f"{self.base_url}{encoded_path}"
async def _mkdir_p(self, directory_path: str):
"""递归创建目录(类似mkdir -p"""
path_obj = Path(unquote(directory_path))
current_path = ""
async with aiohttp.ClientSession(auth=self.auth) as session:
# 逐级检查目录是否存在
for part in path_obj.parts:
current_path = str(Path(current_path) / part)
url = self._build_url(current_path)
# 检查目录是否存在
async with session.head(url) as resp:
if resp.status == 404:
# 创建目录
async with session.request("MKCOL", url) as mkcol_resp:
if mkcol_resp.status not in (200, 201, 409):
content = await mkcol_resp.text()
raise HTTPException(
status_code=mkcol_resp.status,
detail=f"目录创建失败: {content[:200]}",
)
async def save_file(self, file: UploadFile, save_path: str):
"""保存文件(自动创建目录)"""
# 分离文件名和目录路径
path_obj = Path(save_path)
directory_path = str(path_obj.parent)
file_name = path_obj.name
try:
# 先创建目录结构
await self._mkdir_p(directory_path)
# 上传文件
url = self._build_url(save_path)
async with aiohttp.ClientSession(auth=self.auth) as session:
content = await file.read() # 注意:大文件需要分块读取
async with session.put(
url, data=content, headers={"Content-Type": file.content_type}
) as resp:
if resp.status not in (200, 201, 204):
content = await resp.text()
raise HTTPException(
status_code=resp.status,
detail=f"文件上传失败: {content[:200]}",
)
except aiohttp.ClientError as e:
raise HTTPException(status_code=503, detail=f"WebDAV连接异常: {str(e)}")
async def delete_file(self, file_code: FileCodes):
"""删除WebDAV文件"""
file_path = await file_code.get_file_path()
url = self._build_url(file_path)
try:
async with aiohttp.ClientSession(auth=self.auth) as session:
async with session.delete(url) as resp:
if resp.status not in (200, 204):
content = await resp.text()
raise HTTPException(
status_code=resp.status,
detail=f"WebDAV删除失败: {content[:200]}",
)
except aiohttp.ClientError as e:
raise HTTPException(status_code=503, detail=f"WebDAV连接异常: {str(e)}")
async def get_file_url(self, file_code: FileCodes):
return await get_file_url(file_code.code)
async def get_file_response(self, file_code: FileCodes):
"""获取文件响应(代理模式)"""
try:
filename = file_code.prefix + file_code.suffix
url = self._build_url(await file_code.get_file_path())
async with aiohttp.ClientSession(auth=self.auth) as session:
async with session.get(url) as resp:
if resp.status != 200:
raise HTTPException(
status_code=resp.status,
detail=f"文件获取失败: {await resp.text()}",
)
# 读取内容到内存
content = await resp.read()
return Response(
content=content,
media_type=resp.headers.get(
"Content-Type", "application/octet-stream"
),
headers={
"Content-Disposition": f'attachment; filename="{filename.encode("utf-8").decode()}"'
},
)
except aiohttp.ClientError as e:
raise HTTPException(status_code=503, detail=f"WebDAV连接异常: {str(e)}")
storages = {
'local': SystemFileStorage,
's3': S3FileStorage,
'onedrive': OneDriveFileStorage,
'opendal': OpenDALFileStorage,
"local": SystemFileStorage,
"s3": S3FileStorage,
"onedrive": OneDriveFileStorage,
"opendal": OpenDALFileStorage,
"webdav": WebDAVFileStorage,
}