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
« 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
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
29from open_webui.env import USE_SLIM
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
41log = logging.getLogger(__name__)
44class StorageProvider(ABC):
45 @abstractmethod
46 def get_file(self, file_path: str) -> str:
47 pass
49 @abstractmethod
50 def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]:
51 pass
53 @abstractmethod
54 def delete_all_files(self) -> None:
55 pass
57 @abstractmethod
58 def delete_file(self, file_path: str) -> None:
59 pass
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
73 @staticmethod
74 def get_file(file_path: str) -> str:
75 """Handles downloading of the file from local storage."""
76 return file_path
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.')
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.')
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 )
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 )
137 self.bucket_name = S3_BUCKET_NAME
138 self.key_prefix = S3_KEY_PREFIX if S3_KEY_PREFIX else ''
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)
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}')
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}')
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}')
184 # Always delete from local storage
185 LocalStorageProvider.delete_file(file_path)
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
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}')
201 # Always delete from local storage
202 LocalStorageProvider.delete_all_files()
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:])
208 def _get_local_file_path(self, s3_key: str) -> str:
209 return os.path.join(UPLOAD_DIR, s3_key.split('/')[-1])
212class GCSStorageProvider(StorageProvider):
213 def __init__(self):
214 self.bucket_name = GCS_BUCKET_NAME
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)
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}')
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)
245 return local_file_path
246 except NotFound as e:
247 raise RuntimeError(f'Error downloading file from GCS: {e}')
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}')
258 # Always delete from local storage
259 LocalStorageProvider.delete_file(file_path)
261 def delete_all_files(self) -> None:
262 """Handles deletion of all files from GCS storage."""
263 try:
264 blobs = self.bucket.list_blobs()
266 for blob in blobs:
267 blob.delete()
269 except NotFound as e:
270 raise RuntimeError(f'Error deleting all files from GCS: {e}')
272 # Always delete from local storage
273 LocalStorageProvider.delete_all_files()
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
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)
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}')
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}')
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}')
322 # Always delete from local storage
323 LocalStorageProvider.delete_file(file_path)
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}')
334 # Always delete from local storage
335 LocalStorageProvider.delete_all_files()
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
356Storage = get_storage_provider(STORAGE_PROVIDER)