Coverage for open_webui/storage/provider.py: 36%

216 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import logging 

2import os 

3import re 

4import shutil 

5from abc import ABC, abstractmethod 

6from typing import BinaryIO, Dict, Tuple 

7 

8from open_webui.config import ( 

9 AZURE_STORAGE_CONTAINER_NAME, 

10 AZURE_STORAGE_ENDPOINT, 

11 AZURE_STORAGE_KEY, 

12 GCS_BUCKET_NAME, 

13 GOOGLE_APPLICATION_CREDENTIALS_JSON, 

14 S3_ACCESS_KEY_ID, 

15 S3_ADDRESSING_STYLE, 

16 S3_BUCKET_NAME, 

17 S3_ENABLE_TAGGING, 

18 S3_ENDPOINT_URL, 

19 S3_KEY_PREFIX, 

20 S3_REGION_NAME, 

21 S3_SECRET_ACCESS_KEY, 

22 S3_USE_ACCELERATE_ENDPOINT, 

23 STORAGE_PROVIDER, 

24 UPLOAD_DIR, 

25) 

26from open_webui.constants import ERROR_MESSAGES 

27from open_webui.utils.json_codec import JSONCodec 

28 

29from open_webui.env import USE_SLIM 

30 

31if not USE_SLIM: 31 ↛ 41line 31 didn't jump to line 41 because the condition on line 31 was always true

32 import boto3 

33 from azure.core.exceptions import ResourceNotFoundError 

34 from azure.identity import DefaultAzureCredential 

35 from azure.storage.blob import BlobServiceClient 

36 from botocore.config import Config 

37 from botocore.exceptions import ClientError 

38 from google.cloud import storage 

39 from google.cloud.exceptions import GoogleCloudError, NotFound 

40 

41log = logging.getLogger(__name__) 

42 

43 

44class StorageProvider(ABC): 

45 @abstractmethod 

46 def get_file(self, file_path: str) -> str: 

47 pass 

48 

49 @abstractmethod 

50 def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]: 

51 pass 

52 

53 @abstractmethod 

54 def delete_all_files(self) -> None: 

55 pass 

56 

57 @abstractmethod 

58 def delete_file(self, file_path: str) -> None: 

59 pass 

60 

61 

62class LocalStorageProvider(StorageProvider): 

63 @staticmethod 

64 def upload_file(file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]: 

65 contents = file.read() 

66 if not contents: 

67 raise ValueError(ERROR_MESSAGES.EMPTY_CONTENT) 

68 file_path = os.path.join(UPLOAD_DIR, filename) 

69 with open(file_path, 'wb') as f: 

70 f.write(contents) 

71 return contents, file_path 

72 

73 @staticmethod 

74 def get_file(file_path: str) -> str: 

75 """Handles downloading of the file from local storage.""" 

76 return file_path 

77 

78 @staticmethod 

79 def delete_file(file_path: str) -> None: 

80 """Handles deletion of the file from local storage.""" 

81 filename = os.path.basename(file_path) 

82 file_path = os.path.join(UPLOAD_DIR, filename) 

83 if os.path.isfile(file_path): 83 ↛ 86line 83 didn't jump to line 86 because the condition on line 83 was always true

84 os.remove(file_path) 

85 else: 

86 log.warning(f'File {file_path} not found in local storage.') 

87 

88 @staticmethod 

89 def delete_all_files() -> None: 

90 """Handles deletion of all files from local storage.""" 

91 if os.path.exists(UPLOAD_DIR): 91 ↛ 102line 91 didn't jump to line 102 because the condition on line 91 was always true

92 for filename in os.listdir(UPLOAD_DIR): 

93 file_path = os.path.join(UPLOAD_DIR, filename) 

94 try: 

95 if os.path.isfile(file_path) or os.path.islink(file_path): 95 ↛ 97line 95 didn't jump to line 97 because the condition on line 95 was always true

96 os.unlink(file_path) # Remove the file or link 

97 elif os.path.isdir(file_path): 

98 shutil.rmtree(file_path) # Remove the directory 

99 except Exception as e: 

100 log.exception(f'Failed to delete {file_path}. Reason: {e}') 

101 else: 

102 log.warning(f'Directory {UPLOAD_DIR} not found in local storage.') 

103 

104 

105class S3StorageProvider(StorageProvider): 

106 def __init__(self): 

107 config = Config( 

108 s3={ 

109 'use_accelerate_endpoint': S3_USE_ACCELERATE_ENDPOINT, 

110 'addressing_style': S3_ADDRESSING_STYLE, 

111 }, 

112 # KIT change - see https://github.com/boto/boto3/issues/4400#issuecomment-2600742103∆ 

113 request_checksum_calculation='when_required', 

114 response_checksum_validation='when_required', 

115 ) 

116 

117 # If access key and secret are provided, use them for authentication 

118 if S3_ACCESS_KEY_ID and S3_SECRET_ACCESS_KEY: 

119 self.s3_client = boto3.client( 

120 's3', 

121 region_name=S3_REGION_NAME, 

122 endpoint_url=S3_ENDPOINT_URL, 

123 aws_access_key_id=S3_ACCESS_KEY_ID, 

124 aws_secret_access_key=S3_SECRET_ACCESS_KEY, 

125 config=config, 

126 ) 

127 else: 

128 # If no explicit credentials are provided, fall back to default AWS credentials 

129 # This supports workload identity (IAM roles for EC2, EKS, etc.) 

130 self.s3_client = boto3.client( 

131 's3', 

132 region_name=S3_REGION_NAME, 

133 endpoint_url=S3_ENDPOINT_URL, 

134 config=config, 

135 ) 

136 

137 self.bucket_name = S3_BUCKET_NAME 

138 self.key_prefix = S3_KEY_PREFIX if S3_KEY_PREFIX else '' 

139 

140 @staticmethod 

141 def sanitize_tag_value(s: str) -> str: 

142 """Only include S3 allowed characters.""" 

143 return re.sub(r'[^a-zA-Z0-9 äöüÄÖÜß\+\-=\._:/@]', '', s) 

144 

145 def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]: 

146 """Handles uploading of the file to S3 storage.""" 

147 contents, file_path = LocalStorageProvider.upload_file(file, filename, tags) 

148 s3_key = os.path.join(self.key_prefix, filename) 

149 try: 

150 self.s3_client.upload_file(file_path, self.bucket_name, s3_key) 

151 if S3_ENABLE_TAGGING and tags: 

152 sanitized_tags = {self.sanitize_tag_value(k): self.sanitize_tag_value(v) for k, v in tags.items()} 

153 tagging = {'TagSet': [{'Key': k, 'Value': v} for k, v in sanitized_tags.items()]} 

154 self.s3_client.put_object_tagging( 

155 Bucket=self.bucket_name, 

156 Key=s3_key, 

157 Tagging=tagging, 

158 ) 

159 return ( 

160 contents, 

161 f's3://{self.bucket_name}/{s3_key}', 

162 ) 

163 except ClientError as e: 

164 raise RuntimeError(f'Error uploading file to S3: {e}') 

165 

166 def get_file(self, file_path: str) -> str: 

167 """Handles downloading of the file from S3 storage.""" 

168 try: 

169 s3_key = self._extract_s3_key(file_path) 

170 local_file_path = self._get_local_file_path(s3_key) 

171 self.s3_client.download_file(self.bucket_name, s3_key, local_file_path) 

172 return local_file_path 

173 except ClientError as e: 

174 raise RuntimeError(f'Error downloading file from S3: {e}') 

175 

176 def delete_file(self, file_path: str) -> None: 

177 """Handles deletion of the file from S3 storage.""" 

178 try: 

179 s3_key = self._extract_s3_key(file_path) 

180 self.s3_client.delete_object(Bucket=self.bucket_name, Key=s3_key) 

181 except ClientError as e: 

182 raise RuntimeError(f'Error deleting file from S3: {e}') 

183 

184 # Always delete from local storage 

185 LocalStorageProvider.delete_file(file_path) 

186 

187 def delete_all_files(self) -> None: 

188 """Handles deletion of all files from S3 storage.""" 

189 try: 

190 response = self.s3_client.list_objects_v2(Bucket=self.bucket_name) 

191 if 'Contents' in response: 

192 for content in response['Contents']: 

193 # Skip objects that were not uploaded from open-webui in the first place 

194 if not content['Key'].startswith(self.key_prefix): 

195 continue 

196 

197 self.s3_client.delete_object(Bucket=self.bucket_name, Key=content['Key']) 

198 except ClientError as e: 

199 raise RuntimeError(f'Error deleting all files from S3: {e}') 

200 

201 # Always delete from local storage 

202 LocalStorageProvider.delete_all_files() 

203 

204 # The s3 key is the name assigned to an object. It excludes the bucket name, but includes the internal path and the file name. 

205 def _extract_s3_key(self, full_file_path: str) -> str: 

206 return '/'.join(full_file_path.split('//')[1].split('/')[1:]) 

207 

208 def _get_local_file_path(self, s3_key: str) -> str: 

209 return os.path.join(UPLOAD_DIR, s3_key.split('/')[-1]) 

210 

211 

212class GCSStorageProvider(StorageProvider): 

213 def __init__(self): 

214 self.bucket_name = GCS_BUCKET_NAME 

215 

216 if GOOGLE_APPLICATION_CREDENTIALS_JSON: 

217 self.gcs_client = storage.Client.from_service_account_info( 

218 info=JSONCodec.loads(GOOGLE_APPLICATION_CREDENTIALS_JSON) 

219 ) 

220 else: 

221 # if no credentials json is provided, credentials will be picked up from the environment 

222 # if running on local environment, credentials would be user credentials 

223 # if running on a Compute Engine instance, credentials would be from Google Metadata server 

224 self.gcs_client = storage.Client() 

225 self.bucket = self.gcs_client.bucket(GCS_BUCKET_NAME) 

226 

227 def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]: 

228 """Handles uploading of the file to GCS storage.""" 

229 contents, file_path = LocalStorageProvider.upload_file(file, filename, tags) 

230 try: 

231 blob = self.bucket.blob(filename) 

232 blob.upload_from_filename(file_path) 

233 return contents, 'gs://' + self.bucket_name + '/' + filename 

234 except GoogleCloudError as e: 

235 raise RuntimeError(f'Error uploading file to GCS: {e}') 

236 

237 def get_file(self, file_path: str) -> str: 

238 """Handles downloading of the file from GCS storage.""" 

239 try: 

240 filename = file_path.removeprefix('gs://').split('/')[1] 

241 local_file_path = os.path.join(UPLOAD_DIR, filename) 

242 blob = self.bucket.get_blob(filename) 

243 blob.download_to_filename(local_file_path) 

244 

245 return local_file_path 

246 except NotFound as e: 

247 raise RuntimeError(f'Error downloading file from GCS: {e}') 

248 

249 def delete_file(self, file_path: str) -> None: 

250 """Handles deletion of the file from GCS storage.""" 

251 try: 

252 filename = file_path.removeprefix('gs://').split('/')[1] 

253 blob = self.bucket.get_blob(filename) 

254 blob.delete() 

255 except NotFound as e: 

256 raise RuntimeError(f'Error deleting file from GCS: {e}') 

257 

258 # Always delete from local storage 

259 LocalStorageProvider.delete_file(file_path) 

260 

261 def delete_all_files(self) -> None: 

262 """Handles deletion of all files from GCS storage.""" 

263 try: 

264 blobs = self.bucket.list_blobs() 

265 

266 for blob in blobs: 

267 blob.delete() 

268 

269 except NotFound as e: 

270 raise RuntimeError(f'Error deleting all files from GCS: {e}') 

271 

272 # Always delete from local storage 

273 LocalStorageProvider.delete_all_files() 

274 

275 

276class AzureStorageProvider(StorageProvider): 

277 def __init__(self): 

278 self.endpoint = AZURE_STORAGE_ENDPOINT 

279 self.container_name = AZURE_STORAGE_CONTAINER_NAME 

280 storage_key = AZURE_STORAGE_KEY 

281 

282 if storage_key: 

283 # Configure using the Azure Storage Account Endpoint and Key 

284 self.blob_service_client = BlobServiceClient(account_url=self.endpoint, credential=storage_key) 

285 else: 

286 # Configure using the Azure Storage Account Endpoint and DefaultAzureCredential 

287 # If the key is not configured, then the DefaultAzureCredential will be used to support Managed Identity authentication 

288 self.blob_service_client = BlobServiceClient(account_url=self.endpoint, credential=DefaultAzureCredential()) 

289 self.container_client = self.blob_service_client.get_container_client(self.container_name) 

290 

291 def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]: 

292 """Handles uploading of the file to Azure Blob Storage.""" 

293 contents, file_path = LocalStorageProvider.upload_file(file, filename, tags) 

294 try: 

295 blob_client = self.container_client.get_blob_client(filename) 

296 blob_client.upload_blob(contents, overwrite=True) 

297 return contents, f'{self.endpoint}/{self.container_name}/{filename}' 

298 except Exception as e: 

299 raise RuntimeError(f'Error uploading file to Azure Blob Storage: {e}') 

300 

301 def get_file(self, file_path: str) -> str: 

302 """Handles downloading of the file from Azure Blob Storage.""" 

303 try: 

304 filename = file_path.split('/')[-1] 

305 local_file_path = os.path.join(UPLOAD_DIR, filename) 

306 blob_client = self.container_client.get_blob_client(filename) 

307 with open(local_file_path, 'wb') as download_file: 

308 download_file.write(blob_client.download_blob().readall()) 

309 return local_file_path 

310 except ResourceNotFoundError as e: 

311 raise RuntimeError(f'Error downloading file from Azure Blob Storage: {e}') 

312 

313 def delete_file(self, file_path: str) -> None: 

314 """Handles deletion of the file from Azure Blob Storage.""" 

315 try: 

316 filename = file_path.split('/')[-1] 

317 blob_client = self.container_client.get_blob_client(filename) 

318 blob_client.delete_blob() 

319 except ResourceNotFoundError as e: 

320 raise RuntimeError(f'Error deleting file from Azure Blob Storage: {e}') 

321 

322 # Always delete from local storage 

323 LocalStorageProvider.delete_file(file_path) 

324 

325 def delete_all_files(self) -> None: 

326 """Handles deletion of all files from Azure Blob Storage.""" 

327 try: 

328 blobs = self.container_client.list_blobs() 

329 for blob in blobs: 

330 self.container_client.delete_blob(blob.name) 

331 except Exception as e: 

332 raise RuntimeError(f'Error deleting all files from Azure Blob Storage: {e}') 

333 

334 # Always delete from local storage 

335 LocalStorageProvider.delete_all_files() 

336 

337 

338def get_storage_provider(storage_provider: str): 

339 if USE_SLIM and storage_provider != 'local': 339 ↛ 340line 339 didn't jump to line 340 because the condition on line 339 was never true

340 raise RuntimeError( 

341 'Slim requires local file storage. Set STORAGE_PROVIDER=local, or use the standard image to access cloud storage.' 

342 ) 

343 if storage_provider == 'local': 343 ↛ 345line 343 didn't jump to line 345 because the condition on line 343 was always true

344 Storage = LocalStorageProvider() 

345 elif storage_provider == 's3': 

346 Storage = S3StorageProvider() 

347 elif storage_provider == 'gcs': 

348 Storage = GCSStorageProvider() 

349 elif storage_provider == 'azure': 

350 Storage = AzureStorageProvider() 

351 else: 

352 raise RuntimeError(f'Unsupported storage provider: {storage_provider}') 

353 return Storage 

354 

355 

356Storage = get_storage_provider(STORAGE_PROVIDER)