From 40e80113ca8909d90abe2c145bf3e99dbde99524 Mon Sep 17 00:00:00 2001 From: Farrizal Alchudry Mutaqien Date: Sun, 10 Aug 2025 23:30:27 +0700 Subject: [PATCH] refactor: google drive service in knowledge base route into generic content sources standalone folders --- .../open_webui/content_sources/__init__.py | 24 + backend/open_webui/content_sources/base.py | 245 +++ backend/open_webui/content_sources/factory.py | 84 + .../content_sources/providers/__init__.py | 1 + .../providers/google_drive/__init__.py | 4 + .../providers/google_drive/client.py} | 94 +- .../providers/google_drive/provider.py | 460 ++++++ .../open_webui/content_sources/registry.py | 125 ++ .../open_webui/content_sources/scheduler.py | 261 ++++ backend/open_webui/content_sources/syncer.py | 1385 +++++++++++++++++ backend/open_webui/main.py | 36 +- backend/open_webui/models/files.py | 87 ++ backend/open_webui/routers/content_sources.py | 279 ++++ backend/open_webui/routers/knowledge.py | 456 +++--- backend/open_webui/routers/retrieval.py | 6 +- backend/open_webui/services/__init__.py | 1 - .../services/google_drive_scheduler.py | 238 --- backend/open_webui/utils/content_sources.py | 38 + src/lib/apis/knowledge/index.ts | 82 +- .../workspace/Knowledge/KnowledgeBase.svelte | 48 +- .../KnowledgeBase/AddContentMenu.svelte | 3 + .../KnowledgeBase/ContentSourceStatus.svelte | 41 + .../ContentSourceSyncModal.svelte | 484 ++++++ .../KnowledgeBase/GoogleDriveSyncModal.svelte | 258 --- src/lib/types/content-sources.ts | 161 ++ src/lib/types/google-drive.ts | 25 +- src/lib/types/index.ts | 6 +- src/lib/utils/content-sources.ts | 130 ++ 28 files changed, 4210 insertions(+), 852 deletions(-) create mode 100644 backend/open_webui/content_sources/__init__.py create mode 100644 backend/open_webui/content_sources/base.py create mode 100644 backend/open_webui/content_sources/factory.py create mode 100644 backend/open_webui/content_sources/providers/__init__.py create mode 100644 backend/open_webui/content_sources/providers/google_drive/__init__.py rename backend/open_webui/{services/google_drive.py => content_sources/providers/google_drive/client.py} (83%) create mode 100644 backend/open_webui/content_sources/providers/google_drive/provider.py create mode 100644 backend/open_webui/content_sources/registry.py create mode 100644 backend/open_webui/content_sources/scheduler.py create mode 100644 backend/open_webui/content_sources/syncer.py create mode 100644 backend/open_webui/routers/content_sources.py delete mode 100644 backend/open_webui/services/__init__.py delete mode 100644 backend/open_webui/services/google_drive_scheduler.py create mode 100644 backend/open_webui/utils/content_sources.py create mode 100644 src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceStatus.svelte create mode 100644 src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceSyncModal.svelte delete mode 100644 src/lib/components/workspace/Knowledge/KnowledgeBase/GoogleDriveSyncModal.svelte create mode 100644 src/lib/types/content-sources.ts create mode 100644 src/lib/utils/content-sources.ts diff --git a/backend/open_webui/content_sources/__init__.py b/backend/open_webui/content_sources/__init__.py new file mode 100644 index 0000000000..b4e428be3f --- /dev/null +++ b/backend/open_webui/content_sources/__init__.py @@ -0,0 +1,24 @@ +""" +Content Sources Module + +This module provides a unified interface for integrating external content sources +(Google Drive, OneDrive, Dropbox, etc.) with Open WebUI's knowledge base system. +""" + +from .factory import content_source_factory +from .registry import content_source_registry +from .scheduler import scheduler as content_source_scheduler + +# Initialize Google Drive provider in the registry +try: + gdrive_provider = content_source_factory.get_provider('google_drive') + content_source_registry.register_provider('google_drive', gdrive_provider) +except Exception as e: + import logging + logging.getLogger(__name__).warning(f"Failed to initialize Google Drive provider: {e}") + +__all__ = [ + "content_source_factory", + "content_source_registry", + "content_source_scheduler" +] \ No newline at end of file diff --git a/backend/open_webui/content_sources/base.py b/backend/open_webui/content_sources/base.py new file mode 100644 index 0000000000..afeb374f11 --- /dev/null +++ b/backend/open_webui/content_sources/base.py @@ -0,0 +1,245 @@ +""" +Base Content Source Provider + +Defines the abstract interface for all content source providers. +""" + +from abc import ABC, abstractmethod +from typing import List, Dict, Any, Callable, Optional, AsyncGenerator +import logging + +logger = logging.getLogger(__name__) + + +class ContentSourceProvider(ABC): + """ + Abstract base class for content source providers. + + Providers implement methods to list, download, and sync files from external sources. + They emit hooks to allow the knowledge system to process files without tight coupling. + """ + + def __init__(self): + self._hooks: Dict[str, List[Callable]] = {} + + def register_hook(self, event: str, handler: Callable) -> None: + """ + Register a hook handler for a specific event. + + Args: + event: The event name (e.g., 'file_ready', 'sync_started') + handler: Async callable that will be invoked when the event is emitted + """ + if event not in self._hooks: + self._hooks[event] = [] + self._hooks[event].append(handler) + logger.debug(f"Registered hook for event '{event}'") + + async def emit_hook(self, event: str, data: Dict[str, Any]) -> None: + """ + Emit an event to all registered handlers. + + Args: + event: The event name + data: Data to pass to the handlers + """ + handlers = self._hooks.get(event, []) + for handler in handlers: + try: + await handler(data) + except Exception as e: + logger.error(f"Error in hook handler for event '{event}': {e}") + + @abstractmethod + async def list_files(self, path: str = "", recursive: bool = True) -> List[Dict[str, Any]]: + """ + List files in the content source. + + Args: + path: The path to list files from (provider-specific format) + recursive: Whether to list files recursively + + Returns: + List of file metadata dictionaries + """ + pass + + @abstractmethod + async def download_file(self, file_id: str) -> AsyncGenerator[bytes, None]: + """ + Download a file from the content source. + + Args: + file_id: Provider-specific file identifier + + Yields: + File content in chunks + """ + pass + + @abstractmethod + async def get_service_info(self) -> Dict[str, Any]: + """ + Get information about the content source service. + + Returns: + Dictionary with service information (e.g., account email, quota) + """ + pass + + async def sync_folder(self, folder_id: str, context: Optional[Dict[str, Any]] = None) -> None: + """ + Sync a folder from the content source. + + This is a high-level method that lists files and emits hooks for processing. + Subclasses can override for provider-specific behavior. + + Args: + folder_id: Provider-specific folder identifier + context: Optional context to pass through hooks (e.g., kb_id, user_id) + """ + context = context or {} + + # Emit sync started event + await self.emit_hook('sync_started', { + 'folder_id': folder_id, + 'context': context + }) + + try: + # List files in the folder + files = await self.list_files(folder_id, recursive=True) + + # Process each file + for file_info in files: + try: + # Download file content + content_chunks = [] + async for chunk in self.download_file(file_info['id']): + content_chunks.append(chunk) + content = b''.join(content_chunks) + + # Emit file ready event + await self.emit_hook('file_ready', { + 'file_info': file_info, + 'content': content, + 'context': context + }) + + except Exception as e: + logger.error(f"Error processing file {file_info.get('name', 'unknown')}: {e}") + await self.emit_hook('file_error', { + 'file_info': file_info, + 'error': str(e), + 'context': context + }) + + # Emit sync completed event + await self.emit_hook('sync_completed', { + 'folder_id': folder_id, + 'file_count': len(files), + 'context': context + }) + + except Exception as e: + logger.error(f"Error during sync: {e}") + await self.emit_hook('sync_error', { + 'folder_id': folder_id, + 'error': str(e), + 'context': context + }) + raise + + async def sync_content(self, source_id: str, context: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + """ + Generic sync method that handles file synchronization and returns results. + + This method provides a unified interface for the knowledge router to sync content. + It handles both simple sync (download all) and advanced sync (with change detection). + + Args: + source_id: Provider-specific source identifier (folder ID, etc.) + context: Context including file_ids, user_id, options, etc. + + Returns: + Dictionary with sync results including added, updated, removed files and errors + """ + context = context or {} + sync_results = { + "added_files": [], + "updated_files": [], + "removed_files": [], + "errors": [], + "changes": False + } + + # Build provider file map from existing files + provider_file_map = {} + existing_file_ids = context.get("file_ids", []) + if existing_file_ids: + from open_webui.models.files import Files + existing_files = Files.get_files_by_ids(existing_file_ids) + for file in existing_files: + if file.data and file.data.get("provider") == context.get("provider_name"): + provider_id = file.data.get("provider_file_id") + if provider_id: + provider_file_map[provider_id] = { + "id": file.id, + "name": file.filename, + "modified_time": file.data.get("provider_modified_time"), + "data": file.data + } + + # Track sync progress + async def track_file_processed(data: Dict[str, Any]): + file_id = data.get("file_id") + if file_id: + sync_results["added_files"].append(file_id) + sync_results["changes"] = True + + async def track_file_updated(data: Dict[str, Any]): + sync_results["updated_files"].append(data.get("file_id", "unknown")) + sync_results["changes"] = True + + async def track_file_removed(data: Dict[str, Any]): + file_info = data.get("file_info", {}) + file_id = file_info.get("id") + if file_id: + sync_results["removed_files"].append(file_id) + sync_results["changes"] = True + + async def track_error(data: Dict[str, Any]): + sync_results["errors"].append({ + "file": data.get("file_info", {}).get("name", "unknown"), + "error": data.get("error", "Unknown error") + }) + + # Register internal tracking hooks + self.register_hook("file_processed", track_file_processed) + self.register_hook("file_updated", track_file_updated) + self.register_hook("file_removed", track_file_removed) + self.register_hook("file_error", track_error) + + try: + # Use advanced sync if available, otherwise basic sync + if hasattr(self, 'sync_folder_with_metadata'): + await self.sync_folder_with_metadata( + folder_id=source_id, + existing_files=provider_file_map, + context=context + ) + else: + await self.sync_folder(source_id, context) + + finally: + # Clean up tracking hooks + if "file_processed" in self._hooks: + self._hooks["file_processed"] = [h for h in self._hooks["file_processed"] if h != track_file_processed] + if "file_updated" in self._hooks: + self._hooks["file_updated"] = [h for h in self._hooks["file_updated"] if h != track_file_updated] + if "file_removed" in self._hooks: + self._hooks["file_removed"] = [h for h in self._hooks["file_removed"] if h != track_file_removed] + if "file_error" in self._hooks: + self._hooks["file_error"] = [h for h in self._hooks["file_error"] if h != track_error] + + return sync_results \ No newline at end of file diff --git a/backend/open_webui/content_sources/factory.py b/backend/open_webui/content_sources/factory.py new file mode 100644 index 0000000000..a24cfa8804 --- /dev/null +++ b/backend/open_webui/content_sources/factory.py @@ -0,0 +1,84 @@ +""" +Content Source Factory + +Factory for creating content source provider instances. +""" + +from typing import Dict, Type, Optional +import logging + +from .base import ContentSourceProvider +from .providers.google_drive import GoogleDriveProvider + +logger = logging.getLogger(__name__) + + +class ContentSourceFactory: + """ + Factory for creating content source provider instances. + + Follows the pattern established by storage providers in the codebase. + """ + + # Registry of available providers + _providers: Dict[str, Type[ContentSourceProvider]] = { + 'google_drive': GoogleDriveProvider, + # Future providers can be added here: + # 'onedrive': OneDriveContentSource, + # 'dropbox': DropboxContentSource, + # 'sharepoint': SharePointContentSource, + } + + @classmethod + def get_provider(cls, provider_type: str) -> ContentSourceProvider: + """ + Get a content source provider instance. + + Args: + provider_type: The type of provider to create + + Returns: + ContentSourceProvider instance + + Raises: + ValueError: If provider type is not supported + """ + if provider_type not in cls._providers: + available = ', '.join(cls._providers.keys()) + raise ValueError( + f"Unknown content source provider type: {provider_type}. " + f"Available providers: {available}" + ) + + provider_class = cls._providers[provider_type] + logger.info(f"Creating content source provider: {provider_type}") + + return provider_class() + + @classmethod + def register_provider(cls, provider_type: str, provider_class: Type[ContentSourceProvider]) -> None: + """ + Register a new content source provider. + + This allows for dynamic registration of custom providers. + + Args: + provider_type: The identifier for the provider + provider_class: The provider class to register + """ + cls._providers[provider_type] = provider_class + logger.info(f"Registered content source provider: {provider_type}") + + @classmethod + def get_available_providers(cls) -> Dict[str, Type[ContentSourceProvider]]: + """ + Get all available content source providers. + + Returns: + Dictionary of provider types to their classes + """ + return cls._providers.copy() + + +# Create a singleton instance for easy import +content_source_factory = ContentSourceFactory() \ No newline at end of file diff --git a/backend/open_webui/content_sources/providers/__init__.py b/backend/open_webui/content_sources/providers/__init__.py new file mode 100644 index 0000000000..e89f7b7270 --- /dev/null +++ b/backend/open_webui/content_sources/providers/__init__.py @@ -0,0 +1 @@ +"""Content source providers.""" \ No newline at end of file diff --git a/backend/open_webui/content_sources/providers/google_drive/__init__.py b/backend/open_webui/content_sources/providers/google_drive/__init__.py new file mode 100644 index 0000000000..5b1ad5cd45 --- /dev/null +++ b/backend/open_webui/content_sources/providers/google_drive/__init__.py @@ -0,0 +1,4 @@ +"""Google Drive content source provider.""" +from .provider import GoogleDriveProvider + +__all__ = ["GoogleDriveProvider"] \ No newline at end of file diff --git a/backend/open_webui/services/google_drive.py b/backend/open_webui/content_sources/providers/google_drive/client.py similarity index 83% rename from backend/open_webui/services/google_drive.py rename to backend/open_webui/content_sources/providers/google_drive/client.py index 8e03ecf1e6..5d7ab34476 100644 --- a/backend/open_webui/services/google_drive.py +++ b/backend/open_webui/content_sources/providers/google_drive/client.py @@ -79,10 +79,13 @@ class GoogleDriveService: def __init__(self) -> None: self.service: Optional[Resource] = None - self._initialize_service() + self._credentials: Optional[service_account.Credentials] = None + self._service_initialized: bool = False + # Only prepare credentials, don't build the service yet + self._prepare_credentials() - def _initialize_service(self) -> None: - """Initialize Google Drive service with service account credentials.""" + def _prepare_credentials(self) -> None: + """Prepare credentials without building the service.""" try: if not GOOGLE_DRIVE_SERVICE_ACCOUNT_JSON.value: log.warning("Google Drive service account JSON not configured") @@ -91,28 +94,48 @@ class GoogleDriveService: # Parse service account JSON service_account_info = json.loads(GOOGLE_DRIVE_SERVICE_ACCOUNT_JSON.value) - # Create credentials - credentials = service_account.Credentials.from_service_account_info( + # Create credentials (this is fast and doesn't make network calls) + self._credentials = service_account.Credentials.from_service_account_info( service_account_info, scopes=["https://www.googleapis.com/auth/drive.readonly"], ) - - # Build service - self.service = build("drive", "v3", credentials=credentials) - log.info("Google Drive service initialized successfully") + log.debug("Google Drive credentials prepared") except json.JSONDecodeError as e: log.error(f"Invalid service account JSON: {e}") + except Exception as e: + log.error(f"Failed to prepare Google Drive credentials: {e}") + + def _initialize_service(self) -> None: + """Initialize Google Drive service with service account credentials.""" + try: + if not self._credentials: + log.warning("Google Drive credentials not available") + return + + # Build service (this makes a network call to Google's discovery API) + self.service = build("drive", "v3", credentials=self._credentials) + self._service_initialized = True + log.info("Google Drive service initialized successfully") + except Exception as e: log.error(f"Failed to initialize Google Drive service: {e}") + + def _ensure_service(self) -> None: + """Ensure the service is initialized before use.""" + if not self._service_initialized and self._credentials: + self._initialize_service() def is_configured(self) -> bool: """Check if Google Drive service is properly configured.""" - return self.service is not None + return self._credentials is not None def refresh_configuration(self) -> None: """Refresh the service configuration. Call this when config is updated.""" - self._initialize_service() + self.service = None + self._service_initialized = False + self._credentials = None + self._prepare_credentials() def get_service_account_email(self) -> Optional[str]: """Get the service account email address.""" @@ -139,6 +162,9 @@ class GoogleDriveService: Returns: List of file information dictionaries """ + # Ensure service is initialized before use + self._ensure_service() + if not self.service: raise Exception("Google Drive service not configured") @@ -204,6 +230,9 @@ class GoogleDriveService: self, folder_id: str, path: str = "" ) -> List[GoogleDriveFile]: """Get files directly in a specific folder.""" + # Ensure service is initialized + self._ensure_service() + files = [] page_token = None @@ -270,8 +299,14 @@ class GoogleDriveService: # Filter out folders and only keep files for item in items: if item.get("mimeType") != "application/vnd.google-apps.folder": + file_id = item["id"] + # Ensure file ID is clean (no extra parameters) + if "?" in file_id or "&" in file_id: + log.warning(f"File ID contains query parameters: {file_id}") + file_id = file_id.split("?")[0].split("&")[0] + file_info: GoogleDriveFile = { - "id": item["id"], + "id": file_id, "name": item["name"], "mimeType": item["mimeType"], "modifiedTime": item["modifiedTime"], @@ -283,7 +318,7 @@ class GoogleDriveService: if "parents" in item: file_info["parents"] = item["parents"] files.append(file_info) - log.info(f"Google Drive: Added file '{item['name']}' to results") + log.info(f"Google Drive: Added file '{item['name']}' (ID: {file_id}) to results") page_token = results.get("nextPageToken") if not page_token: @@ -293,6 +328,9 @@ class GoogleDriveService: def _get_subfolders(self, folder_id: str) -> List[Dict[str, str]]: """Get subfolders in a specific folder.""" + # Ensure service is initialized + self._ensure_service() + subfolders = [] page_token = None @@ -451,12 +489,39 @@ class GoogleDriveService: Returns: Tuple of (file_content_bytes, filename) """ + # Ensure service is initialized + self._ensure_service() + if not self.service: raise Exception("Google Drive service not configured") try: mime_type = file_info.get("mimeType", "") file_name = file_info.get("name", "") + + # Log file details for debugging + log.info(f"Attempting to download file: '{file_name}' (ID: {file_id}, Type: {mime_type})") + + # First, verify we can access the file's metadata with full details + try: + assert self.service is not None + file_check = self.service.files().get( + fileId=file_id, + fields="id,name,mimeType,permissions,capabilities,driveId,teamDriveId,parents", + supportsAllDrives=True + ).execute() + + log.info(f"File metadata check successful. Capabilities: {file_check.get('capabilities', {})}") + + # Check if we have export/download capability + can_download = file_check.get('capabilities', {}).get('canDownload', True) + if not can_download: + log.warning(f"File '{file_name}' cannot be downloaded due to permissions") + raise Exception("File cannot be downloaded due to permissions") + + except HttpError as e: + log.error(f"Failed to access file metadata: {e}") + raise # Handle Google Workspace files (export) if mime_type.startswith("application/vnd.google-apps"): @@ -519,5 +584,4 @@ class GoogleDriveService: return file_content, file_name -# Global instance -google_drive_service = GoogleDriveService() +# Note: Provider should create instance as needed, not use singleton diff --git a/backend/open_webui/content_sources/providers/google_drive/provider.py b/backend/open_webui/content_sources/providers/google_drive/provider.py new file mode 100644 index 0000000000..42c523eb26 --- /dev/null +++ b/backend/open_webui/content_sources/providers/google_drive/provider.py @@ -0,0 +1,460 @@ +""" +Google Drive Content Source Provider + +Implements the ContentSourceProvider interface for Google Drive integration. +""" + +import logging +from typing import List, Dict, Any, Optional, AsyncGenerator +import io + +from .client import GoogleDriveService +from open_webui.content_sources.base import ContentSourceProvider + +logger = logging.getLogger(__name__) + + +class GoogleDriveProvider(ContentSourceProvider): + """ + Google Drive content source provider implementation. + + Uses the existing GoogleDriveService for API interactions and + emits hooks for knowledge base integration. + """ + + def __init__(self): + super().__init__() + self.service = GoogleDriveService() + + def refresh_configuration(self) -> None: + """Refresh the Google Drive service configuration.""" + self.service.refresh_configuration() + + async def list_files(self, path: str = "", recursive: bool = True) -> List[Dict[str, Any]]: + """ + List files in a Google Drive folder. + + Args: + path: The folder ID to list files from + recursive: Whether to include files from nested folders + + Returns: + List of file metadata dictionaries + """ + if not self.service.is_configured(): + raise ValueError("Google Drive service is not configured") + + # Use the existing service method + gdrive_files = self.service.list_folder_files(path, include_nested=recursive) + + # Convert to our standard format + files = [] + for gdrive_file in gdrive_files: + files.append({ + 'id': gdrive_file['id'], + 'name': gdrive_file['name'], + 'mime_type': gdrive_file['mimeType'], + 'modified_time': gdrive_file['modifiedTime'], + 'size': gdrive_file.get('size'), + 'path': gdrive_file.get('path', ''), + 'provider': 'google_drive', + 'metadata': gdrive_file # Keep original metadata + }) + + return files + + async def download_file(self, file_id: str, chunk_size: int = 1024 * 1024) -> AsyncGenerator[bytes, None]: + """ + Download a file from Google Drive with chunked streaming. + + Args: + file_id: Google Drive file ID + chunk_size: Size of chunks to yield (default 1MB) + + Yields: + File content in chunks + """ + if not self.service.is_configured(): + raise ValueError("Google Drive service is not configured") + + try: + # Use the existing service method to get file info with shared drive support + file_info = self.service.service.files().get( + fileId=file_id, + supportsAllDrives=True + ).execute() + + # Check if it's a Google Workspace file that needs export + mime_type = file_info.get('mimeType', '') + file_name = file_info.get('name', 'unknown') + file_size = int(file_info.get('size', 0)) + + # Log file info for debugging + logger.info(f"Downloading file: {file_name} (ID: {file_id}, Type: {mime_type}, Size: {file_size} bytes)") + except Exception as e: + logger.error(f"Failed to get file info for {file_id}: {e}") + raise + + # Define export formats for Google Workspace files + # Prioritize text formats for knowledge base usage - text is searchable, PDF would be base64 + export_formats = { + 'application/vnd.google-apps.document': [ + 'text/plain', # Primary: Plain text for knowledge base search + 'application/pdf' # Fallback: PDF if text export fails + ], + 'application/vnd.google-apps.spreadsheet': [ + 'text/csv' # CSV is already a text format + ], + 'application/vnd.google-apps.presentation': [ + 'application/pdf', # Keep PDF for presentations (no text export available) + 'text/plain' # Fallback attempt (may not work for all slides) + ], + } + + if mime_type in export_formats: + # Google Workspace files need to be exported + async for chunk in self._download_workspace_file(file_id, file_name, export_formats[mime_type], chunk_size): + yield chunk + else: + # Regular files can be streamed directly + async for chunk in self._download_regular_file(file_id, file_name, chunk_size): + yield chunk + + async def _download_workspace_file(self, file_id: str, file_name: str, export_mime_types: list, chunk_size: int) -> AsyncGenerator[bytes, None]: + """ + Download and export a Google Workspace file with streaming. + + Args: + file_id: Google Drive file ID + file_name: Name of the file for logging + export_mime_types: List of MIME types to try for export + chunk_size: Size of chunks to yield + + Yields: + File content in chunks + """ + from googleapiclient.http import MediaIoBaseDownload + + last_error = None + content_downloaded = False + + for export_mime_type in export_mime_types: + try: + request = self.service.service.files().export_media( + fileId=file_id, + mimeType=export_mime_type + ) + logger.info(f"Attempting to export '{file_name}' (ID: {file_id}) as {export_mime_type}") + + # Use a custom IO stream that yields chunks + class ChunkedDownloadStream: + def __init__(self, chunk_size): + self.chunks = [] + self.chunk_size = chunk_size + self.position = 0 + + def write(self, data): + self.chunks.append(data) + return len(data) + + def read_chunks(self): + """Read and yield accumulated chunks.""" + if self.chunks: + data = b''.join(self.chunks) + self.chunks = [] + + # Yield data in specified chunk sizes + for i in range(0, len(data), self.chunk_size): + yield data[i:i + self.chunk_size] + + stream = ChunkedDownloadStream(chunk_size) + downloader = MediaIoBaseDownload(stream, request, chunksize=chunk_size) + + done = False + while not done: + status, done = downloader.next_chunk() + if status: + logger.debug(f"Export progress for '{file_name}': {int(status.progress() * 100)}%") + + # Yield accumulated chunks + for chunk in stream.read_chunks(): + yield chunk + + # Yield any remaining data + for chunk in stream.read_chunks(): + yield chunk + + logger.info(f"Successfully exported '{file_name}' as {export_mime_type}") + content_downloaded = True + break + + except Exception as e: + last_error = e + # Check if it's a size limit error + if hasattr(e, 'resp') and e.resp.status == 403: + error_content = e.content.decode('utf-8') if hasattr(e, 'content') else str(e) + if 'exportSizeLimitExceeded' in error_content: + logger.warning(f"File '{file_name}' too large for {export_mime_type} export (limit ~10MB), trying next format...") + continue + # For other errors, log but try next format + logger.warning(f"Failed to export '{file_name}' as {export_mime_type}: {e}") + + if not content_downloaded: + # All formats failed + error_msg = f"Failed to export '{file_name}' in any format" + if last_error: + error_msg += f": {last_error}" + raise ValueError(error_msg) + + async def _download_regular_file(self, file_id: str, file_name: str, chunk_size: int) -> AsyncGenerator[bytes, None]: + """ + Download a regular (non-Workspace) file from Google Drive with streaming. + + Args: + file_id: Google Drive file ID + file_name: Name of the file for logging + chunk_size: Size of chunks to yield + + Yields: + File content in chunks + """ + from googleapiclient.http import MediaIoBaseDownload + + request = self.service.service.files().get_media(fileId=file_id) + + # Use chunked download for memory efficiency + class ChunkedDownloadStream: + def __init__(self, chunk_size): + self.buffer = bytearray() + self.chunk_size = chunk_size + + def write(self, data): + self.buffer.extend(data) + return len(data) + + def read_chunks(self): + """Read and yield accumulated chunks.""" + while len(self.buffer) >= self.chunk_size: + chunk = bytes(self.buffer[:self.chunk_size]) + del self.buffer[:self.chunk_size] + yield chunk + + def read_remaining(self): + """Read any remaining data.""" + if self.buffer: + yield bytes(self.buffer) + self.buffer.clear() + + stream = ChunkedDownloadStream(chunk_size) + downloader = MediaIoBaseDownload(stream, request, chunksize=chunk_size) + + done = False + while not done: + status, done = downloader.next_chunk() + if status: + logger.debug(f"Download progress for '{file_name}': {int(status.progress() * 100)}%") + + # Yield accumulated chunks + for chunk in stream.read_chunks(): + yield chunk + + # Yield any remaining data + for chunk in stream.read_remaining(): + yield chunk + + logger.info(f"Successfully downloaded '{file_name}' (ID: {file_id})") + + async def get_service_info(self) -> Dict[str, Any]: + """ + Get information about the Google Drive service. + + Returns: + Dictionary with service information + """ + if not self.service.is_configured(): + return { + 'configured': False, + 'error': 'Google Drive service is not configured' + } + + email = self.service.get_service_account_email() + + return { + 'configured': True, + 'provider': 'google_drive', + 'service_account_email': email, + 'scopes': ['https://www.googleapis.com/auth/drive.readonly'] + } + + async def sync_folder_with_metadata( + self, + folder_id: str, + existing_files: Dict[str, Any], + context: Optional[Dict[str, Any]] = None + ) -> None: + """ + Sync a Google Drive folder with change detection. + + This method extends the base sync_folder to handle: + - Change detection based on modification times + - File updates and removals + - Progress tracking + + Args: + folder_id: Google Drive folder ID + existing_files: Map of Google Drive file IDs to local file metadata + context: Optional context to pass through hooks + """ + context = context or {} + + # Emit sync started event + await self.emit_hook('sync_started', { + 'folder_id': folder_id, + 'context': context + }) + + try: + # List files in the folder + gdrive_files = await self.list_files(folder_id, recursive=context.get('include_nested', True)) + + logger.info(f"Google Drive sync: Found {len(gdrive_files)} files in folder {folder_id}") + + # Track files to process + gdrive_file_ids = set() + files_to_process = [] + + for file_info in gdrive_files: + gdrive_id = file_info['id'] + gdrive_file_ids.add(gdrive_id) + + # Check if file exists and needs updating + if gdrive_id in existing_files: + existing_file = existing_files[gdrive_id] + local_modified = existing_file.get('modified_time') + gdrive_modified = file_info['modified_time'] + + if local_modified != gdrive_modified: + # File was modified, mark for update + await self.emit_hook('file_updated', { + 'file_info': file_info, + 'old_file': existing_file, + 'context': context + }) + files_to_process.append(file_info) + else: + # File is up to date + await self.emit_hook('file_unchanged', { + 'file_info': file_info, + 'context': context + }) + else: + # New file + await self.emit_hook('file_new', { + 'file_info': file_info, + 'context': context + }) + files_to_process.append(file_info) + + # Identify removed files + existing_gdrive_ids = set(existing_files.keys()) + removed_file_ids = existing_gdrive_ids - gdrive_file_ids + + for removed_id in removed_file_ids: + await self.emit_hook('file_removed', { + 'file_id': removed_id, + 'file_info': existing_files[removed_id], + 'context': context + }) + + # Process new and updated files + for i, file_info in enumerate(files_to_process): + try: + # Emit progress + await self.emit_hook('sync_progress', { + 'current': i + 1, + 'total': len(files_to_process), + 'file_name': file_info['name'], + 'context': context + }) + + # Skip Google Workspace files that can't be processed + mime_type = file_info['mime_type'] + if mime_type in [ + 'application/vnd.google-apps.site', + 'application/vnd.google-apps.form', + 'application/vnd.google-apps.map', + 'application/vnd.google-apps.drawing' + ]: + logger.warning(f"Skipping unsupported Google Workspace file: {file_info['name']}") + await self.emit_hook('file_skipped', { + 'file_info': file_info, + 'reason': 'Unsupported Google Workspace file type', + 'context': context + }) + continue + + # Download file content + content_chunks = [] + async for chunk in self.download_file(file_info['id']): + content_chunks.append(chunk) + content = b''.join(content_chunks) + + original_mime = file_info.get('mime_type', '') + if original_mime == 'application/vnd.google-apps.document': + # Google Docs are exported as text for knowledge base + file_info['mime_type'] = 'text/plain' + file_info['export_format'] = 'text' + elif original_mime == 'application/vnd.google-apps.spreadsheet': + # Sheets are exported as CSV + file_info['mime_type'] = 'text/csv' + file_info['export_format'] = 'csv' + elif original_mime == 'application/vnd.google-apps.presentation': + # Presentations are still exported as PDF (no text export available) + file_info['mime_type'] = 'application/pdf' + file_info['export_format'] = 'pdf' + + # Log content info for debugging + logger.info(f"Downloaded {file_info['name']}: {len(content)} bytes, original_mime: {original_mime}, final_mime: {file_info.get('mime_type', 'unknown')}") + + # Skip if content is empty + if not content: + logger.warning(f"File {file_info['name']} has no content, skipping") + await self.emit_hook('file_skipped', { + 'file_info': file_info, + 'reason': 'Empty content', + 'context': context + }) + continue + + # Emit file ready event + await self.emit_hook('file_ready', { + 'file_info': file_info, + 'content': content, + 'context': context + }) + + except Exception as e: + logger.error(f"Error processing file {file_info['name']}: {e}") + await self.emit_hook('file_error', { + 'file_info': file_info, + 'error': str(e), + 'context': context + }) + + # Emit sync completed event + await self.emit_hook('sync_completed', { + 'folder_id': folder_id, + 'total_files': len(gdrive_files), + 'processed_files': len(files_to_process), + 'removed_files': len(removed_file_ids), + 'context': context + }) + + except Exception as e: + logger.error(f"Error during sync: {e}") + await self.emit_hook('sync_error', { + 'folder_id': folder_id, + 'error': str(e), + 'context': context + }) + raise \ No newline at end of file diff --git a/backend/open_webui/content_sources/registry.py b/backend/open_webui/content_sources/registry.py new file mode 100644 index 0000000000..2b3da78d02 --- /dev/null +++ b/backend/open_webui/content_sources/registry.py @@ -0,0 +1,125 @@ +""" +Content Source Registry + +Global registry for content source providers with hook support. +""" + +from typing import Dict, List, Optional, Any, Callable +import logging +import asyncio + +from .base import ContentSourceProvider + +logger = logging.getLogger(__name__) + + +class ContentSourceRegistry: + """Registry for content source providers with global hook support.""" + + def __init__(self): + self._providers: Dict[str, ContentSourceProvider] = {} + self._global_hooks: Dict[str, List[Callable]] = {} + + def register_provider(self, name: str, provider: ContentSourceProvider) -> None: + """Register a content source provider.""" + self._providers[name] = provider + logger.info(f"Registered content source provider: {name}") + + def unregister_provider(self, name: str) -> None: + """Unregister a content source provider.""" + if name in self._providers: + self._providers.pop(name) + logger.info(f"Unregistered content source provider: {name}") + + def get_provider(self, name: str) -> Optional[ContentSourceProvider]: + """Get a content source provider by name.""" + return self._providers.get(name) + + def get_all_providers(self) -> Dict[str, ContentSourceProvider]: + """Get all registered providers.""" + return self._providers.copy() + + def register_global_hook(self, event: str, handler: Callable) -> None: + """ + Register a global hook that will be called for all providers. + + Args: + event: The event name to listen for + handler: Async callable that will be invoked when the event is emitted + """ + if event not in self._global_hooks: + self._global_hooks[event] = [] + self._global_hooks[event].append(handler) + logger.debug(f"Registered global hook for event '{event}'") + + async def emit_hook(self, event: str, data: Dict[str, Any]) -> None: + """ + Emit a hook event to all global handlers and registered providers. + + Args: + event: The event name + data: Data to pass to the handlers + """ + # Call global hooks first + for handler in self._global_hooks.get(event, []): + try: + if asyncio.iscoroutinefunction(handler): + await handler(data) + else: + handler(data) + except Exception as e: + logger.error(f"Error in global hook handler for event '{event}': {e}") + + # Then emit to all registered providers + for provider_name, provider in self._providers.items(): + try: + await provider.emit_hook(event, data) + except Exception as e: + logger.error(f"Error emitting hook '{event}' to provider '{provider_name}': {e}") + + +# Global registry instance +content_source_registry = ContentSourceRegistry() + +# Define standard hook events +HOOK_EVENTS = { + # File operations + 'before_file_add': 'Before file is added to knowledge base', + 'after_file_add': 'After file is added to knowledge base', + 'before_file_remove': 'Before file is removed from knowledge base', + 'after_file_remove': 'After file is removed from knowledge base', + 'before_file_update': 'Before file is updated in knowledge base', + 'after_file_update': 'After file is updated in knowledge base', + + # Knowledge base operations + 'before_knowledge_create': 'Before knowledge base is created', + 'after_knowledge_create': 'After knowledge base is created', + 'before_knowledge_update': 'Before knowledge base is updated', + 'after_knowledge_update': 'After knowledge base is updated', + 'before_knowledge_delete': 'Before knowledge base is deleted', + 'after_knowledge_delete': 'After knowledge base is deleted', + 'before_knowledge_reset': 'Before knowledge base is reset', + 'after_knowledge_reset': 'After knowledge base is reset', + + # Batch operations + 'before_files_batch_add': 'Before multiple files are added', + 'after_files_batch_add': 'After multiple files are added', + + # Sync operations (for content sources) + 'sync_started': 'Sync operation started', + 'sync_completed': 'Sync operation completed', + 'sync_error': 'Error during sync operation', + 'file_ready': 'File downloaded and ready for processing', + 'file_error': 'Error processing file', + + # Content sync operations + 'before_content_sync': 'Before content sync starts', + 'after_content_sync': 'After content sync completes', + 'content_sync_error': 'Error during content sync', + 'sync_progress': 'Progress update during sync', + 'file_new': 'New file detected during sync', + 'file_updated': 'File updated during sync', + 'file_unchanged': 'File unchanged during sync', + 'file_removed': 'File removed detected during sync', + 'file_skipped': 'File skipped during sync', +} \ No newline at end of file diff --git a/backend/open_webui/content_sources/scheduler.py b/backend/open_webui/content_sources/scheduler.py new file mode 100644 index 0000000000..12a6abe45b --- /dev/null +++ b/backend/open_webui/content_sources/scheduler.py @@ -0,0 +1,261 @@ +"""Generic content source sync scheduler for all providers.""" +import asyncio +import logging +import time +from typing import Dict, List, Any, Optional +from datetime import datetime, timedelta + +from open_webui.env import ENV +from open_webui.models.knowledge import Knowledges +from open_webui.content_sources import content_source_factory, content_source_registry +from open_webui.models.files import Files + +log = logging.getLogger(__name__) + + +class ContentSourceScheduler: + """Generic background scheduler for automatic content source sync. + + This scheduler handles sync for all content source providers (Google Drive, + Dropbox, OneDrive, etc.) in a unified way. + """ + + def __init__(self) -> None: + self.running: bool = False + self.sync_tasks: Dict[str, Any] = {} + self.scheduler_task: Optional[asyncio.Task] = None + # In dev environment, check every minute for faster testing + # In production, check every hour + self.check_interval: int = 60 if ENV == "dev" else 3600 + + async def start(self) -> None: + """Start the background sync scheduler.""" + if self.running: + log.warning("Sync scheduler already running") + return + + self.running = True + log.info("Starting content source sync scheduler") + + # Start the main scheduler loop as a background task + self.scheduler_task = asyncio.create_task(self._scheduler_loop()) + + async def stop(self) -> None: + """Stop the background sync scheduler.""" + log.info("Stopping content source sync scheduler") + self.running = False + + # Cancel the scheduler task if it exists + if self.scheduler_task and not self.scheduler_task.done(): + self.scheduler_task.cancel() + try: + await self.scheduler_task + except asyncio.CancelledError: + pass + + # Cancel all running sync tasks + for task_id, task in self.sync_tasks.items(): + if not task.done(): + task.cancel() + log.info(f"Cancelled sync task: {task_id}") + + async def _scheduler_loop(self) -> None: + """Main scheduler loop that checks for sync tasks.""" + while self.running: + try: + await self._check_and_sync_knowledge_bases() + except Exception as e: + log.error(f"Error in sync scheduler loop: {e}", exc_info=True) + + # Wait before next check + await asyncio.sleep(self.check_interval) + + async def _check_and_sync_knowledge_bases(self) -> None: + """Check all knowledge bases for content source sync needs.""" + try: + # Run synchronous database operation in a thread pool to avoid blocking + import asyncio + loop = asyncio.get_event_loop() + knowledge_bases = await loop.run_in_executor(None, Knowledges.get_knowledge_bases) + except Exception as e: + log.error(f"Error fetching knowledge bases: {e}") + return + + if ENV == "dev": + log.info(f"Checking {len(knowledge_bases)} knowledge bases for sync needs") + + for kb in knowledge_bases: + if not kb.data: + continue + + # Check for sync metadata (new structure) + sync_metadata = kb.data.get("sync_metadata", {}) + if not sync_metadata: + continue + + if ENV == "dev": + log.info(f"KB {kb.id} has sync_metadata for providers: {list(sync_metadata.keys())}") + + # Check each provider's sync configuration + for provider_name, provider_config in sync_metadata.items(): + if not provider_config: + continue + + # Get sync interval from options + options = provider_config.get("options", {}) + sync_interval_days = options.get("sync_interval_days", 1) + last_sync = provider_config.get("last_sync", 0) + + # Convert sync interval to seconds + sync_interval_seconds = sync_interval_days * 24 * 60 * 60 + current_time = time.time() + + # Log for debugging in dev mode + if ENV == "dev": + time_since_last_sync = current_time - last_sync + time_until_next_sync = sync_interval_seconds - time_since_last_sync + log.info(f"KB {kb.id} - Provider {provider_name}: " + f"sync_interval_days={sync_interval_days} ({sync_interval_seconds}s), " + f"last_sync={datetime.fromtimestamp(last_sync).isoformat() if last_sync else 'never'}, " + f"time_since_last={time_since_last_sync:.1f}s, " + f"time_until_next={time_until_next_sync:.1f}s, " + f"needs_sync={time_since_last_sync >= sync_interval_seconds}") + + if current_time - last_sync >= sync_interval_seconds: + # Schedule sync task + task_id = f"{kb.id}_{provider_name}_{int(current_time)}" + + if task_id not in self.sync_tasks or self.sync_tasks[task_id].done(): + log.info(f"Scheduling sync for knowledge base {kb.id} with provider {provider_name}") + + # Build content source config for sync + content_source_config = { + "provider": provider_name, + "source_id": provider_config.get("source_id"), + "options": options + } + + task = asyncio.create_task( + self._sync_knowledge_base(kb.id, provider_name, content_source_config) + ) + self.sync_tasks[task_id] = task + + async def _sync_knowledge_base( + self, + kb_id: str, + provider_name: str, + content_source_config: Dict[str, Any] + ) -> None: + """Sync a specific knowledge base with its content source provider.""" + try: + log.info(f"Starting scheduled sync for knowledge base {kb_id} with {provider_name}") + + # Import the sync service + from open_webui.content_sources.syncer import content_syncer + + # Get knowledge base to get user ID (run in thread pool to avoid blocking) + loop = asyncio.get_event_loop() + kb = await loop.run_in_executor(None, Knowledges.get_knowledge_by_id, kb_id) + if not kb: + log.error(f"Knowledge base {kb_id} not found") + return + + # Perform actual sync using the sync service + source_id = content_source_config.get("source_id") + options = content_source_config.get("options", {}) + + # Add auto_sync flag for scheduled syncs + options["auto_sync"] = True + options["rollback_on_error"] = False # Don't rollback on scheduled syncs + + try: + # Use the sync service to perform the actual sync + sync_result = await content_syncer.sync_provider_files( + provider_name=provider_name, + source_id=source_id, + options=options, + request=None, # No request object in scheduler context + user_id=kb.user_id, + knowledge_base_id=kb_id + ) + + # Update knowledge base with sync results + data = kb.data or {} + + # Get event loop for async operations + loop = asyncio.get_event_loop() + + # Update file IDs with successfully synced files + existing_file_ids = set(data.get("file_ids", [])) + new_file_ids = set(sync_result.successful_files) + updated_file_ids = list(existing_file_ids | new_file_ids) + data["file_ids"] = updated_file_ids + + # Update sync metadata in new structure + data.setdefault("sync_metadata", {}) + data["sync_metadata"].setdefault(provider_name, {}) + + # Preserve existing configuration + existing_config = data["sync_metadata"][provider_name] + + # Update with new sync results + data["sync_metadata"][provider_name].update({ + "source_id": source_id, + "last_sync": time.time(), + "options": options, + "results": { + "status": sync_result.status.value, + "added": len(sync_result.added_files), + "updated": len(sync_result.updated_files), + "failed": len(sync_result.failed_files), + "duplicates": len(sync_result.duplicate_files), + "changes": sync_result.changes + } + }) + + # Store any errors for monitoring + if sync_result.errors: + data["sync_metadata"][provider_name]["last_sync_errors"] = sync_result.errors[:5] # Keep last 5 errors + + await loop.run_in_executor(None, Knowledges.update_knowledge_data_by_id, kb_id, data) + + log.info(f"Completed scheduled sync for knowledge base {kb_id}: " + f"status={sync_result.status.value}, " + f"added={len(sync_result.added_files)}, " + f"updated={len(sync_result.updated_files)}, " + f"failed={len(sync_result.failed_files)}") + + except Exception as sync_error: + log.error(f"Sync service error for knowledge base {kb_id}: {sync_error}") + + # Get event loop for async operations + loop = asyncio.get_event_loop() + + # Update knowledge base with error status + data = kb.data or {} + data.setdefault("sync_metadata", {}) + data["sync_metadata"].setdefault(provider_name, {}) + + # Preserve existing configuration + existing_config = data["sync_metadata"][provider_name] + source_id = existing_config.get("source_id", source_id) + existing_options = existing_config.get("options", {}) + + data["sync_metadata"][provider_name].update({ + "source_id": source_id, + "last_sync": time.time(), + "options": existing_options, + "results": { + "status": "failed", + "error": str(sync_error) + } + }) + await loop.run_in_executor(None, Knowledges.update_knowledge_data_by_id, kb_id, data) + raise + + except Exception as e: + log.error(f"Error in scheduled sync for knowledge base {kb_id}: {e}", exc_info=True) + + +# Global scheduler instance +scheduler = ContentSourceScheduler() \ No newline at end of file diff --git a/backend/open_webui/content_sources/syncer.py b/backend/open_webui/content_sources/syncer.py new file mode 100644 index 0000000000..a5b29fff4e --- /dev/null +++ b/backend/open_webui/content_sources/syncer.py @@ -0,0 +1,1385 @@ +""" +Syncs content from external providers (Google Drive, OneDrive, etc.) + +Handles: +- Downloading files from providers with retry logic +- Creating/updating files with deterministic IDs +- Batch operations with rollback on failure +- Simple duplicate detection via file hashes +- Knowledge base synchronization +""" + +import logging +import time +import hashlib +import io +import asyncio +from typing import Dict, List, Optional, Any, Tuple, Union, Set +from dataclasses import dataclass +from enum import Enum +from contextlib import asynccontextmanager +from fastapi import HTTPException, Request, status + +from open_webui.models.files import Files, FileModel, FileForm +from open_webui.models.knowledge import Knowledges +from open_webui.storage.provider import Storage +from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT +from open_webui.routers.retrieval import process_file, ProcessFileForm +from open_webui.utils.misc import calculate_sha256_string +from open_webui.constants import ERROR_MESSAGES +from open_webui.content_sources import content_source_registry, content_source_factory + +# Import PDF text extraction capability +try: + from langchain_community.document_loaders import PyPDFLoader + import tempfile + PDF_EXTRACTION_AVAILABLE = True +except ImportError: + PDF_EXTRACTION_AVAILABLE = False + +# Configure logging +log = logging.getLogger(__name__) + + +class SyncAction(Enum): + """Enumeration of possible sync actions.""" + CREATED = "created" + UPDATED = "updated" + UNCHANGED = "unchanged" + SKIPPED = "skipped" + FAILED = "failed" + + +class SyncStatus(Enum): + """Enumeration of sync operation statuses.""" + SUCCESS = "success" + PARTIAL_SUCCESS = "partial_success" + FAILED = "failed" + ROLLED_BACK = "rolled_back" + + +@dataclass +class SyncResult: + """Result of a single file sync operation.""" + file_id: Optional[str] + action: SyncAction + error: Optional[str] = None + warnings: List[str] = None + metadata: Optional[Dict[str, Any]] = None + + def __post_init__(self): + if self.warnings is None: + self.warnings = [] + + +@dataclass +class BatchSyncResult: + """Result of a batch sync operation.""" + status: SyncStatus + added_files: List[str] + updated_files: List[str] + unchanged_files: List[str] + removed_files: List[str] # Track files that were removed + duplicate_files: List[Dict[str, Any]] + failed_files: List[Dict[str, Any]] + errors: List[Dict[str, str]] + warnings: List[str] + total_processed: int + changes: bool + rollback_performed: bool = False + + @property + def successful_files(self) -> List[str]: + """Get list of all successfully processed files.""" + return self.added_files + self.updated_files + self.unchanged_files + + +@dataclass +class KnowledgeSyncResult: + """Result of a knowledge base sync operation.""" + knowledge_base_id: str + sync_results: BatchSyncResult + file_ids_updated: List[str] + metadata_updated: bool + sync_timestamp: int + + +@dataclass +class DuplicateCheckResult: + """Result of duplicate content checking.""" + is_duplicate: bool + existing_file_ids: List[str] + hash_checked: Optional[str] = None + collection_name: Optional[str] = None + + +class ContentSourceSyncer: + """ + Central service for content synchronization operations. + + This service handles: + - File creation and updates from various content sources + - Knowledge base synchronization with providers + - Duplicate detection and prevention + - Rollback capabilities for failed operations + - Comprehensive error handling and logging + """ + + def __init__(self): + self.storage = Storage + self._active_operations: Set[str] = set() + + def _extract_text_from_pdf(self, pdf_content: bytes) -> Optional[str]: + """ + Extract text from PDF content using PyPDFLoader. + + Args: + pdf_content: Binary PDF content + + Returns: + Extracted text as string, or None if extraction fails + """ + if not PDF_EXTRACTION_AVAILABLE: + log.warning("PDF text extraction not available - PyPDFLoader not installed") + return None + + try: + # Create a temporary file to store PDF content + with tempfile.NamedTemporaryFile(suffix='.pdf', delete=False) as tmp_file: + tmp_file.write(pdf_content) + tmp_file_path = tmp_file.name + + try: + # Use PyPDFLoader to extract text + loader = PyPDFLoader(tmp_file_path) + documents = loader.load() + + # Combine all pages into single text + text_content = "\n\n".join([doc.page_content for doc in documents]) + + # Clean up empty lines and excessive whitespace + text_content = "\n".join(line for line in text_content.split("\n") if line.strip()) + + if text_content.strip(): + log.info(f"Successfully extracted {len(text_content)} characters from PDF") + return text_content + else: + log.warning("PDF text extraction produced empty content") + return None + + finally: + # Clean up temporary file + import os + try: + os.unlink(tmp_file_path) + except: + pass + + except Exception as e: + log.error(f"Failed to extract text from PDF: {e}") + return None + + async def sync_file(self, file_id: str, request: Request, user_id: str) -> SyncResult: + """ + Synchronize a single file by re-processing it into the vector database. + + Args: + file_id: The ID of the file to sync + request: FastAPI request object for processing context + user_id: ID of the user performing the sync + + Returns: + SyncResult containing the outcome of the sync operation + + Raises: + HTTPException: If file not found or sync fails + """ + operation_id = f"sync_file_{file_id}_{int(time.time())}" + + try: + # Prevent concurrent operations on the same file + if file_id in self._active_operations: + log.warning(f"Sync already in progress for file {file_id}") + return SyncResult( + file_id=file_id, + action=SyncAction.SKIPPED, + error="Sync already in progress" + ) + + self._active_operations.add(file_id) + log.info(f"Starting sync for file {file_id} (operation: {operation_id})") + + # Get the file + file = Files.get_file_by_id(file_id) + if not file: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="File not found" + ) + + # Check if file has content to process + if not file.data or not file.data.get("content"): + log.warning(f"File {file_id} has no content to process") + return SyncResult( + file_id=file_id, + action=SyncAction.SKIPPED, + warnings=["File has no content to process"] + ) + + # Process the file into vector database + try: + collection_name = f"file-{file_id}" + process_file( + request, + ProcessFileForm(file_id=file_id, collection_name=collection_name), + user=type('User', (), {'id': user_id})() + ) + + log.info(f"Successfully synced file {file_id}") + return SyncResult( + file_id=file_id, + action=SyncAction.UPDATED, + metadata={"collection_name": collection_name} + ) + + except Exception as e: + error_msg = str(e) + log.error(f"Failed to process file {file_id}: {error_msg}") + return SyncResult( + file_id=file_id, + action=SyncAction.FAILED, + error=error_msg + ) + + except HTTPException: + raise + except Exception as e: + error_context = { + 'file_id': file_id, + 'operation': 'sync_file', + 'user_id': user_id, + 'error_type': type(e).__name__, + 'error_details': str(e) + } + log.error(f"Unexpected error syncing file {file_id}: {error_context}", exc_info=True) + return SyncResult( + file_id=file_id, + action=SyncAction.FAILED, + error=f"Sync failed ({type(e).__name__}): {str(e)}", + metadata={'error_context': error_context} + ) + finally: + self._active_operations.discard(file_id) + + async def sync_provider_files( + self, + provider_name: str, + source_id: str, + options: Dict[str, Any], + request: Request, + user_id: str, + knowledge_base_id: Optional[str] = None + ) -> BatchSyncResult: + """ + Synchronize files from a content provider. + + Args: + provider_name: Name of the content provider (e.g., 'google_drive') + source_id: Provider-specific source identifier (e.g., folder ID) + options: Sync options and configuration + request: FastAPI request object + user_id: ID of the user performing the sync + knowledge_base_id: Optional knowledge base ID for direct sync + + Returns: + BatchSyncResult containing detailed sync results + + Raises: + HTTPException: If provider not found or not configured + """ + operation_id = f"sync_provider_{provider_name}_{source_id}_{int(time.time())}" + log.info(f"Starting provider sync (operation: {operation_id})") + + # Initialize result tracking + result = BatchSyncResult( + status=SyncStatus.SUCCESS, + added_files=[], + updated_files=[], + unchanged_files=[], + removed_files=[], + duplicate_files=[], + failed_files=[], + errors=[], + warnings=[], + total_processed=0, + changes=False + ) + + processed_files = [] # Track for potential rollback + + try: + # Get and validate provider with enhanced error context + try: + provider = await self._get_configured_provider(provider_name) + except Exception as e: + error_msg = f"Provider '{provider_name}' initialization failed: {str(e)}" + log.error(error_msg, extra={ + 'provider': provider_name, + 'source_id': source_id, + 'operation_id': operation_id + }) + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=error_msg + ) + + # Track existing files in knowledge base if provided + existing_provider_files = {} + if knowledge_base_id: + knowledge = Knowledges.get_knowledge_by_id(id=knowledge_base_id) + if knowledge and knowledge.data: + file_ids = knowledge.data.get("file_ids", []) + # Get files that belong to this provider + existing_files = Files.get_files_by_ids(file_ids) + for file in existing_files: + if file.provider == provider_name: + existing_provider_files[file.provider_file_id] = file.id + + # List files from provider with enhanced logging + log.info(f"Listing files from {provider_name} source {source_id}", extra={ + 'provider': provider_name, + 'source_id': source_id, + 'options': options, + 'operation_id': operation_id + }) + + try: + files = await provider.list_files(source_id, recursive=True) + result.total_processed = len(files) + + log.info(f"Found {len(files)} files to process from {provider_name}", extra={ + 'file_count': len(files), + 'provider': provider_name, + 'source_id': source_id + }) + except Exception as e: + provider_error = f"{provider_name} list_files failed" + if hasattr(e, 'resp') and hasattr(e.resp, 'status'): + provider_error += f" (HTTP {e.resp.status})" + provider_error += f": {str(e)}" + + log.error(provider_error, extra={ + 'provider': provider_name, + 'source_id': source_id, + 'error_type': type(e).__name__ + }) + raise HTTPException( + status_code=status.HTTP_502_BAD_GATEWAY, + detail=provider_error + ) + + # Track current provider file IDs + current_provider_file_ids = set() + + # Process each file + for file_info in files: + # Track this file as still existing + current_provider_file_ids.add(file_info['id']) + try: + sync_result = await self._process_single_provider_file( + provider, provider_name, file_info, options, user_id + ) + + # Track the result + if sync_result.file_id: + processed_files.append(sync_result.file_id) + + # Update batch result based on individual result + if sync_result.action == SyncAction.CREATED: + result.added_files.append(sync_result.file_id) + result.changes = True + elif sync_result.action == SyncAction.UPDATED: + result.updated_files.append(sync_result.file_id) + result.changes = True + elif sync_result.action == SyncAction.UNCHANGED: + result.unchanged_files.append(sync_result.file_id) + elif sync_result.action == SyncAction.SKIPPED: + if "duplicate" in sync_result.error.lower() if sync_result.error else False: + result.duplicate_files.append({ + 'file_id': sync_result.file_id, + 'name': file_info.get('name', 'unknown'), + 'reason': sync_result.error + }) + else: + result.warnings.extend(sync_result.warnings) + else: # FAILED + result.failed_files.append({ + 'file_id': sync_result.file_id, + 'name': file_info.get('name', 'unknown'), + 'error': sync_result.error + }) + result.errors.append({ + 'file': file_info.get('name', 'unknown'), + 'error': sync_result.error + }) + + except Exception as e: + file_name = file_info.get('name', 'unknown') + file_id = file_info.get('id', 'unknown') + + # Enhanced error with provider context and error codes + error_details = { + 'provider': provider_name, + 'file_name': file_name, + 'file_id': file_id, + 'error_type': type(e).__name__, + 'error_message': str(e) + } + + # Check for specific provider error codes + if hasattr(e, 'resp') and hasattr(e.resp, 'status'): + error_details['http_status'] = e.resp.status + if e.resp.status == 403: + error_details['error_category'] = 'permission_denied' + elif e.resp.status == 404: + error_details['error_category'] = 'not_found' + elif e.resp.status == 429: + error_details['error_category'] = 'rate_limited' + elif e.resp.status >= 500: + error_details['error_category'] = 'provider_error' + + log.error(f"Error processing {provider_name} file '{file_name}': {error_details}", + extra=error_details) + + result.failed_files.append({ + 'file_id': file_id, + 'name': file_name, + 'error': str(e), + 'error_details': error_details + }) + result.errors.append({ + 'file': file_name, + 'error': str(e), + 'provider': provider_name, + 'error_category': error_details.get('error_category', 'unknown') + }) + + # Detect and handle removed files + if knowledge_base_id and existing_provider_files: + # Find files that exist in knowledge base but not in provider anymore + removed_provider_file_ids = set(existing_provider_files.keys()) - current_provider_file_ids + + for provider_file_id in removed_provider_file_ids: + file_id = existing_provider_files[provider_file_id] + try: + # Remove file from storage and database + file = Files.get_file_by_id(file_id) + if file: + # Delete from vector database + try: + VECTOR_DB_CLIENT.delete_collection(collection_name=f"file-{file_id}") + except Exception as e: + log.warning(f"Failed to delete vector collection for file {file_id}: {e}") + + # Delete file record + if Files.delete_file_by_id(file_id): + result.removed_files.append(file_id) + result.changes = True + log.info(f"Removed file {file_id} (provider file {provider_file_id}) - no longer in source") + else: + log.warning(f"Failed to delete file record {file_id}") + except Exception as e: + log.error(f"Error removing file {file_id}: {e}") + result.errors.append({ + 'file': file_id, + 'error': f"Failed to remove: {str(e)}", + 'provider': provider_name + }) + + # Determine overall status + if result.failed_files and not (result.added_files or result.updated_files or result.removed_files): + result.status = SyncStatus.FAILED + elif result.failed_files: + result.status = SyncStatus.PARTIAL_SUCCESS + else: + result.status = SyncStatus.SUCCESS + + log.info(f"Provider sync completed: {result.status.value}, " + f"added: {len(result.added_files)}, " + f"updated: {len(result.updated_files)}, " + f"removed: {len(result.removed_files)}, " + f"failed: {len(result.failed_files)}") + + return result + + except Exception as e: + # Enhanced batch operation error with full context + batch_error = { + 'provider': provider_name, + 'source_id': source_id, + 'operation_id': operation_id, + 'files_processed': len(processed_files), + 'error_type': type(e).__name__, + 'error_message': str(e) + } + + log.error(f"Provider sync batch operation failed: {batch_error}", + exc_info=True, extra=batch_error) + + # Perform rollback if configured and we have processed files + if options.get('rollback_on_error', True) and processed_files: + log.info(f"Initiating rollback for {len(processed_files)} processed files", extra={ + 'operation_id': operation_id, + 'rollback_count': len(processed_files) + }) + rollback_success = await self._rollback_files(processed_files) + result.rollback_performed = rollback_success + + if rollback_success: + log.info(f"Rollback completed successfully", extra={'operation_id': operation_id}) + else: + log.error(f"Rollback failed or partially completed", extra={'operation_id': operation_id}) + + result.status = SyncStatus.FAILED + result.errors.append({ + 'file': 'batch_operation', + 'error': str(e), + 'provider': provider_name, + 'error_details': batch_error + }) + + return result + + async def sync_knowledge_base( + self, + kb_id: str, + request: Request, + user_id: str, + provider_configs: Optional[List[Dict[str, Any]]] = None + ) -> KnowledgeSyncResult: + """ + Synchronize an entire knowledge base, either reprocessing existing files + or syncing from configured providers. + + Args: + kb_id: Knowledge base ID + request: FastAPI request object + user_id: ID of the user performing the sync + provider_configs: Optional list of provider configurations for sync + + Returns: + KnowledgeSyncResult containing sync results and metadata + + Raises: + HTTPException: If knowledge base not found or access denied + """ + log.info(f"Starting knowledge base sync for {kb_id}") + + # Get knowledge base + knowledge = Knowledges.get_knowledge_by_id(kb_id) + if not knowledge: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Knowledge base not found" + ) + + # Get existing file IDs + data = knowledge.data or {} + existing_file_ids = data.get("file_ids", []) + + sync_timestamp = int(time.time()) + + if provider_configs: + # Sync from providers + all_results = [] + all_new_file_ids = set(existing_file_ids) + + for provider_config in provider_configs: + provider_result = await self.sync_provider_files( + provider_name=provider_config['provider'], + source_id=provider_config['source_id'], + options=provider_config.get('options', {}), + request=request, + user_id=user_id, + knowledge_base_id=kb_id + ) + all_results.append(provider_result) + all_new_file_ids.update(provider_result.successful_files) + + # Combine all results + combined_result = self._combine_batch_results(all_results) + + # Update knowledge base file IDs + updated_file_ids = list(all_new_file_ids) + data["file_ids"] = updated_file_ids + + # Update sync metadata + data.setdefault("sync_metadata", {}) + for i, provider_config in enumerate(provider_configs): + provider_name = provider_config['provider'] + data["sync_metadata"][provider_name] = { + "source_id": provider_config['source_id'], + "last_sync": sync_timestamp, + "options": provider_config.get('options', {}), + "results": { + "added": all_results[i].added_files, + "updated": all_results[i].updated_files, + "errors": len(all_results[i].errors) + } + } + + # Update knowledge base + knowledge = Knowledges.update_knowledge_data_by_id(kb_id, data) + + return KnowledgeSyncResult( + knowledge_base_id=kb_id, + sync_results=combined_result, + file_ids_updated=updated_file_ids, + metadata_updated=True, + sync_timestamp=sync_timestamp + ) + + else: + # Reprocess existing files + log.info(f"Reprocessing {len(existing_file_ids)} existing files in knowledge base {kb_id}") + + results = [] + for file_id in existing_file_ids: + sync_result = await self.sync_file(file_id, request, user_id) + results.append(sync_result) + + # Convert to batch result format + batch_result = self._convert_sync_results_to_batch(results) + + return KnowledgeSyncResult( + knowledge_base_id=kb_id, + sync_results=batch_result, + file_ids_updated=existing_file_ids, + metadata_updated=False, + sync_timestamp=sync_timestamp + ) + + def check_duplicate(self, file_hash: str, collection_name: str) -> DuplicateCheckResult: + """ + Check if a file with the given hash already exists in the vector database. + + Args: + file_hash: SHA256 hash of the file content + collection_name: Vector database collection name to check + + Returns: + DuplicateCheckResult with duplicate status and details + """ + if not file_hash or not collection_name: + return DuplicateCheckResult( + is_duplicate=False, + existing_file_ids=[], + hash_checked=file_hash, + collection_name=collection_name + ) + + try: + result = VECTOR_DB_CLIENT.query( + collection_name=collection_name, + filter={"hash": file_hash}, + ) + + existing_ids = [] + if result is not None and result.ids and result.ids[0]: + existing_ids = result.ids[0] + + return DuplicateCheckResult( + is_duplicate=len(existing_ids) > 0, + existing_file_ids=existing_ids, + hash_checked=file_hash, + collection_name=collection_name + ) + + except Exception as e: + log.error(f"Error checking duplicate in vector DB: {e}") + return DuplicateCheckResult( + is_duplicate=False, + existing_file_ids=[], + hash_checked=file_hash, + collection_name=collection_name + ) + + # Private helper methods + + async def _get_configured_provider(self, provider_name: str): + """Get and validate a content source provider.""" + # Check if provider exists + available_providers = content_source_factory.get_available_providers() + if provider_name not in available_providers: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Unknown provider: {provider_name}. Available: {', '.join(available_providers.keys())}" + ) + + # Get or create provider instance + provider = content_source_registry.get_provider(provider_name) + if not provider: + provider = content_source_factory.get_provider(provider_name) + content_source_registry.register_provider(provider_name, provider) + + # Verify provider is configured + try: + service_info = await provider.get_service_info() + if not service_info.get('configured', False): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Provider '{provider_name}' is not configured" + ) + except Exception as e: + log.error(f"Failed to check provider configuration: {e}") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"Failed to check provider configuration: {str(e)}" + ) + + return provider + + async def _process_single_provider_file( + self, + provider, + provider_name: str, + file_info: Dict[str, Any], + sync_options: Dict[str, Any], + user_id: str + ) -> SyncResult: + """Process a single file from a content provider with retry logic.""" + try: + # Download file content with exponential backoff retry + max_retries = 3 + retry_delay = 1.0 # Initial delay in seconds + + content = None + last_error = None + + for attempt in range(max_retries): + try: + content_chunks = [] + async for chunk in provider.download_file(file_info['id']): + content_chunks.append(chunk) + content = b''.join(content_chunks) + break # Success, exit retry loop + + except Exception as e: + last_error = e + if attempt < max_retries - 1: + # Exponential backoff: 1s, 2s, 4s + delay = retry_delay * (2 ** attempt) + log.warning(f"Download failed for {provider_name}/{file_info.get('name', 'unknown')} " + f"(attempt {attempt + 1}/{max_retries}): {e}. Retrying in {delay}s...") + await asyncio.sleep(delay) + else: + log.error(f"Download failed after {max_retries} attempts for " + f"{provider_name}/{file_info.get('name', 'unknown')}: {e}") + raise + + if content is None: + raise last_error or Exception("Failed to download file content") + + # Create or update file + file_id, action = await self._create_or_update_file( + user_id, provider_name, file_info, content, sync_options + ) + + if not file_id: + if action == 'skipped': + return SyncResult( + file_id=None, + action=SyncAction.SKIPPED, + error="Binary file skipped - cannot extract text for knowledge base" + ) + else: + return SyncResult( + file_id=None, + action=SyncAction.FAILED, + error="Failed to create file record" + ) + + + # Convert action string to enum + if action == 'created': + sync_action = SyncAction.CREATED + elif action == 'updated': + sync_action = SyncAction.UPDATED + elif action == 'unchanged': + sync_action = SyncAction.UNCHANGED + else: + sync_action = SyncAction.FAILED + + return SyncResult( + file_id=file_id, + action=sync_action, + metadata={ + 'provider': provider_name, + 'provider_file_id': file_info['id'], + 'name': file_info.get('name', 'unknown') + } + ) + + except Exception as e: + return SyncResult( + file_id=None, + action=SyncAction.FAILED, + error=str(e) + ) + + async def _create_or_update_file( + self, + user_id: str, + provider_name: str, + file_info: Dict[str, Any], + content: bytes, + sync_options: Optional[Dict[str, Any]] = None + ) -> Tuple[Optional[str], str]: + """ + Create or update a file with provider metadata using existing file model. + + Returns: + tuple: (file_id, action) where action is 'created', 'updated', or 'unchanged' + """ + # Generate consistent file ID based on provider and file ID + base_id = f"{provider_name}_{file_info['id']}" + file_id = hashlib.sha256(base_id.encode()).hexdigest()[:16] + + # Check if file already exists + existing_file = Files.get_file_by_id(file_id) + if existing_file: + # Check if update is needed + provider_info = existing_file.provider_info + if provider_info and provider_info.get('provider_modified_time') == file_info.get('modified_time'): + return existing_file.id, 'unchanged' + + # Update file content and metadata + file_obj = io.BytesIO(content) + _, file_path = self.storage.upload_file(file_obj, existing_file.filename, {}) + + # Update file data - only save text content for knowledge base + try: + text_content = content.decode("utf-8") if isinstance(content, bytes) else str(content) + except UnicodeDecodeError: + # Binary content - try to extract text if it's a PDF + extracted_text = None + if file_info and file_info.get('mime_type') == 'application/pdf': + extracted_text = self._extract_text_from_pdf(content) + + if extracted_text: + # Successfully extracted text from PDF + text_content = extracted_text + log.info(f"Extracted text from PDF for file {file_id}") + else: + # Skip binary files that we can't extract text from + log.warning(f"Skipping binary file {file_id} - cannot extract text for knowledge base") + return existing_file.id, 'unchanged' # Don't update with unusable content + + updated_data = { + **existing_file.data, + "content": text_content + } + Files.update_file_data_by_id(file_id, updated_data) + + # Calculate and update hash + content_hash = calculate_sha256_string(text_content) + Files.update_file_hash_by_id(file_id, content_hash) + + # Update provider info + Files.update_file_provider_info_by_id( + file_id, + provider_modified_time=file_info.get('modified_time'), + provider_metadata=file_info.get('metadata', {}) + ) + + return file_id, 'updated' + + # Create new file + filename = f"{provider_name}_{file_info['name']}" + + # Store file + file_obj = io.BytesIO(content) + _, file_path = self.storage.upload_file(file_obj, filename, {}) + + # Convert content for storage - only save text content for knowledge base + try: + text_content = content.decode("utf-8") if isinstance(content, bytes) else str(content) + except UnicodeDecodeError: + # Binary content - try to extract text if it's a PDF + extracted_text = None + if file_info and file_info.get('mime_type') == 'application/pdf': + extracted_text = self._extract_text_from_pdf(content) + + if extracted_text: + # Successfully extracted text from PDF + text_content = extracted_text + log.info(f"Extracted text from PDF for new file {filename}") + else: + # Skip binary files that we can't extract text from + log.warning(f"Skipping binary file {filename} - cannot extract text for knowledge base") + return None, 'skipped' # Don't create file with unusable content + + # Determine if sync should be enabled + enable_sync = sync_options.get('auto_sync', False) if sync_options else False + + # Create FileForm with provider fields + file_form = FileForm( + id=file_id, + filename=filename, + path=file_path, + data={"content": text_content}, + meta={ + "name": file_info['name'], + "content_type": file_info.get('mime_type', 'application/octet-stream'), + "size": len(content) + }, + provider=provider_name, + provider_file_id=file_info['id'], + provider_modified_time=file_info.get('modified_time'), + provider_sync_enabled=enable_sync, + provider_metadata=file_info.get('metadata', {}) + ) + + file_model = Files.insert_new_file(user_id=user_id, form_data=file_form) + + # Calculate and update hash immediately after creation + if file_model: + content_hash = calculate_sha256_string(text_content) + Files.update_file_hash_by_id(file_model.id, content_hash) + + return (file_model.id, 'created') if file_model else (None, None) + + async def _rollback_files(self, file_ids: List[str]) -> bool: + """ + Rollback file changes by removing them from the system. + + Args: + file_ids: List of file IDs to rollback + + Returns: + bool: True if rollback was successful, False otherwise + """ + try: + log.info(f"Rolling back {len(file_ids)} files") + + for file_id in file_ids: + try: + # Remove from vector database + collection_name = f"file-{file_id}" + if VECTOR_DB_CLIENT.has_collection(collection_name): + VECTOR_DB_CLIENT.delete_collection(collection_name) + + # Remove file record + Files.delete_file_by_id(file_id) + log.debug(f"Rolled back file {file_id}") + + except Exception as e: + log.error(f"Failed to rollback file {file_id}: {e}") + # Continue with other files + + log.info("Rollback completed") + return True + + except Exception as e: + log.error(f"Rollback operation failed: {e}") + return False + + def _combine_batch_results(self, results: List[BatchSyncResult]) -> BatchSyncResult: + """Combine multiple batch sync results into a single result.""" + combined = BatchSyncResult( + status=SyncStatus.SUCCESS, + added_files=[], + updated_files=[], + unchanged_files=[], + removed_files=[], + duplicate_files=[], + failed_files=[], + errors=[], + warnings=[], + total_processed=0, + changes=False + ) + + has_failures = False + has_successes = False + + for result in results: + combined.added_files.extend(result.added_files) + combined.updated_files.extend(result.updated_files) + combined.unchanged_files.extend(result.unchanged_files) + combined.removed_files.extend(result.removed_files) + combined.duplicate_files.extend(result.duplicate_files) + combined.failed_files.extend(result.failed_files) + combined.errors.extend(result.errors) + combined.warnings.extend(result.warnings) + combined.total_processed += result.total_processed + + if result.changes: + combined.changes = True + + if result.status == SyncStatus.FAILED: + has_failures = True + else: + has_successes = True + + # Determine combined status + if has_failures and has_successes: + combined.status = SyncStatus.PARTIAL_SUCCESS + elif has_failures: + combined.status = SyncStatus.FAILED + else: + combined.status = SyncStatus.SUCCESS + + return combined + + def _convert_sync_results_to_batch(self, sync_results: List[SyncResult]) -> BatchSyncResult: + """Convert individual sync results to a batch result format.""" + batch_result = BatchSyncResult( + status=SyncStatus.SUCCESS, + added_files=[], + updated_files=[], + unchanged_files=[], + removed_files=[], + duplicate_files=[], + failed_files=[], + errors=[], + warnings=[], + total_processed=len(sync_results), + changes=False + ) + + has_failures = False + + for result in sync_results: + if result.action == SyncAction.CREATED: + batch_result.added_files.append(result.file_id) + batch_result.changes = True + elif result.action == SyncAction.UPDATED: + batch_result.updated_files.append(result.file_id) + batch_result.changes = True + elif result.action == SyncAction.UNCHANGED: + batch_result.unchanged_files.append(result.file_id) + elif result.action == SyncAction.FAILED: + batch_result.failed_files.append({ + 'file_id': result.file_id, + 'name': result.metadata.get('name', 'unknown') if result.metadata else 'unknown', + 'error': result.error + }) + batch_result.errors.append({ + 'file': result.metadata.get('name', 'unknown') if result.metadata else 'unknown', + 'error': result.error + }) + has_failures = True + + if result.warnings: + batch_result.warnings.extend(result.warnings) + + # Set overall status + if has_failures and (batch_result.added_files or batch_result.updated_files): + batch_result.status = SyncStatus.PARTIAL_SUCCESS + elif has_failures: + batch_result.status = SyncStatus.FAILED + + return batch_result + + async def sync_file_from_provider( + self, + file_id: str, + request: Request, + user_id: str + ) -> Dict[str, Any]: + """ + Sync a single file from its configured content source provider. + + This method: + 1. Validates the file has provider information + 2. Downloads latest content from the provider + 3. Updates local file if content changed + 4. Re-processes the file for vector database + + Args: + file_id: ID of the file to sync + request: FastAPI request object + user_id: ID of the user performing the sync + + Returns: + Dict containing sync result with success status, messages, and metadata + """ + log.info(f"Starting provider sync for file {file_id}") + + # Get the file + file = Files.get_file_by_id(file_id) + if not file: + return { + "success": False, + "file_id": file_id, + "filename": "unknown", + "message": "File not found", + "error": "File not found" + } + + # Check if file has provider info + provider_info = file.provider_info + if not provider_info or not provider_info.get("provider"): + return { + "success": False, + "file_id": file_id, + "filename": file.filename, + "message": "File does not have provider information", + "error": "No provider configured for this file" + } + + provider_name = provider_info.get("provider") + provider_file_id = provider_info.get("provider_file_id") + + if not provider_file_id: + return { + "success": False, + "file_id": file_id, + "filename": file.filename, + "message": "File does not have provider file ID", + "error": "No provider file ID found" + } + + try: + # Get the content source provider + provider = await self._get_configured_provider(provider_name) + + # Get file metadata from provider to check if sync needed + files = await provider.list_files("", recursive=False) + provider_file = None + for f in files: + if f.get("id") == provider_file_id: + provider_file = f + break + + # Check if file needs sync (if we found metadata) + if provider_file: + provider_modified_time = provider_file.get("modified_time") + local_modified_time = provider_info.get("provider_modified_time") + + if provider_modified_time == local_modified_time: + return { + "success": True, + "file_id": file_id, + "filename": file.filename, + "message": "File is already up to date", + "updated": False, + "content_changed": False, + "provider_modified_time": provider_modified_time + } + + # Download the file content + content_chunks = [] + async for chunk in provider.download_file(provider_file_id): + content_chunks.append(chunk) + + if not content_chunks: + return { + "success": False, + "file_id": file_id, + "filename": file.filename, + "message": "No content received from provider", + "error": "Empty file download" + } + + # Combine chunks + new_content = b''.join(content_chunks) + + # Save the new content to storage + user = type('User', (), { + 'id': user_id, + 'email': f'user_{user_id}@example.com', + 'name': f'User {user_id}' + })() + + tags = { + "OpenWebUI-User-Email": user.email, + "OpenWebUI-User-Id": user.id, + "OpenWebUI-User-Name": user.name, + "OpenWebUI-File-Id": file_id, + } + + # Upload the new content (this will overwrite the existing file) + filename = f"{file_id}_{file.filename}" + try: + _, file_path = Storage.upload_file( + file=io.BytesIO(new_content), + filename=filename, + tags=tags + ) + except Exception as e: + log.error(f"Failed to upload synced file to storage: {e}") + return { + "success": False, + "file_id": file_id, + "filename": file.filename, + "message": "Failed to save synced content", + "error": str(e) + } + + # Update file metadata with new provider info + updated_provider_info = { + "provider": provider_name, + "provider_file_id": provider_file_id, + "provider_modified_time": provider_file.get("modified_time") if provider_file else None, + "provider_sync_enabled": provider_info.get("provider_sync_enabled", False), + "provider_metadata": provider_file.get("metadata") if provider_file else provider_info.get("provider_metadata") + } + + # Update the file's provider info + updated_file = Files.update_file_provider_info_by_id( + id=file_id, + **updated_provider_info + ) + + if not updated_file: + log.warning(f"Failed to update file provider info for file {file_id}") + + # Re-process the file to update embeddings + content_changed = False + try: + # Process the file with new content + process_file( + request, + ProcessFileForm(file_id=file_id), + user=user + ) + content_changed = True + except Exception as e: + log.error(f"Error processing synced file: {e}") + # Don't fail the sync if processing fails + content_changed = False + + return { + "success": True, + "file_id": file_id, + "filename": file.filename, + "message": "File synced successfully", + "updated": True, + "content_changed": content_changed, + "provider_modified_time": updated_provider_info.get("provider_modified_time") + } + + except Exception as e: + log.exception(f"Error syncing file {file_id}: {e}") + return { + "success": False, + "file_id": file_id, + "filename": file.filename, + "message": "Failed to sync file", + "error": str(e) + } + + async def sync_files_by_provider_batch( + self, + provider_name: str, + force_sync: bool, + user_id: str, + user_role: str, + request: Request, + permission_check_fn=None + ) -> Dict[str, Any]: + """ + Sync all files from a specific provider in batch. + + Args: + provider_name: Name of the provider to sync files from + force_sync: Whether to force sync even if sync is disabled + user_id: ID of the user performing the sync + user_role: Role of the user (admin, user, etc.) + request: FastAPI request object + permission_check_fn: Optional function to check file permissions + + Returns: + Dict containing batch sync results + """ + log.info(f"Starting batch sync for provider {provider_name}") + + try: + # Validate provider exists + provider = content_source_factory.get_provider(provider_name) + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Invalid provider: {str(e)}" + ) + + # Get all files from this provider + if user_role == "admin": + files = Files.get_files_by_provider(provider_name) + else: + files = Files.get_files_by_provider(provider_name, user_id=user_id) + + # Track results + results = [] + errors = [] + synced_count = 0 + failed_count = 0 + skipped_count = 0 + + for file in files: + # Check if file has sync enabled (unless force sync) + if not force_sync and not file.provider_sync_enabled: + skipped_count += 1 + results.append({ + "success": True, + "file_id": file.id, + "filename": file.filename, + "message": "Sync not enabled for this file", + "updated": False + }) + continue + + # Check permissions for each file if permission function provided + if permission_check_fn and not permission_check_fn(file.id, user_id, user_role): + skipped_count += 1 + results.append({ + "success": False, + "file_id": file.id, + "filename": file.filename, + "message": "No permission to sync this file", + "error": "Unauthorized" + }) + continue + + # Sync the file + try: + sync_result = await self.sync_file_from_provider(file.id, request, user_id) + results.append(sync_result) + + if sync_result["success"]: + if sync_result.get("updated", False): + synced_count += 1 + else: + skipped_count += 1 + else: + failed_count += 1 + if sync_result.get("error"): + errors.append(f"File {file.filename}: {sync_result['error']}") + except Exception as e: + failed_count += 1 + error_msg = f"Failed to sync file {file.filename}: {str(e)}" + errors.append(error_msg) + results.append({ + "success": False, + "file_id": file.id, + "filename": file.filename, + "message": "Sync failed", + "error": str(e) + }) + + return { + "success": failed_count == 0, + "provider": provider_name, + "total_files": len(files), + "synced_files": synced_count, + "failed_files": failed_count, + "skipped_files": skipped_count, + "results": results, + "errors": errors + } + + +# Global service instance +content_syncer = ContentSourceSyncer() \ No newline at end of file diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 642d6c3568..7818dfcda1 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -89,6 +89,7 @@ from open_webui.routers import ( memories, models, knowledge, + content_sources, prompts, evaluations, skills, @@ -564,7 +565,7 @@ from open_webui.tasks import ( list_tasks, ) # Import from tasks.py -from open_webui.services.google_drive_scheduler import google_drive_scheduler +from open_webui.content_sources.scheduler import scheduler as content_source_scheduler from open_webui.utils.redis import get_sentinels_from_env @@ -707,30 +708,26 @@ async def lifespan(app: FastAPI): # Mark application as ready to accept traffic from a startup perspective. app.state.startup_complete = True - # Start Google Drive sync scheduler if enabled - if app.state.config.ENABLE_GOOGLE_DRIVE_FOLDER_SYNC: - try: - await google_drive_scheduler.start() - log.info("Google Drive sync scheduler started successfully") - except Exception as e: - log.error(f"Failed to start Google Drive sync scheduler: {e}") - # Non-critical failure - continue application startup + # Start content source sync scheduler if any providers are configured + # The scheduler will automatically check which providers need syncing + try: + await content_source_scheduler.start() + log.info("Content source sync scheduler started successfully") + except Exception as e: + log.error(f"Failed to start content source sync scheduler: {e}") + # Non-critical failure - continue application startup yield if hasattr(app.state, 'redis_task_command_listener'): app.state.redis_task_command_listener.cancel() - # Stop Google Drive sync scheduler - if ( - hasattr(app.state, "config") - and app.state.config.ENABLE_GOOGLE_DRIVE_FOLDER_SYNC - ): - try: - await google_drive_scheduler.stop() - log.info("Google Drive sync scheduler stopped") - except Exception as e: - log.error(f"Error stopping Google Drive sync scheduler: {e}") + # Stop content source sync scheduler + try: + await content_source_scheduler.stop() + log.info("Content source sync scheduler stopped") + except Exception as e: + log.error(f"Error stopping content source sync scheduler: {e}") app = FastAPI( @@ -1537,6 +1534,7 @@ app.include_router(notes.router, prefix='/api/v1/notes', tags=['notes']) app.include_router(models.router, prefix='/api/v1/models', tags=['models']) app.include_router(knowledge.router, prefix='/api/v1/knowledge', tags=['knowledge']) +app.include_router(content_sources.router, prefix='/api/v1/content-sources', tags=['content_sources']) app.include_router(prompts.router, prefix='/api/v1/prompts', tags=['prompts']) app.include_router(tools.router, prefix='/api/v1/tools', tags=['tools']) app.include_router(skills.router, prefix='/api/v1/skills', tags=['skills']) diff --git a/backend/open_webui/models/files.py b/backend/open_webui/models/files.py index 9a5b8fa400..45e36b94bf 100644 --- a/backend/open_webui/models/files.py +++ b/backend/open_webui/models/files.py @@ -48,6 +48,28 @@ class FileModel(BaseModel): created_at: Optional[int] # timestamp in epoch updated_at: Optional[int] # timestamp in epoch + + @property + def provider_info(self) -> dict: + """Get provider information from the data field.""" + if self.data and isinstance(self.data, dict): + return self.data.get("provider_info", {}) + return {} + + @property + def provider(self) -> Optional[str]: + """Get the provider name.""" + return self.provider_info.get("provider") + + @property + def provider_file_id(self) -> Optional[str]: + """Get the provider's file ID.""" + return self.provider_info.get("provider_file_id") + + @property + def provider_sync_enabled(self) -> bool: + """Check if provider sync is enabled.""" + return self.provider_info.get("provider_sync_enabled", False) #################### @@ -177,6 +199,14 @@ class FilesTable: except Exception: return None + def get_file_provider_info(self, id: str, db: Optional[Session] = None) -> Optional[dict]: + """Get provider information for a file.""" + with get_db_context(db) as db: + file = db.query(File).filter_by(id=id).first() + if file and file.data and isinstance(file.data, dict): + return file.data.get("provider_info", {}) + return None + def get_file_metadata_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileMetadataResponse]: with get_db_context(db) as db: try: @@ -315,6 +345,24 @@ class FilesTable: self, id: str, form_data: FileUpdateForm, db: Optional[Session] = None ) -> Optional[FileModel]: with get_db_context(db) as db: + + def get_files_by_provider( + self, provider: str, user_id: Optional[str] = None, db: Optional[Session] = None + ) -> list[FileModel]: + """Get all files from a specific provider, optionally filtered by user.""" + with get_db_context(db) as db: + query = db.query(File) + if user_id: + query = query.filter_by(user_id=user_id) + files = [] + for file in query.all(): + if file.data and isinstance(file.data, dict): + provider_info = file.data.get("provider_info", {}) + if provider_info.get("provider") == provider: + files.append(FileModel.model_validate(file)) + return files + + try: file = db.query(File).filter_by(id=id).first() @@ -372,6 +420,45 @@ class FilesTable: def delete_file_by_id(self, id: str, db: Optional[Session] = None) -> bool: with get_db_context(db) as db: + + def update_file_provider_info_by_id( + self, + id: str, + provider: Optional[str] = None, + provider_file_id: Optional[str] = None, + provider_modified_time: Optional[str] = None, + provider_sync_enabled: Optional[bool] = None, + provider_metadata: Optional[dict] = None, + db: Optional[Session] = None + ) -> Optional[FileModel]: + """Update provider information for a file.""" + with get_db_context(db) as db: + try: + file = db.query(File).filter_by(id=id).first() + if not file: + return None + if not file.data: + file.data = {} + if "provider_info" not in file.data: + file.data["provider_info"] = {} + updates = { + "provider": provider, + "provider_file_id": provider_file_id, + "provider_modified_time": provider_modified_time, + "provider_sync_enabled": provider_sync_enabled, + "provider_metadata": provider_metadata, + } + for key, value in updates.items(): + if value is not None: + file.data["provider_info"][key] = value + file.updated_at = int(time.time()) + db.commit() + return FileModel.model_validate(file) + except Exception as e: + log.exception(f'Error updating file provider info: {e}') + return None + + try: db.query(File).filter_by(id=id).delete() db.commit() diff --git a/backend/open_webui/routers/content_sources.py b/backend/open_webui/routers/content_sources.py new file mode 100644 index 0000000000..0ed11c6bf0 --- /dev/null +++ b/backend/open_webui/routers/content_sources.py @@ -0,0 +1,279 @@ +""" +Content Sources Router + +Generic router for managing content source providers. +Provides provider-agnostic endpoints for listing providers and getting provider information. +""" + +from typing import List, Dict, Any, Optional +from fastapi import APIRouter, Depends, HTTPException, status +from pydantic import BaseModel +import logging + +from open_webui.models.users import UserModel +from open_webui.utils.auth import get_verified_user +from open_webui.content_sources import content_source_factory, content_source_registry +from open_webui.constants import ERROR_MESSAGES +from open_webui.env import SRC_LOG_LEVELS + +log = logging.getLogger(__name__) +log.setLevel(SRC_LOG_LEVELS["MODELS"]) + +router = APIRouter() + + +############################ +# Response Models +############################ + + +class ProviderInfo(BaseModel): + """Information about a content source provider.""" + name: str + display_name: str + description: str + configured: bool + metadata: Optional[Dict[str, Any]] = None + + +class ProviderServiceInfo(BaseModel): + """Service information for a specific provider.""" + provider: str + configured: bool + metadata: Dict[str, Any] + + +class ProviderListResponse(BaseModel): + """Response containing list of available providers.""" + providers: List[ProviderInfo] + + +############################ +# Helper Functions +############################ + + +def get_provider_display_info(provider_name: str) -> Dict[str, str]: + """Get display information for providers.""" + # This could be expanded with more providers + provider_info = { + "google_drive": { + "display_name": "Google Drive", + "description": "Sync files from Google Drive folders" + }, + "onedrive": { + "display_name": "Microsoft OneDrive", + "description": "Sync files from OneDrive folders" + }, + "dropbox": { + "display_name": "Dropbox", + "description": "Sync files from Dropbox folders" + }, + "sharepoint": { + "display_name": "SharePoint", + "description": "Sync files from SharePoint document libraries" + } + } + + return provider_info.get(provider_name, { + "display_name": provider_name.replace("_", " ").title(), + "description": f"Content source provider: {provider_name}" + }) + + +############################ +# Endpoints +############################ + + +@router.get("/", response_model=ProviderListResponse) +async def list_content_source_providers( + user: UserModel = Depends(get_verified_user) +) -> ProviderListResponse: + """ + List all available content source providers. + + Returns information about each provider including whether it's configured. + """ + providers = [] + + # Get available provider types from factory + available_providers = content_source_factory.get_available_providers() + + for provider_name in available_providers: + try: + # Try to get the provider instance from registry + provider = content_source_registry.get_provider(provider_name) + + if provider: + # Provider is registered, check if it's configured + try: + service_info = await provider.get_service_info() + configured = service_info.get('configured', False) + except Exception as e: + log.warning(f"Failed to get service info for {provider_name}: {e}") + configured = False + else: + # Provider is available but not registered + configured = False + + # Get display information + display_info = get_provider_display_info(provider_name) + + providers.append(ProviderInfo( + name=provider_name, + display_name=display_info["display_name"], + description=display_info["description"], + configured=configured + )) + + except Exception as e: + log.error(f"Error processing provider {provider_name}: {e}") + continue + + return ProviderListResponse(providers=providers) + + +@router.get("/{provider}/info", response_model=ProviderServiceInfo) +async def get_provider_info( + provider: str, + user: UserModel = Depends(get_verified_user) +) -> ProviderServiceInfo: + """ + Get provider-specific information. + + This endpoint returns provider-specific metadata such as service account emails, + API endpoints, quotas, or any other provider-specific information. + + Args: + provider: The provider name (e.g., 'google_drive', 'onedrive') + + Returns: + Provider service information including configuration status and metadata + """ + # Check if provider exists in factory + available_providers = content_source_factory.get_available_providers() + if provider not in available_providers: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Unknown content source provider: {provider}. Available providers: {', '.join(available_providers.keys())}" + ) + + # Try to get provider from registry + provider_instance = content_source_registry.get_provider(provider) + + if not provider_instance: + # Provider is available but not initialized/registered + try: + # Try to initialize the provider + provider_instance = content_source_factory.get_provider(provider) + content_source_registry.register_provider(provider, provider_instance) + except Exception as e: + log.error(f"Failed to initialize provider {provider}: {e}") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"Failed to initialize provider {provider}: {str(e)}" + ) + + try: + # Get service information from the provider + service_info = await provider_instance.get_service_info() + + # Extract metadata, removing the 'configured' field to avoid duplication + metadata = {k: v for k, v in service_info.items() if k not in ['configured', 'provider']} + + return ProviderServiceInfo( + provider=provider, + configured=service_info.get('configured', False), + metadata=metadata + ) + + except Exception as e: + log.error(f"Error getting service info for provider {provider}: {e}") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"Failed to get service information: {str(e)}" + ) + + +@router.get("/{provider}/capabilities", response_model=Dict[str, Any]) +async def get_provider_capabilities( + provider: str, + user: UserModel = Depends(get_verified_user) +) -> Dict[str, Any]: + """ + Get capabilities and features supported by a provider. + + This endpoint returns information about what operations the provider supports, + such as folder sync, file filtering, nested folder support, etc. + + Args: + provider: The provider name + + Returns: + Dictionary of provider capabilities + """ + # Check if provider exists + available_providers = content_source_factory.get_available_providers() + if provider not in available_providers: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Unknown content source provider: {provider}" + ) + + # Define capabilities for known providers + # This could be moved to each provider class as a method + capabilities = { + "google_drive": { + "supports_folder_sync": True, + "supports_nested_folders": True, + "supports_file_filtering": True, + "supports_incremental_sync": True, + "supports_oauth": False, # Currently using service account + "supports_webhooks": False, + "file_size_limit": "5TB", + "supported_file_types": ["documents", "spreadsheets", "presentations", "pdfs", "text", "images"], + "export_formats": { + "google-docs": ["docx", "pdf", "txt", "html"], + "google-sheets": ["xlsx", "csv", "pdf"], + "google-slides": ["pptx", "pdf"] + } + }, + # Add more providers as they are implemented + # This is an example, MUST be changed in the real implementation + "onedrive": { + "supports_folder_sync": True, + "supports_nested_folders": True, + "supports_file_filtering": True, + "supports_incremental_sync": True, + "supports_oauth": True, + "supports_webhooks": True, + "file_size_limit": "250GB", + "supported_file_types": ["documents", "spreadsheets", "presentations", "pdfs", "text", "images"] + }, + "dropbox": { + "supports_folder_sync": True, + "supports_nested_folders": True, + "supports_file_filtering": True, + "supports_incremental_sync": True, + "supports_oauth": True, + "supports_webhooks": True, + "file_size_limit": "50GB", + "supported_file_types": ["documents", "spreadsheets", "presentations", "pdfs", "text", "images"] + } + } + + provider_capabilities = capabilities.get(provider, { + "supports_folder_sync": False, + "supports_nested_folders": False, + "supports_file_filtering": False, + "supports_incremental_sync": False, + "supports_oauth": False, + "supports_webhooks": False, + "message": "Capabilities not defined for this provider" + }) + + return { + "provider": provider, + "capabilities": provider_capabilities + } \ No newline at end of file diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index a19490b5d0..f971630754 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -1,4 +1,4 @@ -from typing import List, Optional +from typing import List, Optional, Dict, Any from pydantic import BaseModel from fastapi import APIRouter, Depends, HTTPException, status, Request, Query from fastapi.responses import StreamingResponse @@ -8,6 +8,7 @@ import time import io import zipfile from urllib.parse import quote +import base64 from sqlalchemy.orm import Session from open_webui.internal.db import get_session @@ -19,7 +20,7 @@ from open_webui.models.knowledge import ( KnowledgeResponse, KnowledgeUserResponse, ) -from open_webui.models.files import Files, FileModel, FileMetadataResponse +from open_webui.models.files import Files, FileModel, FileMetadataResponse, FileForm from open_webui.models.users import UserModel from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT from open_webui.routers.retrieval import ( @@ -28,13 +29,14 @@ from open_webui.routers.retrieval import ( process_files_batch, BatchProcessFilesForm, ) -from open_webui.services.google_drive import google_drive_service from open_webui.storage.provider import Storage from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_verified_user, get_admin_user from open_webui.utils.access_control import has_permission, filter_allowed_access_grants from open_webui.models.access_grants import AccessGrants +from open_webui.utils.misc import calculate_sha256_string +from open_webui.content_sources import content_source_registry, content_source_factory from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL @@ -44,6 +46,41 @@ log = logging.getLogger(__name__) router = APIRouter() + +############################ +# Helper Functions +############################ + +def check_duplicate_in_vector_db(collection_name: str, file_hash: str) -> bool: + """Check if a file with the given hash already exists in the vector DB.""" + if not file_hash or not collection_name: + return False + + try: + result = VECTOR_DB_CLIENT.query( + collection_name=collection_name, + filter={"hash": file_hash}, + ) + + if result is not None and result.ids and result.ids[0]: + return True + except Exception as e: + log.debug(f"Error checking duplicate in vector DB: {e}") + + return False + + +############################ +# Response Models +############################ + + +class KnowledgeFilesResponse(KnowledgeResponse): + """Knowledge base response with files list.""" + files: List[FileMetadataResponse] = [] + warnings: Optional[Dict[str, Any]] = None + sync_results: Optional[Dict[str, Any]] = None + ############################ # getKnowledgeBases ############################ @@ -270,6 +307,12 @@ async def create_new_knowledge( 'sharing.public_knowledge', ) +# Emit before hook + await content_source_registry.emit_hook('before_knowledge_create', { + 'form_data': form_data.model_dump() if hasattr(form_data, 'model_dump') else form_data.__dict__, + 'user_id': user.id + }) + knowledge = Knowledges.insert_new_knowledge(user.id, form_data) if knowledge: @@ -280,6 +323,13 @@ async def create_new_knowledge( knowledge.name, knowledge.description, ) + +# Emit after hook + await content_source_registry.emit_hook('after_knowledge_create', { + 'knowledge_base_id': knowledge.id, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__, + 'user_id': user.id + }) return knowledge else: raise HTTPException( @@ -385,6 +435,7 @@ async def reindex_knowledge_base_metadata_embeddings( class KnowledgeFilesResponse(KnowledgeResponse): files: Optional[list[FileMetadataResponse]] = None write_access: Optional[bool] = False + sync_results: Optional[Dict[str, Any]] = None @router.get('/{id}', response_model=Optional[KnowledgeFilesResponse]) @@ -660,6 +711,15 @@ def add_file_to_knowledge_by_id( detail=ERROR_MESSAGES.FILE_NOT_PROCESSED, ) + # Emit before hook + await content_source_registry.emit_hook('before_file_add', { + 'knowledge_base_id': id, + 'file_id': form_data.file_id, + 'user_id': user.id, + 'file': file.model_dump() if hasattr(file, 'model_dump') else file.__dict__, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__ + }) + # Add content to the vector database try: process_file( @@ -888,6 +948,13 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S log.info(f'Deleting knowledge base: {id} (name: {knowledge.name})') + # Emit before hook + await content_source_registry.emit_hook('before_knowledge_delete', { + 'knowledge_base_id': id, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__, + 'user_id': user.id + }) + # Get all models models = Models.get_all_models(db=db) log.info(f'Found {len(models)} models to check for knowledge base {id}') @@ -959,6 +1026,13 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Se detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) + # Emit before hook + await content_source_registry.emit_hook('before_knowledge_reset', { + 'knowledge_base_id': id, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__, + 'user_id': user.id + }) + try: VECTOR_DB_CLIENT.delete_collection(collection_name=id) except Exception as e: @@ -1022,6 +1096,14 @@ async def add_files_to_knowledge_batch( detail=f'File {missing_ids[0]} not found', ) + # Emit before hook + await content_source_registry.emit_hook('before_files_batch_add', { + 'knowledge_base_id': id, + 'file_ids': [form.file_id for form in form_data], + 'user_id': user.id, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__ + }) + # Process files try: result = await process_files_batch( @@ -1039,6 +1121,14 @@ async def add_files_to_knowledge_batch( for file_id in successful_file_ids: Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db) + # Emit after hook + await content_source_registry.emit_hook('after_files_batch_add', { + 'knowledge_base_id': id, + 'file_ids': successful_file_ids, + 'user_id': user.id, + 'knowledge_base': knowledge.model_dump() if hasattr(knowledge, 'model_dump') else knowledge.__dict__ + }) + # If there were any errors, include them in the response if result.errors: error_details = [f'{err.file_id}: {err.error}' for err in result.errors] @@ -1109,59 +1199,42 @@ async def export_knowledge_by_id(id: str, user=Depends(get_admin_user), db: Sess ############################ -# Google Drive Sync +# Sync Content from Provider ############################ -class GoogleDriveSyncForm(BaseModel): - folder_id: str - include_nested: bool = True - sync_interval_days: float = 1.0 # Float to support fractional days +# Import the sync service +from open_webui.content_sources.syncer import content_syncer -class GoogleDriveServiceAccountResponse(BaseModel): - email: str +class ContentSourceSyncForm(BaseModel): + """Form for syncing content from a content source provider.""" + provider: str + source_id: str + options: Optional[Dict[str, Any]] = None -@router.get( - "/google-drive/service-account-email", - response_model=GoogleDriveServiceAccountResponse, -) -async def get_google_drive_service_account_email( - user: UserModel = Depends(get_verified_user), -) -> GoogleDriveServiceAccountResponse: - """Get the Google Drive service account email for sharing folders.""" - if not google_drive_service.is_configured(): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Google Drive service account not configured", - ) - - email = google_drive_service.get_service_account_email() - if not email: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Failed to get service account email", - ) - - return {"email": email} - - -@router.post("/{id}/google-drive/sync", response_model=Optional[KnowledgeFilesResponse]) -async def sync_google_drive_folder( - id: str, - form_data: GoogleDriveSyncForm, +@router.post("/{id}/sync", response_model=Optional[KnowledgeFilesResponse]) +async def sync_content_from_provider( request: Request, - user: UserModel = Depends(get_verified_user), -) -> Optional[KnowledgeFilesResponse]: - """Sync a Google Drive folder with the knowledge base.""" + id: str, + form_data: ContentSourceSyncForm, + user=Depends(get_verified_user), +): + """ + Sync content from a content source provider to a knowledge base. + + This endpoint now acts as a thin orchestrator, delegating the heavy + sync logic to the ContentSyncService. + """ + # Validate knowledge base and permissions knowledge = Knowledges.get_knowledge_by_id(id=id) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.NOT_FOUND, ) - + if ( knowledge.user_id != user.id and not has_access(user.id, "write", knowledge.access_control) @@ -1171,240 +1244,79 @@ async def sync_google_drive_folder( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - - if not google_drive_service.is_configured(): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Google Drive service account not configured", - ) - + try: - # Get current files in knowledge base - current_file_ids = knowledge.data.get("file_ids", []) if knowledge.data else [] - current_files = Files.get_files_by_ids(current_file_ids) - - # Create a map of Google Drive file IDs to local file IDs - gdrive_file_map = {} - for file in current_files: - if file.meta and file.meta.get("google_drive_id"): - gdrive_file_map[file.meta["google_drive_id"]] = file.id - - # Get files from Google Drive folder - gdrive_files = google_drive_service.list_folder_files( - form_data.folder_id, form_data.include_nested + # Use the sync service to perform the sync + log.info(f"Starting sync for knowledge base {id} with provider {form_data.provider}") + + sync_result = await content_syncer.sync_provider_files( + provider_name=form_data.provider, + source_id=form_data.source_id, + options=form_data.options or {}, + request=request, + user_id=user.id, + knowledge_base_id=id ) - - log.info( - f"Google Drive sync: Found {len(gdrive_files)} files in folder {form_data.folder_id}" - ) - for gdrive_file in gdrive_files: - log.info( - f"Google Drive file: {gdrive_file.get('name', 'Unknown')} (ID: {gdrive_file.get('id', 'Unknown')}, Type: {gdrive_file.get('mimeType', 'Unknown')})" - ) - - # Track files to keep and files to add - files_to_keep = set() - files_to_add = [] - - for gdrive_file in gdrive_files: - gdrive_id = gdrive_file["id"] - - if gdrive_id in gdrive_file_map: - # File exists, check if it needs updating - local_file_id = gdrive_file_map[gdrive_id] - local_file = Files.get_file_by_id(local_file_id) - - if local_file and local_file.meta: - local_modified = local_file.meta.get("google_drive_modified") - gdrive_modified = gdrive_file["modifiedTime"] - - if local_modified != gdrive_modified: - # File was modified, re-download and update - files_to_add.append(gdrive_file) - # Remove old file - Files.delete_file_by_id(local_file_id) - else: - # File is up to date, keep it - files_to_keep.add(local_file_id) - else: - # Local file metadata is missing, re-download - files_to_add.append(gdrive_file) - else: - # New file, add it - files_to_add.append(gdrive_file) - - # Remove files that are no longer in Google Drive - files_to_remove = set(current_file_ids) - files_to_keep - for file_id in files_to_remove: - Files.delete_file_by_id(file_id) - - # Download and add new/updated files - new_file_ids = list(files_to_keep) - - log.info( - f"Google Drive sync: Processing {len(files_to_add)} files to add/update" - ) - - for gdrive_file in files_to_add: - try: - log.info( - f"Google Drive sync: Processing file '{gdrive_file.get('name', 'Unknown')}' (ID: {gdrive_file.get('id', 'Unknown')})" - ) - - # Download file from Google Drive - file_content, filename = google_drive_service.download_file( - gdrive_file["id"], gdrive_file - ) - - log.info( - f"Google Drive sync: Downloaded file '{filename}', size: {len(file_content)} bytes" - ) - - # Debug: Check if content is actually text - try: - content_preview = ( - file_content.decode("utf-8")[:200] - if isinstance(file_content, bytes) - else str(file_content)[:200] - ) - log.info( - f"Google Drive sync: File content preview: {content_preview}..." - ) - except Exception as e: - log.error(f"Google Drive sync: Could not decode file content: {e}") - - # Create file object - file_obj = io.BytesIO(file_content) - file_obj.name = filename - - # Upload to storage - import uuid - - file_id = str(uuid.uuid4()) - tags = { - "OpenWebUI-User-Email": user.email, - "OpenWebUI-User-Id": user.id, - "OpenWebUI-User-Name": user.name, - "OpenWebUI-File-Id": file_id, - } - - contents, file_path = Storage.upload_file( - file_obj, f"{file_id}_{filename}", tags - ) - log.info( - f"Google Drive sync: Uploaded to storage, contents size: {len(contents)} bytes" - ) - - # Create file record with content in data field - from open_webui.models.files import FileForm - - # Convert bytes to string for text content - text_content = ( - file_content.decode("utf-8") - if isinstance(file_content, bytes) - else str(file_content) - ) - - file_item = Files.insert_new_file( - user.id, - FileForm( - id=file_id, - filename=filename, - path=file_path, - data={ - "content": text_content, # Store the actual text content - }, - meta={ - "name": filename, - "content_type": "text/plain", # Changed from application/octet-stream - "size": len(contents), - "google_drive_id": gdrive_file["id"], - "google_drive_modified": gdrive_file["modifiedTime"], - "google_drive_path": gdrive_file["path"], - "collection_name": id, - }, - ), - ) - - if file_item: - new_file_ids.append(file_id) - log.info( - f"Google Drive sync: Created file record for '{filename}' (File ID: {file_id})" - ) - - # Process file for vector storage - try: - from open_webui.routers.retrieval import process_file - - process_file( - request, - ProcessFileForm( - file_id=file_id, - collection_name=id, - ), - user, - ) - log.info( - f"Google Drive sync: Successfully processed file '{filename}' for vector storage" - ) - except Exception as e: - log.error( - f"Failed to process file {file_id} ('{filename}') for vector storage: {e}" - ) - else: - log.error( - f"Google Drive sync: Failed to create file record for '{filename}'" - ) - - except Exception as e: - log.error(f"Failed to sync file {gdrive_file['name']}: {e}") - continue - - # Update knowledge base data - sync_data = { - "google_drive_folder_id": form_data.folder_id, - "google_drive_include_nested": form_data.include_nested, - "google_drive_sync_interval_days": form_data.sync_interval_days, - "google_drive_last_sync": int(time.time()), + + # Update knowledge base with sync results + data = knowledge.data or {} + + # Update file IDs with successfully synced files and remove deleted ones + existing_file_ids = set(data.get("file_ids", [])) + new_file_ids = set(sync_result.successful_files) + removed_file_ids = set(sync_result.removed_files) if hasattr(sync_result, 'removed_files') else set() + + # Add new files and remove deleted ones + updated_file_ids = list((existing_file_ids | new_file_ids) - removed_file_ids) + data["file_ids"] = updated_file_ids + + # Update sync metadata + data.setdefault("sync_metadata", {})[form_data.provider] = { + "source_id": form_data.source_id, + "last_sync": time.time(), + "options": form_data.options, + "results": { + "status": sync_result.status.value, + "added": len(sync_result.added_files), + "updated": len(sync_result.updated_files), + "removed": len(sync_result.removed_files) if hasattr(sync_result, 'removed_files') else 0, + "failed": len(sync_result.failed_files), + "duplicates": len(sync_result.duplicate_files), + "changes": sync_result.changes + } } - - updated_data = knowledge.data.copy() if knowledge.data else {} - updated_data.update(sync_data) - updated_data["file_ids"] = new_file_ids - - knowledge = Knowledges.update_knowledge_data_by_id(id=id, data=updated_data) - - log.info( - f"Google Drive sync completed: {len(new_file_ids)} total files in knowledge base" + + knowledge = Knowledges.update_knowledge_data_by_id(id=id, data=data) + + # Prepare response + files = Files.get_file_metadatas_by_ids(updated_file_ids) + + # Build sync results for response + sync_results_dict = None + if sync_result.changes or sync_result.errors: + sync_results_dict = { + "status": sync_result.status.value, + "added_files": sync_result.added_files, + "updated_files": sync_result.updated_files, + "duplicate_files": sync_result.duplicate_files, + "errors": sync_result.errors, + "warnings": sync_result.warnings, + "changes": sync_result.changes, + "total_processed": sync_result.total_processed + } + + return KnowledgeFilesResponse( + **knowledge.model_dump(), + files=files, + sync_results=sync_results_dict ) - log.info(f"Google Drive sync: Knowledge object type: {type(knowledge)}") - log.info(f"Google Drive sync: Knowledge object: {knowledge}") - - if knowledge: - files = Files.get_file_metadatas_by_ids(new_file_ids) - log.info(f"Google Drive sync: Returning {len(files)} file metadata records") - - # Handle both dict and model cases - if hasattr(knowledge, "model_dump"): - knowledge_dict = knowledge.model_dump() - else: - knowledge_dict = ( - knowledge if isinstance(knowledge, dict) else knowledge.__dict__ - ) - - return KnowledgeFilesResponse( - **knowledge_dict, - files=files, - ) - else: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Failed to update knowledge base", - ) - + + except HTTPException: + raise except Exception as e: - log.error(f"Google Drive sync error: {e}") + log.error(f"Sync failed for {form_data.provider}: {e}", exc_info=True) raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Failed to sync Google Drive folder: {str(e)}", + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"Sync failed: {str(e)}" ) + diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 054c02cae9..2662472706 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -97,7 +97,7 @@ from open_webui.utils.misc import ( ) from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_permission -from open_webui.services.google_drive import google_drive_service +from open_webui.utils.content_sources import refresh_provider_configuration from open_webui.config import ( ENV, @@ -1061,9 +1061,9 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend else request.app.state.config.ENABLE_ONEDRIVE_INTEGRATION ) - # Refresh Google Drive service if service account JSON was updated + # Refresh Google Drive provider if service account JSON was updated if form_data.GOOGLE_DRIVE_SERVICE_ACCOUNT_JSON is not None: - google_drive_service.refresh_configuration() + refresh_provider_configuration('google_drive') if form_data.web is not None: # Web search settings diff --git a/backend/open_webui/services/__init__.py b/backend/open_webui/services/__init__.py deleted file mode 100644 index a15012dfa8..0000000000 --- a/backend/open_webui/services/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Services module for Open WebUI diff --git a/backend/open_webui/services/google_drive_scheduler.py b/backend/open_webui/services/google_drive_scheduler.py deleted file mode 100644 index 573d82ce72..0000000000 --- a/backend/open_webui/services/google_drive_scheduler.py +++ /dev/null @@ -1,238 +0,0 @@ -import asyncio -import logging -import time -from typing import Dict, List, Any, Optional -from datetime import datetime, timedelta -import uuid -import io - -from open_webui.env import ENV -from open_webui.models.knowledge import Knowledges, KnowledgeModel -from open_webui.services.google_drive import google_drive_service, GoogleDriveFile -from open_webui.models.files import Files, FileModel, FileForm -from open_webui.routers.retrieval import process_file, ProcessFileForm -from open_webui.storage.provider import Storage - -log = logging.getLogger(__name__) - - -class GoogleDriveSyncScheduler: - """Background scheduler for automatic Google Drive folder sync.""" - - def __init__(self) -> None: - self.running: bool = False - self.sync_tasks: Dict[str, Any] = {} - # In dev environment, check every minute for faster testing - # In production, check every hour - self.check_interval: int = 60 if ENV == "dev" else 3600 - - async def start(self) -> None: - """Start the background sync scheduler.""" - if self.running: - return - - self.running = True - log.info(f"Google Drive sync scheduler started (ENV: {ENV}, check_interval: {self.check_interval}s)") - - # Start the background task - asyncio.create_task(self._sync_loop()) - - async def stop(self) -> None: - """Stop the background sync scheduler.""" - self.running = False - log.info("Google Drive sync scheduler stopped") - - async def _sync_loop(self) -> None: - """Main sync loop that runs in the background.""" - while self.running: - try: - await self._check_and_sync_knowledge_bases() - await asyncio.sleep(self.check_interval) - except Exception as e: - log.error(f"Error in sync loop: {e}") - await asyncio.sleep(self.check_interval) - - async def _check_and_sync_knowledge_bases(self) -> None: - """Check all knowledge bases for Google Drive sync requirements.""" - if not google_drive_service.is_configured(): - return - - try: - # Get all knowledge bases - knowledge_bases = Knowledges.get_knowledge_bases() - - for kb in knowledge_bases: - if not kb.data: - continue - - # Check if this knowledge base has Google Drive sync configured - folder_id = kb.data.get("google_drive_folder_id") - sync_interval_days = kb.data.get("google_drive_sync_interval_days", 1) - last_sync = kb.data.get("google_drive_last_sync", 0) - - if not folder_id: - continue - - # Check if sync is due - current_time = int(time.time()) - sync_interval_seconds = sync_interval_days * 24 * 3600 - - if current_time - last_sync >= sync_interval_seconds: - log.info( - f"Auto-syncing knowledge base {kb.id} with Google Drive folder {folder_id}" - ) - await self._sync_knowledge_base(kb) - - except Exception as e: - log.error(f"Error checking knowledge bases for sync: {e}") - - async def _sync_knowledge_base(self, knowledge_base: KnowledgeModel) -> None: - """Sync a specific knowledge base with its Google Drive folder.""" - try: - if not knowledge_base.data: - return - - # Type guard for mypy - assert knowledge_base.data is not None - - folder_id = knowledge_base.data.get("google_drive_folder_id") - include_nested = knowledge_base.data.get( - "google_drive_include_nested", True - ) - - if not folder_id: - return - - # Get current files in knowledge base - current_file_ids = knowledge_base.data.get("file_ids", []) - current_files = Files.get_files_by_ids(current_file_ids) - - # Create a map of Google Drive file IDs to local file IDs - gdrive_file_map = {} - for file in current_files: - if file.meta and file.meta.get("google_drive_id"): - gdrive_file_map[file.meta["google_drive_id"]] = file.id - - # Get files from Google Drive folder - gdrive_files = google_drive_service.list_folder_files( - folder_id, include_nested - ) - - # Track files to keep and files to add - files_to_keep = set() - files_to_add = [] - - for gdrive_file in gdrive_files: - gdrive_id = gdrive_file["id"] - - if gdrive_id in gdrive_file_map: - # File exists, check if it needs updating - local_file_id = gdrive_file_map[gdrive_id] - local_file = Files.get_file_by_id(local_file_id) - - if local_file and local_file.meta: - local_modified = local_file.meta.get("google_drive_modified") - gdrive_modified = gdrive_file["modifiedTime"] - - if local_modified != gdrive_modified: - # File was modified, re-download and update - files_to_add.append(gdrive_file) - # Remove old file - Files.delete_file_by_id(local_file_id) - else: - # File is up to date, keep it - files_to_keep.add(local_file_id) - else: - # Local file metadata is missing, re-download - files_to_add.append(gdrive_file) - else: - # New file, add it - files_to_add.append(gdrive_file) - - # Remove files that are no longer in Google Drive - files_to_remove = set(current_file_ids) - files_to_keep - for file_id in files_to_remove: - Files.delete_file_by_id(file_id) - - # Download and add new/updated files - new_file_ids = list(files_to_keep) - - for gdrive_file in files_to_add: - try: - # Download file from Google Drive - file_content, filename = google_drive_service.download_file( - gdrive_file["id"], gdrive_file - ) - - # Create file object - file_obj = io.BytesIO(file_content) - file_obj.name = filename - - # Upload to storage - file_id = str(uuid.uuid4()) - tags = { - "OpenWebUI-User-Email": "system", - "OpenWebUI-User-Id": "system", - "OpenWebUI-User-Name": "Google Drive Sync", - "OpenWebUI-File-Id": file_id, - } - - contents, file_path = Storage.upload_file( - file_obj, f"{file_id}_{filename}", tags - ) - - # Create file record - file_form = FileForm( - id=file_id, - filename=filename, - path=file_path, - meta={ - "name": filename, - "content_type": "application/octet-stream", - "size": len(contents), - "google_drive_id": gdrive_file["id"], - "google_drive_modified": gdrive_file["modifiedTime"], - "google_drive_path": gdrive_file["path"], - "collection_name": knowledge_base.id, - }, - ) - file_item = Files.insert_new_file(knowledge_base.user_id, file_form) - - if file_item: - new_file_ids.append(file_id) - - # Process file for vector storage - try: - process_file( - ProcessFileForm( - file_id=file_id, - collection_name=knowledge_base.id, - ) - ) - except Exception as e: - log.error(f"Failed to process file {file_id}: {e}") - - except Exception as e: - log.error(f"Failed to sync file {gdrive_file['name']}: {e}") - continue - - # Update knowledge base data - assert knowledge_base.data is not None # Type guard for mypy - updated_data = knowledge_base.data.copy() - updated_data["file_ids"] = new_file_ids - updated_data["google_drive_last_sync"] = int(time.time()) - - Knowledges.update_knowledge_data_by_id( - id=knowledge_base.id, data=updated_data - ) - - log.info( - f"Successfully synced knowledge base {knowledge_base.id} with {len(files_to_add)} new/updated files" - ) - - except Exception as e: - log.error(f"Error syncing knowledge base {knowledge_base.id}: {e}") - - -# Global scheduler instance -google_drive_scheduler = GoogleDriveSyncScheduler() diff --git a/backend/open_webui/utils/content_sources.py b/backend/open_webui/utils/content_sources.py new file mode 100644 index 0000000000..d088a8f23b --- /dev/null +++ b/backend/open_webui/utils/content_sources.py @@ -0,0 +1,38 @@ +"""Utility functions for content source management.""" +import logging +from typing import Optional +from open_webui.content_sources.registry import content_source_registry +from open_webui.content_sources.factory import content_source_factory + +logger = logging.getLogger(__name__) + + +def refresh_provider_configuration(provider_name: str) -> bool: + """ + Refresh a content source provider's configuration. + + This is typically called when configuration settings are updated + in the admin panel. + + Args: + provider_name: Name of the provider to refresh + + Returns: + True if refresh was successful, False otherwise + """ + try: + # Get existing provider instance + provider = content_source_registry.get_provider(provider_name) + + if provider and hasattr(provider, 'refresh_configuration'): + provider.refresh_configuration() + logger.info(f"Refreshed configuration for provider: {provider_name}") + return True + else: + # Provider not loaded yet or doesn't support refresh + logger.debug(f"Provider {provider_name} not loaded or doesn't support refresh") + return False + + except Exception as e: + logger.error(f"Error refreshing provider {provider_name}: {e}") + return False \ No newline at end of file diff --git a/src/lib/apis/knowledge/index.ts b/src/lib/apis/knowledge/index.ts index 3576b25232..755361426a 100644 --- a/src/lib/apis/knowledge/index.ts +++ b/src/lib/apis/knowledge/index.ts @@ -1,5 +1,25 @@ import { WEBUI_API_BASE_URL } from '$lib/constants'; -import type { GoogleDriveServiceAccount, GoogleDriveSyncResponse } from '$lib/types/google-drive'; +import type { ContentSourceProvider } from '$lib/types'; + +/** + * Knowledge Base API + * + * This file contains generic content source endpoints for managing knowledge base content. + * + * Generic Endpoints (Provider-agnostic): + * - getContentSourceInfo(): Get information about any content source provider + * - syncContentSource(): Sync content from any provider to a knowledge base + * - getContentSourceProviders(): Get list of available content source providers + * + * All endpoints support multiple content source providers (Google Drive, OneDrive, etc.) + */ + +// Import generic content source types from centralized types +import type { + ContentSourceServiceInfo as ContentSourceInfo, + ContentSourceSyncConfig, + ContentSourceSyncResults +} from '$lib/types'; export const createNewKnowledge = async ( token: string, @@ -492,12 +512,48 @@ export const deleteKnowledgeById = async (token: string, id: string) => { return res; }; -export const getGoogleDriveServiceAccountEmail = async ( +// Get list of available content source providers +export const getContentSourceProviders = async ( token: string -): Promise => { +): Promise => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/knowledge/google-drive/service-account-email`, { + const res = await fetch(`${WEBUI_API_BASE_URL}/content-sources/`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json.providers || []; + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +// Generic content source functions +export const getContentSourceInfo = async ( + token: string, + provider: string +): Promise => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/content-sources/${provider}/info`, { method: 'GET', headers: { Accept: 'application/json', @@ -525,27 +581,23 @@ export const getGoogleDriveServiceAccountEmail = async ( return res; }; -export const syncGoogleDriveFolder = async ( + +// Generic content source sync function +export const syncContentSource = async ( token: string, knowledgeId: string, - folderId: string, - includeNested: boolean = true, - syncIntervalDays: number = 1 -): Promise => { + config: ContentSourceSyncConfig +): Promise => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/knowledge/${knowledgeId}/google-drive/sync`, { + const res = await fetch(`${WEBUI_API_BASE_URL}/knowledge/${knowledgeId}/sync`, { method: 'POST', headers: { Accept: 'application/json', 'Content-Type': 'application/json', authorization: `Bearer ${token}` }, - body: JSON.stringify({ - folder_id: folderId, - include_nested: includeNested, - sync_interval_days: syncIntervalDays - }) + body: JSON.stringify(config) }) .then(async (res) => { if (!res.ok) throw await res.json(); diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase.svelte index 138ede3f22..96eb5dde96 100644 --- a/src/lib/components/workspace/Knowledge/KnowledgeBase.svelte +++ b/src/lib/components/workspace/Knowledge/KnowledgeBase.svelte @@ -32,6 +32,7 @@ import AddContentMenu from './KnowledgeBase/AddContentMenu.svelte'; import AddTextContentModal from './KnowledgeBase/AddTextContentModal.svelte'; + import ContentSourceStatus from './KnowledgeBase/ContentSourceStatus.svelte'; import Drawer from '$lib/components/common/Drawer.svelte'; @@ -46,6 +47,8 @@ import Pagination from '$lib/components/common/Pagination.svelte'; import AttachWebpageModal from '$lib/components/chat/MessageInput/AttachWebpageModal.svelte'; import GoogleDriveSyncModal from './KnowledgeBase/GoogleDriveSyncModal.svelte'; + import ContentSourceSyncModal from './KnowledgeBase/ContentSourceSyncModal.svelte'; + import { getContentSourceProviders } from '$lib/apis/knowledge'; let largeScreen = true; @@ -65,11 +68,6 @@ description: string; data: { file_ids: string[]; - // Google Drive sync fields - google_drive_folder_id?: string; - google_drive_include_nested?: boolean; - google_drive_sync_interval_days?: number; - google_drive_last_sync?: number; }; files: any[]; access_grants?: any[]; @@ -84,6 +82,9 @@ let selectedFile = null; let selectedFileContent = ''; let showGoogleDriveSyncModal = false; + let showContentSourceSyncModal = false; + let availableProviders = []; + let selectedContentProvider = null; let inputFiles = null; @@ -701,6 +702,16 @@ }; onMount(async () => { + // Load available content source providers + try { + const providers = await getContentSourceProviders(localStorage.token); + if (providers) { + availableProviders = providers.filter(p => p.configured); + } + } catch (error) { + console.error('Failed to load content source providers:', error); + } + // listen to resize 1024px mediaQuery = window.matchMedia('(min-width: 1024px)'); @@ -803,15 +814,18 @@ }} /> - { - knowledge = e.detail; - toast.success($i18n.t('Google Drive folder synced successfully')); - }} -/> +{#if selectedContentProvider} + { + knowledge = e.detail; + toast.success($i18n.t(`${selectedContentProvider.display_name} synced successfully`)); + }} + /> +{/if} + + {#if knowledge} +
+ +
+ {/if} diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase/AddContentMenu.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase/AddContentMenu.svelte index bac0971ef9..6dee31d426 100644 --- a/src/lib/components/workspace/Knowledge/KnowledgeBase/AddContentMenu.svelte +++ b/src/lib/components/workspace/Knowledge/KnowledgeBase/AddContentMenu.svelte @@ -11,10 +11,13 @@ import ArrowPath from '$lib/components/icons/ArrowPath.svelte'; import GlobeAlt from '$lib/components/icons/GlobeAlt.svelte'; import { config } from '$lib/stores'; + import type { ContentSourceProvider } from '$lib/types'; + import { getProviderIcon } from '$lib/utils/content-sources'; const i18n = getContext('i18n'); export let onClose: Function = () => {}; + export let availableProviders: ContentSourceProvider[] = []; export let onSync: Function = () => {}; export let onUpload: Function = (data) => {}; diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceStatus.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceStatus.svelte new file mode 100644 index 0000000000..7cd18e3866 --- /dev/null +++ b/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceStatus.svelte @@ -0,0 +1,41 @@ + +{#if hasSync && provider} +
+ {#if providerIcon} + + + + {/if} + + {getProviderDisplayName(provider)} + {#if syncStatus.sourceId} + • {$i18n.t('Last sync')}: {formatLastSync(syncStatus.lastSync, $i18n)} + {/if} + +
+{/if} \ No newline at end of file diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceSyncModal.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceSyncModal.svelte new file mode 100644 index 0000000000..d392cd08fd --- /dev/null +++ b/src/lib/components/workspace/Knowledge/KnowledgeBase/ContentSourceSyncModal.svelte @@ -0,0 +1,484 @@ + + + +
+
+
+ {$i18n.t(`Sync ${provider.display_name}`)} +
+ +
+
+ +
+ + {#if providerInfo?.metadata} + + {#if provider.name === 'google_drive' && providerInfo.metadata.service_account_email} +
+
+ {$i18n.t('Step 1: Share folder with service account')} +
+
+ {$i18n.t( + 'Copy the service account email below and share your Google Drive folder with this email address:' + )} +
+
+ + +
+
+ {/if} + + + {#if ['onedrive', 'dropbox'].includes(provider.name) && providerInfo.metadata.oauth_status} +
+
+ {$i18n.t('Authentication Status')} +
+
+ {#if providerInfo.metadata.oauth_status === 'connected'} + ✓ {$i18n.t('Connected')} + {:else} + ⚠ {$i18n.t('Not connected')} + + {/if} +
+
+ {/if} + + + {#if providerInfo.metadata.requires_api_key && !providerInfo.metadata.api_key_configured} +
+
+ {$i18n.t('API Key Required')} +
+
+ {$i18n.t('Please configure the API key in the admin settings before syncing.')} +
+
+ {/if} + {/if} + + +
+
+ {providerInfo?.metadata?.service_account_email + ? $i18n.t('Step 2: Enter source') + : $i18n.t('Enter source')} + : {$i18n.t(getSourceLabel(provider.name))} +
+ +
+ {$i18n.t(`You can paste the full ${provider.display_name} URL or just the ID`)} +
+
+ + +
+
{$i18n.t('Sync Options')}
+ + + {#if provider.name === 'google_drive'} +
+ + +
+ {/if} + + + {#if provider.name === 'onedrive'} +
+ + +
+ {/if} + + + {#if provider.name === 'dropbox'} +
+ + +
+ {/if} + + + {#if provider.name === 'sharepoint'} +
+ + +
+ {/if} + + +
+ + +
+ + +
+ + +
+ + +
+
+
+
\ No newline at end of file diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase/GoogleDriveSyncModal.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase/GoogleDriveSyncModal.svelte deleted file mode 100644 index dfd5c6a8fc..0000000000 --- a/src/lib/components/workspace/Knowledge/KnowledgeBase/GoogleDriveSyncModal.svelte +++ /dev/null @@ -1,258 +0,0 @@ - - - -
-
-
{$i18n.t('Sync Google Drive Folder')}
- -
-
- -
- - {#if serviceAccountEmail} -
-
- {$i18n.t('Step 1: Share folder with service account')} -
-
- {$i18n.t( - 'Copy the service account email below and share your Google Drive folder with this email address:' - )} -
-
- - -
-
- {/if} - - -
-
{$i18n.t('Step 2: Enter folder URL or ID')}
- -
- {$i18n.t('You can paste the full Google Drive folder URL or just the folder ID')} -
-
- - -
-
{$i18n.t('Sync Options')}
- - -
- - -
- - -
- - -
-
- - -
- - -
-
-
-
diff --git a/src/lib/types/content-sources.ts b/src/lib/types/content-sources.ts new file mode 100644 index 0000000000..802919612b --- /dev/null +++ b/src/lib/types/content-sources.ts @@ -0,0 +1,161 @@ +// Content Source TypeScript Type Definitions +// Generic types for provider-agnostic content sources +// Mirrors backend Python types for consistency + +/** + * Content source provider information from backend + */ +export interface ContentSourceProvider { + name: string; + display_name: string; + description: string; + configured: boolean; + metadata?: Record; +} + +/** + * Generic content source sync configuration + */ +export interface ContentSourceSyncConfig { + provider: string; + source_id: string; + options?: Record; +} + +/** + * Generic content source sync results + */ +export interface ContentSourceSyncResults { + added_files: string[]; + updated_files: string[]; + removed_files: string[]; + errors: string[]; +} + +/** + * Generic content source file information + */ +export interface ContentSourceFile { + id: string; + name: string; + mimeType: string; + modifiedTime: string; + size?: string | number; + path: string; + webViewLink?: string; + provider: string; + metadata?: Record; +} + +/** + * Generic content source folder information + */ +export interface ContentSourceFolder { + id: string; + name: string; + provider: string; + path?: string; + metadata?: Record; +} + +/** + * Content source service information + */ +export interface ContentSourceServiceInfo { + provider: string; + configured: boolean; + metadata: Record; +} + +/** + * Content source sync status + */ +export interface ContentSourceSyncStatus { + provider: string; + source_id: string; + last_sync: number; // Unix timestamp + sync_interval_days: number; + status: 'idle' | 'syncing' | 'error'; + error_message?: string; +} + +/** + * Knowledge base data with generic content source fields + */ +export interface KnowledgeDataWithContentSource { + id?: string; + name: string; + description?: string; + data?: { + file_ids?: string[]; + // Sync metadata structure + sync_metadata?: { + [provider: string]: { + source_id: string; + last_sync: number; + options?: Record; + results?: { + status: string; + added: number; + updated: number; + failed: number; + duplicates: number; + changes: any; + }; + }; + }; + }; + files?: any[]; +} + +/** + * Provider capabilities information + */ +export interface ContentSourceCapabilities { + supports_folders: boolean; + supports_recursive_sync: boolean; + supports_selective_sync: boolean; + supports_incremental_sync: boolean; + supports_webhooks: boolean; + max_file_size?: number; + supported_mime_types?: string[]; +} + +/** + * Content source authentication configuration + */ +export interface ContentSourceAuthConfig { + type: 'oauth2' | 'api_key' | 'service_account' | 'none'; + oauth2?: { + client_id: string; + scopes: string[]; + auth_url?: string; + token_url?: string; + }; + api_key?: { + header_name: string; + required: boolean; + }; + service_account?: { + email: string; + status: 'active' | 'inactive'; + }; +} + +/** + * Helper type guards + */ +export function isContentSourceConfigured(source: ContentSourceProvider): boolean { + return source.configured; +} + +export function hasContentSourceSync(data: KnowledgeDataWithContentSource): boolean { + return !!data.data?.content_source?.source_id; +} + +export function getContentSourceProvider(data: KnowledgeDataWithContentSource): string | undefined { + if (data.data?.content_source?.provider) { + return data.data.content_source.provider; + } + return undefined; +} \ No newline at end of file diff --git a/src/lib/types/google-drive.ts b/src/lib/types/google-drive.ts index 222b36e601..44a725dd42 100644 --- a/src/lib/types/google-drive.ts +++ b/src/lib/types/google-drive.ts @@ -1,6 +1,15 @@ // Google Drive TypeScript Type Definitions +// Specific types for Google Drive integration // Mirrors backend Python types for consistency +import type { KnowledgeDataWithContentSource } from './content-sources'; + +/** + * Google Drive Specific Types + * These types are specific to Google Drive integration + * For generic content source types, see content-sources.ts + */ + /** * Google Drive file information from Google Drive API */ @@ -166,22 +175,6 @@ export interface GoogleOAuthError { message?: string; } -/** - * Knowledge base data with Google Drive fields - */ -export interface KnowledgeDataWithGoogleDrive { - id?: string; - name: string; - description?: string; - data?: Record; - // Google Drive specific fields - google_drive_folder_id?: string; - google_drive_include_nested?: boolean; - google_drive_sync_interval_days?: number; - google_drive_last_sync?: number; - google_drive_file_ids?: string[]; -} - /** * API error response structure */ diff --git a/src/lib/types/index.ts b/src/lib/types/index.ts index a48fb58bec..03c331b59c 100644 --- a/src/lib/types/index.ts +++ b/src/lib/types/index.ts @@ -14,5 +14,9 @@ export enum TTS_RESPONSE_SPLIT { NONE = 'none' } -// Re-export Google Drive types for convenience +// Re-export content source types (generic provider-agnostic types) +export * from './content-sources'; + +// Re-export Google Drive specific types export * from './google-drive'; + diff --git a/src/lib/utils/content-sources.ts b/src/lib/utils/content-sources.ts new file mode 100644 index 0000000000..4a24fe3a2a --- /dev/null +++ b/src/lib/utils/content-sources.ts @@ -0,0 +1,130 @@ +/** + * Utility functions for content source providers + */ + +import type { ContentSourceProvider, KnowledgeDataWithContentSource } from '$lib/types'; + +// Provider icon SVG paths - these match the paths used in AddContentMenu.svelte +export const providerIcons: Record = { + google_drive: 'M12.01 1.485c-2.082 0-3.754.02-3.743.047.01.02 1.708 3.001 3.774 6.62l3.76 6.574h3.76c2.081 0 3.753-.02 3.742-.047-.005-.02-1.708-3.001-3.775-6.62l-3.76-6.574zm-4.76 1.73a789.828 789.861 0 0 0-3.63 6.319L0 15.868l1.89 3.298 1.885 3.297 3.62-6.335 3.618-6.33-1.88-3.287C8.1 4.704 7.255 3.22 7.25 3.214zm2.259 12.653-.203.348c-.114.198-.96 1.672-1.88 3.287a423.93 423.948 0 0 1-1.698 2.97c-.01.026 3.24.042 7.222.042h7.244l1.796-3.157c.992-1.734 1.85-3.23 1.906-3.323l.104-.167h-7.249z', + onedrive: 'M12.188 5.813q-1.325 0-2.52.544-1.195.545-2.04 1.565.446.117.85.299.405.181.792.416l4.78 2.86 2.731-1.15q.27-.117.545-.204.276-.088.58-.147-.293-.937-.855-1.705-.563-.768-1.319-1.318-.755-.551-1.658-.856-.902-.304-1.886-.304zM2.414 16.395l9.914-4.184-3.832-2.297q-.586-.351-1.23-.539-.645-.188-1.325-.188-.914 0-1.722.364-.809.363-1.412.978-.603.615-.967 1.424-.363.808-.363 1.722 0 .62.163 1.201.164.58.469 1.09.305.509.738.897.434.387.967.65.533.262 1.13.387.598.126 1.225.126h13.125q.773 0 1.453-.3.68-.299 1.19-.808.51-.51.809-1.19.299-.68.299-1.453 0-.738-.28-1.389-.282-.65-.768-1.136-.486-.486-1.143-.768-.656-.281-1.4-.281h-.047q-.023 0-.047.006-.316-1.242-1.008-2.28-.691-1.036-1.658-1.78-.967-.744-2.144-1.16-1.178-.417-2.456-.417-.949 0-1.845.229-.897.228-1.705.668-.809.439-1.5 1.066-.692.627-1.207 1.413h-.012q-.445.022-.861.082-.416.058-.85.187-.937.293-1.711.861-.774.569-1.324 1.325-.551.755-.862 1.658-.31.902-.31 1.887 0 1.242.474 2.332.475 1.09 1.29 1.904.814.815 1.903 1.29 1.09.475 2.332.475z', + dropbox: 'M12 5L6 9l6 4-6 4-6-4 6-4L0 5l6-4zm-6 14l6-4 6 4-6 4zm6-6l6-4-6-4 6-4 6 4-6 4 6 4-6 4z', + sharepoint: 'M24 13.5q0 1.242-.475 2.332-.474 1.09-1.289 1.904-.814.815-1.904 1.29-1.09.474-2.332.474-.762 0-1.523-.2-.106.997-.557 1.858-.451.862-1.154 1.494-.704.633-1.606.99-.902.358-1.91.358-1.09 0-2.045-.416-.955-.416-1.664-1.125-.709-.709-1.125-1.664Q6 19.84 6 18.75q0-.188.018-.375.017-.188.04-.375H.997q-.41 0-.703-.293T0 17.004V6.996q0-.41.293-.703T.996 6h3.54q.14-1.277.726-2.373.586-1.096 1.488-1.904Q7.652.914 8.807.457 9.96 0 11.25 0q1.395 0 2.625.533T16.02 1.98q.914.915 1.447 2.145T18 6.75q0 .188-.012.375-.011.188-.035.375 1.242 0 2.344.469 1.101.468 1.928 1.277.826.809 1.3 1.904Q24 12.246 24 13.5zm-12.75-12q-.973 0-1.857.34-.885.34-1.577.943-.691.604-1.154 1.43Q6.2 5.039 6.06 6h4.945q.41 0 .703.293t.293.703v4.945l.21-.035q.212-.75.61-1.424.399-.673.944-1.218.545-.545 1.213-.944.668-.398 1.43-.61.093-.503.093-.96 0-1.09-.416-2.045-.416-.955-1.125-1.664' +}; + +/** + * Get provider icon SVG path + */ +export function getProviderIcon(providerName: string): string | null { + return providerIcons[providerName] || null; +} + +/** + * Get provider display name with fallback + */ +export function getProviderDisplayName(provider: ContentSourceProvider | string): string { + if (typeof provider === 'string') { + // Fallback display names for common providers + const fallbackNames: Record = { + google_drive: 'Google Drive', + onedrive: 'OneDrive', + dropbox: 'Dropbox', + sharepoint: 'SharePoint' + }; + return fallbackNames[provider] || provider; + } + return provider.display_name || provider.name; +} + +/** + * Check if a knowledge base has content source sync configured + */ +export function hasContentSourceSync(data: KnowledgeDataWithContentSource): boolean { + // Check both old format (sync_metadata) and new format (content_source) + if (data.data?.sync_metadata) { + // Check if any provider has a source_id + const providers = Object.keys(data.data.sync_metadata); + return providers.some(provider => + data.data.sync_metadata[provider]?.source_id + ); + } + return !!data.data?.content_source?.source_id; +} + +/** + * Get the configured content source provider name + */ +export function getConfiguredProvider(data: KnowledgeDataWithContentSource): string | null { + // Check old format first (sync_metadata) + if (data.data?.sync_metadata) { + // Return the first provider that has a source_id + const providers = Object.keys(data.data.sync_metadata); + for (const provider of providers) { + if (data.data.sync_metadata[provider]?.source_id) { + return provider; + } + } + } + // Check new format + return data.data?.content_source?.provider || null; +} + +/** + * Get sync status information + */ +export function getSyncStatus(data: KnowledgeDataWithContentSource): { + lastSync: number | null; + intervalDays: number; + sourceId: string | null; +} { + // Check old format first (sync_metadata) + const provider = getConfiguredProvider(data); + if (provider && data.data?.sync_metadata?.[provider]) { + const syncData = data.data.sync_metadata[provider]; + return { + lastSync: syncData.last_sync || null, + intervalDays: syncData.options?.sync_interval_days || 1, + sourceId: syncData.source_id || null + }; + } + + // Check new format + if (data.data?.content_source) { + return { + lastSync: data.data.content_source.last_sync || null, + intervalDays: data.data.content_source.sync_interval_days || 1, + sourceId: data.data.content_source.source_id || null + }; + } + + return { + lastSync: null, + intervalDays: 1, + sourceId: null + }; +} + +/** + * Format last sync time for display + */ +export function formatLastSync(timestamp: number | null, i18n: any): string { + if (!timestamp) { + return i18n.t('Never synced'); + } + + const now = Date.now(); + const diff = now - timestamp * 1000; // Convert to milliseconds + const minutes = Math.floor(diff / 60000); + const hours = Math.floor(diff / 3600000); + const days = Math.floor(diff / 86400000); + + if (minutes < 1) { + return i18n.t('Just now'); + } else if (minutes < 60) { + return i18n.t('{{count}} minutes ago', { count: minutes }); + } else if (hours < 24) { + return i18n.t('{{count}} hours ago', { count: hours }); + } else { + return i18n.t('{{count}} days ago', { count: days }); + } +} \ No newline at end of file