diff --git a/api/apps/restful_apis/connector_api.py b/api/apps/restful_apis/connector_api.py index 7a86876fc6..9a62c116e0 100644 --- a/api/apps/restful_apis/connector_api.py +++ b/api/apps/restful_apis/connector_api.py @@ -28,7 +28,7 @@ from api.db import InputType from api.db.services.connector_service import ConnectorService, SyncLogsService from api.utils.api_utils import get_data_error_result, get_json_result, get_request_json, validate_request from api.utils.pagination_utils import DEFAULT_PAGE, DEFAULT_PAGE_SIZE, validate_rest_api_page, validate_rest_api_page_size -from common.constants import RetCode, TaskStatus +from common.constants import FileSource, RetCode, TaskStatus from common.data_source.config import GOOGLE_DRIVE_WEB_OAUTH_REDIRECT_URI, GMAIL_WEB_OAUTH_REDIRECT_URI, BOX_WEB_OAUTH_REDIRECT_URI, DocumentSource from common.data_source.google_util.constant import WEB_OAUTH_POPUP_TEMPLATE, GOOGLE_SCOPES from common.misc_utils import get_uuid @@ -180,102 +180,48 @@ def rm_connector(connector_id): @manager.route("/connectors//test", methods=["POST"]) # noqa: F821 @login_required +@validate_request("source") async def test_connector(connector_id): - """Validate connector configuration without persisting changes or triggering sync. - - For the REST API connector, this uses `RestAPIConnector.validate_config` - against the existing saved configuration. For BigQuery, it runs `SELECT 1` - plus a free dry-run of the configured base query under `maximum_bytes_billed` - so bad credentials, wrong location, or runaway scans surface before scheduled - syncs run (and incur cost). - """ - if not ConnectorService.accessible(connector_id, current_user.id): + """Validate connector configuration from the request body without persisting.""" + unsaved = connector_id in {source.value for source in FileSource if source.value} + if not unsaved and not ConnectorService.accessible(connector_id, current_user.id): return _connector_auth_error(connector_id, current_user.id) + from common.data_source import build_connector_for_source from common.data_source.exceptions import ConnectorMissingCredentialError, ConnectorValidationError - ok, conn = ConnectorService.get_by_id(connector_id) - if not ok: - return get_data_error_result(message="Can't find this Connector!") + req = await get_request_json() + source = req["source"] + config = req.get("config") or {} + if not isinstance(config, dict): + return get_json_result(code=RetCode.ARGUMENT_ERROR, message="config must be an object.") - config = conn.config or {} - credentials = config.get("credentials") or {} + if not unsaved: + ok, conn = ConnectorService.get_by_id(connector_id) + if ok and conn.tenant_id != current_user.id: + return get_json_result(code=RetCode.PERMISSION_ERROR, message="You don't own this connector.") - if conn.source == DocumentSource.REST_API: - from common.data_source.rest_api_connector import RestAPIConnector + def _validate() -> None: + connector = build_connector_for_source(source, config) + connector.validate_connector_settings() - try: - await asyncio.to_thread( - RestAPIConnector.validate_config, - config=config, - credentials=credentials, - ) - except (ConnectorValidationError, ConnectorMissingCredentialError) as exc: - return get_json_result( - code=RetCode.DATA_ERROR, - message=str(exc), - data=False, - ) - except Exception as exc: - logging.exception("REST API connector validation failed: %s", exc) - return get_json_result( - code=RetCode.SERVER_ERROR, - message="REST API connector validation failed, please check logs.", - data=False, - ) + try: + await asyncio.to_thread(_validate) + except (ConnectorValidationError, ConnectorMissingCredentialError) as exc: + return get_json_result( + code=RetCode.DATA_ERROR, + message=str(exc), + data=False, + ) + except Exception as exc: + logging.exception("Connector validation failed for %s: %s", connector_id, exc) + return get_json_result( + code=RetCode.SERVER_ERROR, + message="Connector validation failed, please check logs.", + data=False, + ) - return get_json_result(data=True) - - if conn.source == DocumentSource.BIGQUERY: - from common.data_source.bigquery_connector import BigQueryConnector - - def _validate_bigquery(): - connector_kwargs = { - "project_id": config.get("project_id", ""), - "dataset_id": config.get("dataset_id") or None, - "table_id": config.get("table_id") or None, - "location": config.get("location") or None, - "query": config.get("query", ""), - "content_columns": config.get("content_columns", ""), - "metadata_columns": config.get("metadata_columns", ""), - "id_column": config.get("id_column") or None, - "timestamp_column": config.get("timestamp_column") or None, - "use_query_cache": config.get("use_query_cache", True), - } - if config.get("page_size") is not None: - connector_kwargs["page_size"] = int(config["page_size"]) - if config.get("maximum_bytes_billed") is not None: - connector_kwargs["maximum_bytes_billed"] = int(config["maximum_bytes_billed"]) - if config.get("job_timeout_ms") is not None: - connector_kwargs["job_timeout_ms"] = int(config["job_timeout_ms"]) - - connector = BigQueryConnector(**connector_kwargs) - connector.load_credentials(credentials) - connector.validate_connector_settings() - - try: - await asyncio.to_thread(_validate_bigquery) - except (ConnectorValidationError, ConnectorMissingCredentialError) as exc: - return get_json_result( - code=RetCode.DATA_ERROR, - message=str(exc), - data=False, - ) - except Exception as exc: - logging.exception("BigQuery connector validation failed: %s", exc) - return get_json_result( - code=RetCode.SERVER_ERROR, - message="BigQuery connector validation failed, please check logs.", - data=False, - ) - - return get_json_result(data=True) - - return get_json_result( - code=RetCode.ARGUMENT_ERROR, - message="Test endpoint currently supports only REST API and BigQuery connectors.", - data=False, - ) + return get_json_result(data=True) WEB_FLOW_TTL_SECS = 15 * 60 diff --git a/common/data_source/__init__.py b/common/data_source/__init__.py index 158e5450b5..f0865a3ead 100644 --- a/common/data_source/__init__.py +++ b/common/data_source/__init__.py @@ -22,36 +22,100 @@ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. """ +from typing import Any + +from common.constants import FileSource + +from .airtable_connector import AirtableConnector +from .asana_connector import AsanaConnector +from .azure_blob_connector import AzureBlobConnector +from .bigquery_connector import BigQueryConnector +from .bitbucket.connector import BitbucketConnector from .blob_connector import BlobStorageConnector -from .rss_connector import RSSConnector -from .slack_connector import SlackConnector -from .gmail_connector import GmailConnector -from .notion_connector import NotionConnector +from .box_connector import BoxConnector from .confluence_connector import ConfluenceConnector +from .config import BlobType, DocumentSource +from .dingtalk_ai_table_connector import DingTalkAITableConnector from .discord_connector import DiscordConnector from .dropbox_connector import DropboxConnector +from .exceptions import ( + ConnectorMissingCredentialError, + ConnectorValidationError, + CredentialExpiredError, + InsufficientPermissionsError, + UnexpectedValidationError, +) +from .github.connector import GithubConnector +from .gitlab_connector import GitlabConnector +from .gmail_connector import GmailConnector from .google_drive.connector import GoogleDriveConnector +from .imap_connector import ImapConnector from .jira.connector import JiraConnector -from .sharepoint_connector import SharePointConnector +from .models import BasicExpertInfo, Document, ImageSection, TextSection +from .moodle_connector import MoodleConnector +from .notion_connector import NotionConnector from .onedrive_connector import OneDriveConnector from .outlook_connector import OutlookConnector -from .salesforce_connector import SalesforceConnector -from .azure_blob_connector import AzureBlobConnector -from .teams_connector import TeamsConnector -from .moodle_connector import MoodleConnector -from .airtable_connector import AirtableConnector -from .dingtalk_ai_table_connector import DingTalkAITableConnector -from .asana_connector import AsanaConnector -from .imap_connector import ImapConnector -from .zendesk_connector import ZendeskConnector -from .seafile_connector import SeaFileConnector from .rdbms_connector import RDBMSConnector -from .bigquery_connector import BigQueryConnector -from .webdav_connector import WebDAVConnector from .rest_api_connector import RestAPIConnector -from .config import BlobType, DocumentSource -from .models import Document, TextSection, ImageSection, BasicExpertInfo -from .exceptions import ConnectorMissingCredentialError, ConnectorValidationError, CredentialExpiredError, InsufficientPermissionsError, UnexpectedValidationError +from .rss_connector import RSSConnector +from .salesforce_connector import SalesforceConnector +from .seafile_connector import SeaFileConnector +from .sharepoint_connector import SharePointConnector +from .slack_connector import SlackConnector +from .teams_connector import TeamsConnector +from .webdav_connector import WebDAVConnector +from .zendesk_connector import ZendeskConnector + +CONNECTOR_BY_SOURCE: dict[str, type] = { + FileSource.S3: BlobStorageConnector, + FileSource.R2: BlobStorageConnector, + FileSource.OCI_STORAGE: BlobStorageConnector, + FileSource.GOOGLE_CLOUD_STORAGE: BlobStorageConnector, + FileSource.RSS: RSSConnector, + FileSource.CONFLUENCE: ConfluenceConnector, + FileSource.NOTION: NotionConnector, + FileSource.DISCORD: DiscordConnector, + FileSource.GMAIL: GmailConnector, + FileSource.DROPBOX: DropboxConnector, + FileSource.GOOGLE_DRIVE: GoogleDriveConnector, + FileSource.JIRA: JiraConnector, + FileSource.SHAREPOINT: SharePointConnector, + FileSource.SLACK: SlackConnector, + FileSource.TEAMS: TeamsConnector, + FileSource.WEBDAV: WebDAVConnector, + FileSource.MOODLE: MoodleConnector, + FileSource.BOX: BoxConnector, + FileSource.AIRTABLE: AirtableConnector, + FileSource.ASANA: AsanaConnector, + FileSource.GITHUB: GithubConnector, + FileSource.IMAP: ImapConnector, + FileSource.ZENDESK: ZendeskConnector, + FileSource.GITLAB: GitlabConnector, + FileSource.BITBUCKET: BitbucketConnector, + FileSource.SEAFILE: SeaFileConnector, + FileSource.DINGTALK_AI_TABLE: DingTalkAITableConnector, + FileSource.MYSQL: RDBMSConnector, + FileSource.POSTGRESQL: RDBMSConnector, + FileSource.REST_API: RestAPIConnector, + FileSource.BIGQUERY: BigQueryConnector, + FileSource.ONEDRIVE: OneDriveConnector, + FileSource.OUTLOOK: OutlookConnector, + FileSource.SALESFORCE: SalesforceConnector, + FileSource.AZURE_BLOB: AzureBlobConnector, +} + + +def build_connector_for_source(source: str, config: dict[str, Any]) -> Any: + connector_cls = CONNECTOR_BY_SOURCE.get(source) + if connector_cls is None: + raise ConnectorValidationError(f"Unsupported data source type: {source}") + if connector_cls is BlobStorageConnector: + return connector_cls.build_connector(config, bucket_type=source) + if connector_cls is RDBMSConnector: + return connector_cls.build_connector(config, db_type=source) + return connector_cls.build_connector(config) + __all__ = [ "BlobStorageConnector", @@ -65,6 +129,10 @@ __all__ = [ "GoogleDriveConnector", "JiraConnector", "SharePointConnector", + "GithubConnector", + "GitlabConnector", + "BitbucketConnector", + "BoxConnector", "OneDriveConnector", "OutlookConnector", "SalesforceConnector", @@ -92,4 +160,6 @@ __all__ = [ "WebDAVConnector", "DingTalkAITableConnector", "RestAPIConnector", + "CONNECTOR_BY_SOURCE", + "build_connector_for_source", ] diff --git a/common/data_source/airtable_connector.py b/common/data_source/airtable_connector.py index 2ab471191c..ed35bbfa90 100644 --- a/common/data_source/airtable_connector.py +++ b/common/data_source/airtable_connector.py @@ -7,7 +7,7 @@ import requests from pyairtable import Api as AirtableApi from common.data_source.config import AIRTABLE_CONNECTOR_SIZE_THRESHOLD, INDEX_BATCH_SIZE, DocumentSource -from common.data_source.exceptions import ConnectorMissingCredentialError +from common.data_source.exceptions import ConnectorMissingCredentialError, ConnectorValidationError from common.data_source.interfaces import LoadConnector, PollConnector, SlimConnectorWithPermSync from common.data_source.models import ( Document, @@ -81,10 +81,31 @@ class AirtableConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync) # ------------------------- # Credentials # ------------------------- + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "AirtableConnector": + credentials = config.get("credentials") or {} + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + base_id=config.get("base_id"), + table_name_or_id=config.get("table_name_or_id"), + batch_size=batch_size, + ) + connector.load_credentials({"airtable_access_token": credentials["airtable_access_token"]}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: self._airtable_client = AirtableApi(credentials["airtable_access_token"]) return None + def validate_connector_settings(self) -> None: + if not self._airtable_client: + raise ConnectorMissingCredentialError("Airtable credentials not loaded.") + + try: + self.airtable_client.table(self.base_id, self.table_name_or_id).all(max_records=1) + except Exception as e: + raise ConnectorValidationError(f"Failed to validate Airtable connector settings: {e}") from e + @property def airtable_client(self) -> AirtableApi: if not self._airtable_client: diff --git a/common/data_source/asana_connector.py b/common/data_source/asana_connector.py index 2592888807..7e4375af0f 100644 --- a/common/data_source/asana_connector.py +++ b/common/data_source/asana_connector.py @@ -6,6 +6,7 @@ from typing import Any, Dict import asana import requests from common.data_source.config import CONTINUE_ON_CONNECTOR_FAILURE, INDEX_BATCH_SIZE, DocumentSource +from common.data_source.exceptions import ConnectorMissingCredentialError, ConnectorValidationError from common.data_source.interfaces import LoadConnector, PollConnector, SlimConnectorWithPermSync from common.data_source.models import Document, GenerateDocumentsOutput, GenerateSlimDocumentOutput, SecondsSinceUnixEpoch, SlimDocument from common.data_source.utils import extract_size_bytes, get_file_ext @@ -329,6 +330,19 @@ class AsanaConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): self.size_threshold = None logging.info(f"AsanaConnector initialized with workspace_id: {asana_workspace_id}") + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "AsanaConnector": + credentials = config.get("credentials") or {} + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + config.get("asana_workspace_id"), + config.get("asana_project_ids"), + config.get("asana_team_id"), + batch_size=batch_size, + ) + connector.load_credentials({"asana_api_token_secret": credentials["asana_api_token_secret"]}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: self.api_token = credentials["asana_api_token_secret"] self.asana_client = AsanaAPI( @@ -340,6 +354,15 @@ class AsanaConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): logging.info("Asana credentials loaded and API client initialized") return None + def validate_connector_settings(self) -> None: + if not hasattr(self, "asana_client") or self.asana_client is None: + raise ConnectorMissingCredentialError("Asana credentials not loaded.") + + try: + self.asana_client.workspaces_api.get_workspace(self.workspace_id, {}) + except Exception as e: + raise ConnectorValidationError(f"Failed to validate Asana connector settings: {e}") from e + def poll_source(self, start: SecondsSinceUnixEpoch, end: SecondsSinceUnixEpoch | None) -> GenerateDocumentsOutput: start_time = datetime.fromtimestamp(start, tz=timezone.utc).isoformat() end_time = datetime.fromtimestamp(end, tz=timezone.utc) if end is not None else None diff --git a/common/data_source/azure_blob_connector.py b/common/data_source/azure_blob_connector.py index 5ac88b10fc..498c3dd91c 100644 --- a/common/data_source/azure_blob_connector.py +++ b/common/data_source/azure_blob_connector.py @@ -105,6 +105,18 @@ class AzureBlobConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPer self.auth_mode = (auth_mode or "").strip().lower() self._container_client = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "AzureBlobConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + batch_size=batch_size, + prefix=config.get("prefix") or None, + allow_images=bool(config.get("allow_images", False)), + auth_mode=config.get("auth_mode") or None, + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + # ------------------------------------------------------------------ # Auth # ------------------------------------------------------------------ diff --git a/common/data_source/bigquery_connector.py b/common/data_source/bigquery_connector.py index 00489c5301..09ff99db48 100644 --- a/common/data_source/bigquery_connector.py +++ b/common/data_source/bigquery_connector.py @@ -138,6 +138,32 @@ class BigQueryConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync) self._pending_sync_cursor_value: Any = None self._pending_sync_cursor_id: Any = None + @classmethod + def build_connector(cls, config: Dict[str, Any]) -> "BigQueryConnector": + connector_kwargs: Dict[str, Any] = { + "project_id": config.get("project_id", ""), + "dataset_id": config.get("dataset_id") or None, + "table_id": config.get("table_id") or None, + "location": config.get("location") or None, + "query": config.get("query", ""), + "content_columns": config.get("content_columns", ""), + "metadata_columns": config.get("metadata_columns", ""), + "id_column": config.get("id_column") or None, + "timestamp_column": config.get("timestamp_column") or None, + "use_query_cache": config.get("use_query_cache", True), + } + if config.get("batch_size") is not None: + connector_kwargs["batch_size"] = int(config["batch_size"]) + if config.get("page_size") is not None: + connector_kwargs["page_size"] = int(config["page_size"]) + if config.get("maximum_bytes_billed") is not None: + connector_kwargs["maximum_bytes_billed"] = int(config["maximum_bytes_billed"]) + if config.get("job_timeout_ms") is not None: + connector_kwargs["job_timeout_ms"] = int(config["job_timeout_ms"]) + connector = cls(**connector_kwargs) + connector.load_credentials(config.get("credentials") or {}) + return connector + # ------------------------------------------------------------------ # # Credentials & client # ------------------------------------------------------------------ # diff --git a/common/data_source/bitbucket/connector.py b/common/data_source/bitbucket/connector.py index 0e570f7786..798ecd9890 100644 --- a/common/data_source/bitbucket/connector.py +++ b/common/data_source/bitbucket/connector.py @@ -82,6 +82,22 @@ class BitbucketConnector( self.email: str | None = None self.api_token: str | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "BitbucketConnector": + credentials = config.get("credentials") or {} + connector = cls( + workspace=config.get("workspace"), + repositories=config.get("repository_slugs"), + projects=config.get("projects"), + ) + connector.load_credentials( + { + "bitbucket_email": credentials.get("bitbucket_account_email"), + "bitbucket_api_token": credentials.get("bitbucket_api_token"), + } + ) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Load API token-based credentials. diff --git a/common/data_source/blob_connector.py b/common/data_source/blob_connector.py index d66e5a0993..ed6f3f1e1b 100644 --- a/common/data_source/blob_connector.py +++ b/common/data_source/blob_connector.py @@ -76,6 +76,19 @@ class BlobStorageConnector(LoadConnector, PollConnector, FingerprintConnector): logging.info(f"Setting allow_images to {allow_images}.") self._allow_images = allow_images + @classmethod + def build_connector(cls, config: dict[str, Any], *, bucket_type: str) -> "BlobStorageConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + bucket_type=bucket_type, + bucket_name=config["bucket_name"], + prefix=config.get("prefix", ""), + batch_size=batch_size, + ) + connector.set_allow_images(config.get("allow_images", False)) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Load credentials""" logging.debug(f"Loading credentials for {self.bucket_name} of type {self.bucket_type}") diff --git a/common/data_source/box_connector.py b/common/data_source/box_connector.py index 51be0ffbcc..47d38260ac 100644 --- a/common/data_source/box_connector.py +++ b/common/data_source/box_connector.py @@ -1,10 +1,11 @@ """Box connector""" +import json import logging from datetime import datetime, timezone from typing import Any, Generator -from box_sdk_gen import BoxClient +from box_sdk_gen import AccessToken, BoxClient, BoxOAuth, OAuthConfig from common.data_source.config import DocumentSource, INDEX_BATCH_SIZE from common.data_source.exceptions import ( ConnectorMissingCredentialError, @@ -22,6 +23,29 @@ class BoxConnector(LoadConnector, PollConnector): self.use_marker = use_marker self.box_client: BoxClient | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "BoxConnector": + credentials = config.get("credentials") or {} + box_tokens = credentials.get("box_tokens") + if not box_tokens: + raise ConnectorMissingCredentialError("Box tokens are required.") + token_payload = json.loads(box_tokens) if isinstance(box_tokens, str) else box_tokens + auth = BoxOAuth( + OAuthConfig( + client_id=token_payload["client_id"], + client_secret=token_payload["client_secret"], + ) + ) + auth.token_storage.store( + AccessToken( + access_token=token_payload["access_token"], + refresh_token=token_payload["refresh_token"], + ) + ) + connector = cls(folder_id=config.get("folder_id", "0")) + connector.load_credentials(auth) + return connector + def load_credentials(self, auth: Any): self.box_client = BoxClient(auth=auth) return None diff --git a/common/data_source/confluence_connector.py b/common/data_source/confluence_connector.py index bb447ecdcd..af4c74a7d2 100644 --- a/common/data_source/confluence_connector.py +++ b/common/data_source/confluence_connector.py @@ -1284,6 +1284,34 @@ class ConfluenceConnector( raise ConnectorMissingCredentialError("Confluence") return self._low_timeout_confluence_client + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "ConfluenceConnector": + index_mode = (config.get("index_mode") or "everything").lower() + space = "" + page_id = "" + index_recursively = False + if index_mode == "space": + space = (config.get("space") or "").strip() + elif index_mode == "page": + page_id = (config.get("page_id") or "").strip() + index_recursively = bool(config.get("index_recursively", False)) + + connector = cls( + wiki_base=config["wiki_base"], + is_cloud=config.get("is_cloud", True), + space=space, + page_id=page_id, + index_recursively=index_recursively, + ) + connector.set_credentials_provider( + StaticCredentialsProvider( + tenant_id=None, + connector_name=DocumentSource.CONFLUENCE, + credential_json=config.get("credentials") or {}, + ) + ) + return connector + def set_credentials_provider(self, credentials_provider: CredentialsProviderInterface) -> None: self.credentials_provider = credentials_provider diff --git a/common/data_source/dingtalk_ai_table_connector.py b/common/data_source/dingtalk_ai_table_connector.py index 40dc44b61f..983e278ea9 100644 --- a/common/data_source/dingtalk_ai_table_connector.py +++ b/common/data_source/dingtalk_ai_table_connector.py @@ -85,6 +85,18 @@ class DingTalkAITableConnector(LoadConnector, PollConnector, SlimConnectorWithPe config.region_id = "central" return NotableClient(config) + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "DingTalkAITableConnector": + credentials = config.get("credentials") or {} + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + table_id=config.get("table_id"), + operator_id=config.get("operator_id"), + batch_size=batch_size, + ) + connector.load_credentials({"access_token": credentials["access_token"]}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """ Load DingTalk credentials. diff --git a/common/data_source/discord_connector.py b/common/data_source/discord_connector.py index 02deddc746..74bf0622e9 100644 --- a/common/data_source/discord_connector.py +++ b/common/data_source/discord_connector.py @@ -248,6 +248,19 @@ class DiscordConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): self._discord_bot_token: str | None = None self.requested_start_date_string: str = start_date or "" + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "DiscordConnector": + server_ids = config.get("server_ids") + channel_names = config.get("channel_names") + connector = cls( + server_ids=server_ids.split(",") if server_ids else [], + channel_names=channel_names.split(",") if channel_names else [], + start_date=datetime(1970, 1, 1, tzinfo=timezone.utc).strftime("%Y-%m-%d"), + batch_size=int(config.get("batch_size") or INDEX_BATCH_SIZE), + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + @property def discord_bot_token(self) -> str: if self._discord_bot_token is None: diff --git a/common/data_source/dropbox_connector.py b/common/data_source/dropbox_connector.py index 43ab08f4b0..71ae65cb26 100644 --- a/common/data_source/dropbox_connector.py +++ b/common/data_source/dropbox_connector.py @@ -28,6 +28,13 @@ class DropboxConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): self.batch_size = batch_size self.dropbox_client: Dropbox | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "DropboxConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Load Dropbox credentials""" access_token = credentials.get("dropbox_access_token") diff --git a/common/data_source/github/connector.py b/common/data_source/github/connector.py index 1e9be17d1b..76b7c69c9a 100644 --- a/common/data_source/github/connector.py +++ b/common/data_source/github/connector.py @@ -358,6 +358,18 @@ class GithubConnector(CheckpointedConnectorWithPermSyncGH[GithubConnectorCheckpo self.include_issues = include_issues self.github_client: Github | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "GithubConnector": + credentials = config.get("credentials") or {} + connector = cls( + repo_owner=config.get("repository_owner"), + repositories=config.get("repository_name"), + include_prs=config.get("include_pull_requests", True), + include_issues=config.get("include_issues", True), + ) + connector.load_credentials({"github_access_token": credentials["github_access_token"]}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: # defaults to 30 items per page, can be set to as high as 100 token = credentials["github_access_token"] diff --git a/common/data_source/gitlab_connector.py b/common/data_source/gitlab_connector.py index 2547ea54de..ff151f978b 100644 --- a/common/data_source/gitlab_connector.py +++ b/common/data_source/gitlab_connector.py @@ -176,6 +176,24 @@ class GitlabConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): self.include_code_files = include_code_files self.gitlab_client: gitlab.Gitlab | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "GitlabConnector": + credentials = config.get("credentials") or {} + connector = cls( + project_owner=config.get("project_owner"), + project_name=config.get("project_name"), + include_mrs=config.get("include_mrs", False), + include_issues=config.get("include_issues", False), + include_code_files=config.get("include_code_files", False), + ) + connector.load_credentials( + { + "gitlab_access_token": credentials.get("gitlab_access_token"), + "gitlab_url": config.get("gitlab_url"), + } + ) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: self.gitlab_client = gitlab.Gitlab(credentials["gitlab_url"], private_token=credentials["gitlab_access_token"]) return None diff --git a/common/data_source/gmail_connector.py b/common/data_source/gmail_connector.py index bdd03d6c14..4c0e582c14 100644 --- a/common/data_source/gmail_connector.py +++ b/common/data_source/gmail_connector.py @@ -8,6 +8,12 @@ from common.data_source.config import INDEX_BATCH_SIZE, SLIM_BATCH_SIZE, Documen from common.data_source.google_util.auth import get_google_creds from common.data_source.google_util.constant import DB_CREDENTIALS_PRIMARY_ADMIN_KEY, MISSING_SCOPES_ERROR_STR, SCOPE_INSTRUCTIONS, USER_FIELDS from common.data_source.google_util.resource import get_admin_service, get_gmail_service +from common.data_source.exceptions import ( + ConnectorMissingCredentialError, + ConnectorValidationError, + CredentialExpiredError, + InsufficientPermissionsError, +) from common.data_source.google_util.util import _execute_single_retrieval, execute_paginated_retrieval, clean_string from common.data_source.interfaces import LoadConnector, PollConnector, SecondsSinceUnixEpoch, SlimConnectorWithPermSync from common.data_source.models import BasicExpertInfo, Document, ExternalAccess, GenerateDocumentsOutput, GenerateSlimDocumentOutput, SlimDocument, TextSection @@ -168,6 +174,13 @@ class GmailConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): raise RuntimeError("Creds missing, should not call this property before calling load_credentials") return self._creds + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "GmailConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, str] | None: """Load Gmail credentials.""" primary_admin_email = credentials[DB_CREDENTIALS_PRIMARY_ADMIN_KEY] @@ -179,6 +192,28 @@ class GmailConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): ) return new_creds_dict + def validate_connector_settings(self) -> None: + if self._creds is None: + raise ConnectorMissingCredentialError("Gmail credentials not loaded.") + + try: + gmail_service = get_gmail_service(self.creds, self.primary_admin_email) + gmail_service.users().labels().list( + userId=self.primary_admin_email, + maxResults=1, + ).execute() + except HttpError as e: + status_code = e.resp.status if e.resp else None + if status_code == 401: + raise CredentialExpiredError("Invalid or expired Gmail credentials (401).") from e + if status_code == 403: + raise InsufficientPermissionsError("Gmail app lacks required permissions (403).") from e + raise ConnectorValidationError(f"Unexpected Gmail error (status={status_code}): {e}") from e + except Exception as e: + if MISSING_SCOPES_ERROR_STR in str(e): + raise InsufficientPermissionsError("Gmail credentials are missing required scopes.") from e + raise ConnectorValidationError(f"Unexpected error during Gmail validation: {e}") from e + def _get_all_user_emails(self) -> list[str]: """Get all user emails for Google Workspace domain.""" try: diff --git a/common/data_source/google_drive/connector.py b/common/data_source/google_drive/connector.py index 6fdae09bc5..298e846950 100644 --- a/common/data_source/google_drive/connector.py +++ b/common/data_source/google_drive/connector.py @@ -192,6 +192,23 @@ class GoogleDriveConnector(SlimConnectorWithPermSync, CheckpointedConnectorWithP raise RuntimeError("Creds missing, should not call this property before calling load_credentials") return self._creds + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "GoogleDriveConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + include_shared_drives=config.get("include_shared_drives", False), + include_my_drives=config.get("include_my_drives", False), + include_files_shared_with_me=config.get("include_files_shared_with_me", False), + shared_drive_urls=config.get("shared_drive_urls"), + my_drive_emails=config.get("my_drive_emails"), + shared_folder_urls=config.get("shared_folder_urls"), + specific_user_emails=config.get("specific_user_emails"), + batch_size=batch_size, + ) + connector.set_allow_images(config.get("allow_images", False)) + connector.load_credentials(config.get("credentials") or {}) + return connector + # TODO: ensure returned new_creds_dict is actually persisted when this is called? def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: try: diff --git a/common/data_source/imap_connector.py b/common/data_source/imap_connector.py index 1d3560c2c5..f6e38315dd 100644 --- a/common/data_source/imap_connector.py +++ b/common/data_source/imap_connector.py @@ -23,6 +23,7 @@ from common.data_source.interfaces import ( CheckpointedConnectorWithPermSync, CredentialsConnector, CredentialsProviderInterface, + StaticCredentialsProvider, ) from common.data_source.models import ( BasicExpertInfo, @@ -291,6 +292,22 @@ class ImapConnector( # impls for BaseConnector + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "ImapConnector": + connector = cls( + host=config.get("imap_host"), + port=config.get("imap_port"), + mailboxes=config.get("imap_mailbox"), + ) + connector.set_credentials_provider( + StaticCredentialsProvider( + tenant_id=None, + connector_name=DocumentSource.IMAP, + credential_json=config.get("credentials") or {}, + ) + ) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: self._credentials = credentials return None diff --git a/common/data_source/interfaces.py b/common/data_source/interfaces.py index 5c103f0603..18c5614cc7 100644 --- a/common/data_source/interfaces.py +++ b/common/data_source/interfaces.py @@ -45,6 +45,10 @@ class LoadConnector(ABC): """Validate connector settings""" pass + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "LoadConnector": + raise NotImplementedError(f"{cls.__name__} must implement build_connector") + class PollConnector(ABC): """Poll connector interface""" @@ -239,6 +243,10 @@ class BaseConnector(abc.ABC, Generic[CT]): Default is a no-op (always successful). """ + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "BaseConnector": + raise NotImplementedError(f"{cls.__name__} must implement build_connector") + def validate_perm_sync(self) -> None: """ Permission-sync validation hook. diff --git a/common/data_source/jira/connector.py b/common/data_source/jira/connector.py index c562cf15e0..9f3e452019 100644 --- a/common/data_source/jira/connector.py +++ b/common/data_source/jira/connector.py @@ -135,6 +135,26 @@ class JiraConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPermSync # Connector lifecycle helpers # ------------------------------------------------------------------------- + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "JiraConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + jira_base_url=config["base_url"], + project_key=config.get("project_key"), + jql_query=config.get("jql_query"), + batch_size=batch_size, + include_comments=config.get("include_comments", True), + include_attachments=config.get("include_attachments", False), + labels_to_skip=config.get("labels_to_skip"), + comment_email_blacklist=config.get("comment_email_blacklist"), + scoped_token=config.get("scoped_token", False), + attachment_size_limit=config.get("attachment_size_limit"), + timezone_offset=config.get("timezone_offset"), + time_buffer_seconds=config.get("time_buffer_seconds"), + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Instantiate the Jira client using either an API token or username/password.""" jira_url_for_client = self.jira_base_url diff --git a/common/data_source/moodle_connector.py b/common/data_source/moodle_connector.py index 192155e03a..aef473f381 100644 --- a/common/data_source/moodle_connector.py +++ b/common/data_source/moodle_connector.py @@ -67,6 +67,13 @@ class MoodleConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): for batch in batch_generator(generator, self.batch_size): yield batch + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "MoodleConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(moodle_url=config["moodle_url"], batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> None: token = credentials.get("moodle_token") if not token: diff --git a/common/data_source/notion_connector.py b/common/data_source/notion_connector.py index 2972809294..7e9bdadb88 100644 --- a/common/data_source/notion_connector.py +++ b/common/data_source/notion_connector.py @@ -666,6 +666,12 @@ class NotionConnector(LoadConnector, PollConnector): if slim_batch: yield slim_batch + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "NotionConnector": + connector = cls(root_page_id=config["root_page_id"]) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Applies integration token to headers.""" self.headers["Authorization"] = f"Bearer {credentials['notion_integration_token']}" diff --git a/common/data_source/onedrive_connector.py b/common/data_source/onedrive_connector.py index 0d2a614595..fdb2f336f2 100644 --- a/common/data_source/onedrive_connector.py +++ b/common/data_source/onedrive_connector.py @@ -80,6 +80,13 @@ class OneDriveConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPerm self._access_token: str | None = None self._tenant_id: str | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "OneDriveConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(batch_size=batch_size, folder_path=config.get("folder_path") or None) + connector.load_credentials(config.get("credentials") or {}) + return connector + # ------------------------------------------------------------------ # Auth # ------------------------------------------------------------------ diff --git a/common/data_source/outlook_connector.py b/common/data_source/outlook_connector.py index 08aa210f52..6ed2abaef3 100644 --- a/common/data_source/outlook_connector.py +++ b/common/data_source/outlook_connector.py @@ -120,6 +120,20 @@ class OutlookConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPermS self._access_token: str | None = None self._tenant_id: str | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "OutlookConnector": + user_ids = config.get("user_ids") + if isinstance(user_ids, str): + user_ids = [item.strip() for item in user_ids.split(",") if item.strip()] + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + batch_size=batch_size, + folder=config.get("folder") or _DEFAULT_FOLDER, + user_ids=user_ids, + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + # ------------------------------------------------------------------ # Auth # ------------------------------------------------------------------ diff --git a/common/data_source/rdbms_connector.py b/common/data_source/rdbms_connector.py index 33551fb429..4614d95875 100644 --- a/common/data_source/rdbms_connector.py +++ b/common/data_source/rdbms_connector.py @@ -130,6 +130,25 @@ class RDBMSConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): excluded = {self.id_column, self.timestamp_column} return [col for col in row_dict.keys() if col not in excluded] + @classmethod + def build_connector(cls, config: Dict[str, Any], *, db_type: str) -> "RDBMSConnector": + default_port = 3306 if db_type == DatabaseType.MYSQL else 5432 + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + db_type=db_type, + host=config.get("host", "localhost"), + port=int(config.get("port") or default_port), + database=config.get("database", ""), + query=config.get("query", ""), + content_columns=config.get("content_columns", ""), + metadata_columns=config.get("metadata_columns", ""), + id_column=config.get("id_column") or None, + timestamp_column=config.get("timestamp_column") or None, + batch_size=batch_size, + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: Dict[str, Any]) -> Dict[str, Any] | None: """Load database credentials.""" logging.debug(f"Loading credentials for {self.db_type} database: {self.database}") diff --git a/common/data_source/rest_api_connector.py b/common/data_source/rest_api_connector.py index 91d905089f..68279f205d 100644 --- a/common/data_source/rest_api_connector.py +++ b/common/data_source/rest_api_connector.py @@ -428,6 +428,20 @@ class RestAPIConnector(LoadConnector, PollConnector): return cfg + @classmethod + def build_connector(cls, config: Dict[str, Any]) -> "RestAPIConnector": + cfg = cls.parse_storage_config(config) + connector = cls.from_parsed_config(cfg, max_pages=min(cfg.max_pages, 10)) + connector.load_credentials(config.get("credentials") or {}) + return connector + + def validate_connector_settings(self) -> None: + try: + logging.info("Validating REST API connector by fetching first page") + _ = next(self._page_iter_for_validation()) + except StopIteration: + pass + # -- LoadConnector / PollConnector interface ----------------------------- def load_from_state(self) -> Generator[List[Document], None, None]: diff --git a/common/data_source/rss_connector.py b/common/data_source/rss_connector.py index 6fad756d73..fed9eee10d 100644 --- a/common/data_source/rss_connector.py +++ b/common/data_source/rss_connector.py @@ -34,6 +34,13 @@ class RSSConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): self.credentials = credentials or {} return None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "RSSConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(feed_url=config["feed_url"], batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def validate_connector_settings(self) -> None: self._validate_feed_url() if self.batch_size < 1: diff --git a/common/data_source/salesforce_connector.py b/common/data_source/salesforce_connector.py index 5e0c71f1c5..a9af3eb72c 100644 --- a/common/data_source/salesforce_connector.py +++ b/common/data_source/salesforce_connector.py @@ -124,6 +124,20 @@ class SalesforceConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPe self._instance_url: str | None = None self._access_token: str | None = None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "SalesforceConnector": + objects = config.get("objects") + if isinstance(objects, str): + objects = [item.strip() for item in objects.split(",") if item.strip()] + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + batch_size=batch_size, + objects=objects, + api_version=config.get("api_version") or _DEFAULT_API_VERSION, + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + # ------------------------------------------------------------------ # Auth # ------------------------------------------------------------------ diff --git a/common/data_source/seafile_connector.py b/common/data_source/seafile_connector.py index c9ee59e083..90fa402898 100644 --- a/common/data_source/seafile_connector.py +++ b/common/data_source/seafile_connector.py @@ -167,6 +167,20 @@ class SeaFileConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): ) return resp + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "SeaFileConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + seafile_url=config["seafile_url"], + batch_size=batch_size, + include_shared=config.get("include_shared", True), + sync_scope=config.get("sync_scope", SeafileSyncScope.ACCOUNT), + repo_id=config.get("repo_id") or None, + sync_path=config.get("sync_path") or None, + ) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: logger.debug("Loading credentials for SeaFile server %s", self.seafile_url) diff --git a/common/data_source/sharepoint_connector.py b/common/data_source/sharepoint_connector.py index 519c41c37f..67c6203cab 100644 --- a/common/data_source/sharepoint_connector.py +++ b/common/data_source/sharepoint_connector.py @@ -49,6 +49,13 @@ class SharePointConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPe # -- credentials --------------------------------------------------------- + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "SharePointConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls(batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Configure a Microsoft Graph client from app-only credentials. diff --git a/common/data_source/slack_connector.py b/common/data_source/slack_connector.py index fb2e235491..cdb653c71c 100644 --- a/common/data_source/slack_connector.py +++ b/common/data_source/slack_connector.py @@ -14,9 +14,9 @@ from slack_sdk.errors import SlackApiError from slack_sdk.http_retry import ConnectionErrorRetryHandler from slack_sdk.http_retry.builtin_interval_calculators import FixedValueRetryIntervalCalculator -from common.data_source.config import INDEX_BATCH_SIZE, SLACK_NUM_THREADS, ENABLE_EXPENSIVE_EXPERT_CALLS, _SLACK_LIMIT, FAST_TIMEOUT, MAX_RETRIES, MAX_CHANNELS_TO_LOG +from common.data_source.config import DocumentSource, INDEX_BATCH_SIZE, SLACK_NUM_THREADS, ENABLE_EXPENSIVE_EXPERT_CALLS, _SLACK_LIMIT, FAST_TIMEOUT, MAX_RETRIES, MAX_CHANNELS_TO_LOG from common.data_source.exceptions import ConnectorMissingCredentialError, ConnectorValidationError, CredentialExpiredError, InsufficientPermissionsError, UnexpectedValidationError -from common.data_source.interfaces import CheckpointedConnectorWithPermSync, CredentialsConnector, SlimConnectorWithPermSync +from common.data_source.interfaces import CheckpointedConnectorWithPermSync, CredentialsConnector, SlimConnectorWithPermSync, StaticCredentialsProvider from common.data_source.models import ( BasicExpertInfo, ConnectorCheckpoint, @@ -450,6 +450,24 @@ class SlackConnector( def channels(self, channels: list[str] | None) -> None: self._channels = [channel.removeprefix("#") for channel in channels] if channels else None + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "SlackConnector": + channels = config.get("channels") + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + channels=channels.split(",") if isinstance(channels, str) and channels else channels, + channel_regex_enabled=bool(config.get("channel_regex_enabled", False)), + batch_size=batch_size, + ) + connector.set_credentials_provider( + StaticCredentialsProvider( + tenant_id=None, + connector_name=DocumentSource.SLACK, + credential_json=config.get("credentials") or {}, + ) + ) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Load credentials""" raise NotImplementedError("Use set_credentials_provider with this connector.") diff --git a/common/data_source/teams_connector.py b/common/data_source/teams_connector.py index 2ca4604217..c31e03d31a 100644 --- a/common/data_source/teams_connector.py +++ b/common/data_source/teams_connector.py @@ -56,6 +56,13 @@ class TeamsConnector(CheckpointedConnectorWithPermSync, SlimConnectorWithPermSyn # -- credentials --------------------------------------------------------- + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "TeamsConnector": + batch_size = int(config.get("batch_size") or _SLIM_DOC_BATCH_SIZE) + connector = cls(batch_size=batch_size) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Configure a Microsoft Graph client from app-only credentials. diff --git a/common/data_source/webdav_connector.py b/common/data_source/webdav_connector.py index 4ba6bd3372..c205a6b3aa 100644 --- a/common/data_source/webdav_connector.py +++ b/common/data_source/webdav_connector.py @@ -105,6 +105,18 @@ class WebDAVConnector(LoadConnector, PollConnector, SlimConnectorWithPermSync): logging.info(f"Setting allow_images to {allow_images}.") self._allow_images = allow_images + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "WebDAVConnector": + batch_size = int(config.get("batch_size") or INDEX_BATCH_SIZE) + connector = cls( + base_url=config["base_url"], + remote_path=config.get("remote_path", "/"), + batch_size=batch_size, + ) + connector.set_allow_images(config.get("allow_images", False)) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: """Load credentials and initialize WebDAV client diff --git a/common/data_source/zendesk_connector.py b/common/data_source/zendesk_connector.py index c740e994b9..57e5e41ec4 100644 --- a/common/data_source/zendesk_connector.py +++ b/common/data_source/zendesk_connector.py @@ -318,6 +318,12 @@ class ZendeskConnector(SlimConnectorWithPermSync, CheckpointedConnector[ZendeskC self.content_tags: dict[str, str] = {} self.calls_per_minute = calls_per_minute + @classmethod + def build_connector(cls, config: dict[str, Any]) -> "ZendeskConnector": + connector = cls(content_type=config.get("zendesk_content_type")) + connector.load_credentials(config.get("credentials") or {}) + return connector + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: # Subdomain is actually the whole URL subdomain = credentials["zendesk_subdomain"].replace("https://", "").split(".zendesk.com")[0] diff --git a/test/testcases/restful_api/test_connector_routes_unit.py b/test/testcases/restful_api/test_connector_routes_unit.py index 4797e66cdd..7af3d0a5ad 100644 --- a/test/testcases/restful_api/test_connector_routes_unit.py +++ b/test/testcases/restful_api/test_connector_routes_unit.py @@ -18,6 +18,7 @@ import asyncio import importlib.util import json import sys +from enum import StrEnum from pathlib import Path from types import ModuleType, SimpleNamespace @@ -257,6 +258,12 @@ def _load_connector_app(monkeypatch): SCHEDULE="schedule", CANCEL="cancel", ) + + class _FileSource(StrEnum): + LOCAL = "" + RSS = "rss" + + constants_mod.FileSource = _FileSource monkeypatch.setitem(sys.modules, "common.constants", constants_mod) config_mod = ModuleType("common.data_source.config") @@ -276,6 +283,26 @@ def _load_connector_app(monkeypatch): } monkeypatch.setitem(sys.modules, "common.data_source.google_util.constant", google_constants_mod) + data_source_mod = ModuleType("common.data_source") + + def _stub_build_connector_for_source(_source, _config): + raise NotImplementedError("patch build_connector_for_source in test") + + data_source_mod.build_connector_for_source = _stub_build_connector_for_source + monkeypatch.setitem(sys.modules, "common.data_source", data_source_mod) + + data_source_exceptions_mod = ModuleType("common.data_source.exceptions") + + class _ConnectorMissingCredentialError(Exception): + pass + + class _ConnectorValidationError(Exception): + pass + + data_source_exceptions_mod.ConnectorMissingCredentialError = _ConnectorMissingCredentialError + data_source_exceptions_mod.ConnectorValidationError = _ConnectorValidationError + monkeypatch.setitem(sys.modules, "common.data_source.exceptions", data_source_exceptions_mod) + misc_mod = ModuleType("common.misc_utils") misc_mod.get_uuid = lambda: "uuid-from-helper" monkeypatch.setitem(sys.modules, "common.misc_utils", misc_mod) @@ -455,6 +482,33 @@ def test_connector_by_id_routes_reject_cross_tenant_access(monkeypatch): assert all(res["data"] is False for res in responses) assert touched == [] + class _FakeConnector: + def validate_connector_settings(self): + return None + + monkeypatch.setattr( + sys.modules["common.data_source"], + "build_connector_for_source", + lambda source, config: _FakeConnector(), + ) + monkeypatch.setattr(module.ConnectorService, "get_by_id", lambda _connector_id: (False, None)) + + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({"source": "rss", "config": {"feed_url": "https://example.com"}}), + ) + ok_res = _run(module.test_connector("rss")) + assert ok_res["data"] is True + + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({"source": "rss", "config": "bad"}), + ) + bad_config_res = _run(module.test_connector("rss")) + assert bad_config_res["code"] == module.RetCode.ARGUMENT_ERROR + @pytest.mark.p2 def test_connector_oauth_helper_functions(monkeypatch): diff --git a/test/testcases/test_web_api/test_connector_app/test_connector_routes_unit.py b/test/testcases/test_web_api/test_connector_app/test_connector_routes_unit.py index 6f485af334..413c9771a8 100644 --- a/test/testcases/test_web_api/test_connector_app/test_connector_routes_unit.py +++ b/test/testcases/test_web_api/test_connector_app/test_connector_routes_unit.py @@ -18,6 +18,7 @@ import asyncio import importlib.util import json import sys +from enum import StrEnum from pathlib import Path from types import ModuleType, SimpleNamespace @@ -253,6 +254,12 @@ def _load_connector_app(monkeypatch): AUTHENTICATION_ERROR=109, ) constants_mod.TaskStatus = SimpleNamespace(SCHEDULE="schedule", CANCEL="cancel") + + class _FileSource(StrEnum): + LOCAL = "" + RSS = "rss" + + constants_mod.FileSource = _FileSource monkeypatch.setitem(sys.modules, "common.constants", constants_mod) config_mod = ModuleType("common.data_source.config") @@ -272,6 +279,26 @@ def _load_connector_app(monkeypatch): } monkeypatch.setitem(sys.modules, "common.data_source.google_util.constant", google_constants_mod) + data_source_mod = ModuleType("common.data_source") + + def _stub_build_connector_for_source(_source, _config): + raise NotImplementedError("patch build_connector_for_source in test") + + data_source_mod.build_connector_for_source = _stub_build_connector_for_source + monkeypatch.setitem(sys.modules, "common.data_source", data_source_mod) + + data_source_exceptions_mod = ModuleType("common.data_source.exceptions") + + class _ConnectorMissingCredentialError(Exception): + pass + + class _ConnectorValidationError(Exception): + pass + + data_source_exceptions_mod.ConnectorMissingCredentialError = _ConnectorMissingCredentialError + data_source_exceptions_mod.ConnectorValidationError = _ConnectorValidationError + monkeypatch.setitem(sys.modules, "common.data_source.exceptions", data_source_exceptions_mod) + misc_mod = ModuleType("common.misc_utils") misc_mod.get_uuid = lambda: "uuid-from-helper" monkeypatch.setitem(sys.modules, "common.misc_utils", misc_mod) @@ -450,6 +477,33 @@ def test_connector_by_id_routes_reject_cross_tenant_access(monkeypatch): assert all(res["data"] is False for res in responses) assert touched == [] + class _FakeConnector: + def validate_connector_settings(self): + return None + + monkeypatch.setattr( + sys.modules["common.data_source"], + "build_connector_for_source", + lambda source, config: _FakeConnector(), + ) + monkeypatch.setattr(module.ConnectorService, "get_by_id", lambda _connector_id: (False, None)) + + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({"source": "rss", "config": {"feed_url": "https://example.com"}}), + ) + ok_res = _run(module.test_connector("rss")) + assert ok_res["data"] is True + + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({"source": "rss", "config": "bad"}), + ) + bad_config_res = _run(module.test_connector("rss")) + assert bad_config_res["code"] == module.RetCode.ARGUMENT_ERROR + @pytest.mark.p2 def test_connector_oauth_helper_functions(monkeypatch): diff --git a/web/src/components/dynamic-form.tsx b/web/src/components/dynamic-form.tsx index d82599117f..c42510bd85 100644 --- a/web/src/components/dynamic-form.tsx +++ b/web/src/components/dynamic-form.tsx @@ -157,6 +157,7 @@ export interface DynamicFormRef { submit: () => void; isDirty: () => boolean; getValues: (name?: string) => any; + getFilteredValues: () => any; reset: (values?: any) => void; trigger: UseFormTrigger; watch: (field: string, callback: (value: any) => void) => () => void; @@ -861,6 +862,7 @@ const DynamicForm = { }, isDirty: () => form.formState.isDirty, getValues: form.getValues, + getFilteredValues: () => filterActiveValues(form.getValues()), reset: (values?: T) => { if (values) { form.reset(values); diff --git a/web/src/locales/en.ts b/web/src/locales/en.ts index 1e0d62ab4c..931e1ddc5d 100644 --- a/web/src/locales/en.ts +++ b/web/src/locales/en.ts @@ -1800,6 +1800,10 @@ Example: Virtual Hosted Style`, restApiTestSuccess: 'REST API connector validated successfully.', restApiTestFailed: 'REST API connector validation failed. Please check your configuration and logs.', + dataSourceTestConnection: 'Test connection', + dataSourceTestSuccess: 'Data source connection validated successfully.', + dataSourceTestFailed: + 'Data source connection validation failed. Please check your configuration and logs.', availableSourcesDescription: 'Select a data source to add', availableSources: 'Available sources', datasourceDescription: 'Manage your data source and connections', diff --git a/web/src/locales/zh.ts b/web/src/locales/zh.ts index f3f35597e6..0ab839d53c 100644 --- a/web/src/locales/zh.ts +++ b/web/src/locales/zh.ts @@ -1496,6 +1496,9 @@ NER:使用 spaCy NER 和基于规则的关键词提取来抽取 Entities 和 R availableSourcesDescription: '选择要添加的数据源', availableSources: '可用数据源', datasourceDescription: '管理您的数据源和连接', + dataSourceTestConnection: '测试连接', + dataSourceTestSuccess: '数据源连接验证成功。', + dataSourceTestFailed: '数据源连接验证失败,请检查配置和日志。', chatChannels: '聊天渠道', chatChannelsDescription: '管理您的聊天渠道机器人及凭证', channelEmptyTip: '暂未添加任何聊天渠道,请从下方选择一个进行连接。', diff --git a/web/src/pages/user-setting/data-source/add-datasource-modal.tsx b/web/src/pages/user-setting/data-source/add-datasource-modal.tsx index 4aeda8a70f..7ba3536978 100644 --- a/web/src/pages/user-setting/data-source/add-datasource-modal.tsx +++ b/web/src/pages/user-setting/data-source/add-datasource-modal.tsx @@ -14,10 +14,15 @@ * limitations under the License. */ -import { DynamicForm, FormFieldConfig } from '@/components/dynamic-form'; +import { + DynamicForm, + DynamicFormRef, + FormFieldConfig, +} from '@/components/dynamic-form'; +import { Button } from '@/components/ui/button'; import { Modal } from '@/components/ui/modal/modal'; import { IModalProps } from '@/interfaces/common'; -import { useMemo } from 'react'; +import { useMemo, useRef } from 'react'; import { FieldValues } from 'react-hook-form'; import { useTranslation } from 'react-i18next'; import { @@ -27,6 +32,7 @@ import { getDataSourceFieldsWithExtras, mergeDataSourceFormValues, } from './constant'; +import { useTestDataSource } from './hooks'; import { IDataSorceInfo } from './interface'; const AddDataSourceModal = ({ @@ -37,6 +43,9 @@ const AddDataSourceModal = ({ onOk, }: IModalProps & { sourceData?: IDataSorceInfo }) => { const { t } = useTranslation(); + const formRef = useRef(null); + const { loading: testLoading, handleTest } = useTestDataSource(formRef); + const fields = useMemo(() => { if (!sourceData) { return []; @@ -80,6 +89,7 @@ const AddDataSourceModal = ({ footer={
} > { console.log(data); @@ -93,6 +103,15 @@ const AddDataSourceModal = ({ hideModal?.(); }} /> + { const { t } = useTranslation(); const formRef = useRef(null); + const [searchParams] = useSearchParams(); + const connectorId = searchParams.get('id')!; const { data: detail } = useFetchDataSourceDetail(); const { updateStatus, loading: statusUpdateLoading } = @@ -158,7 +160,10 @@ const SourceDetailPage = () => { }, [detail]); const { addLoading, handleAddOk } = useAddDataSource({ isEdit: true }); - const { loading: testLoading, handleTest } = useTestDataSource(); + const { loading: testLoading, handleTest } = useTestDataSource( + formRef, + connectorId, + ); const onSubmit = useCallback(() => { formRef?.current?.submit(); @@ -248,18 +253,15 @@ const SourceDetailPage = () => { />
- {(detail?.source === DataSourceKey.REST_API || - detail?.source === DataSourceKey.BIGQUERY) && ( - - )} +