diff --git a/CHANGELOG.md b/CHANGELOG.md index 55ff47ba..040c0f55 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +- **Databricks Connector (Unity Catalog + Delta Lake ingestion)** (#747) by @KaifAhmad1 + - Added `DatabricksIngestor` (`semantica/ingest/databricks_ingestor.py`), mirroring `SnowflakeIngestor`'s structure and public API shape: a `DatabricksConnector` connection handler, a `DatabricksData` dataclass, and an optional-import guard for `databricks-sdk`/`databricks-sql-connector` + - Supports personal access token and OAuth M2M (service principal `client_id`/`client_secret`) authentication, configurable via constructor args or `DATABRICKS_*` environment variables + - `ingest_table()`/`ingest_query()` run against a SQL warehouse or cluster via `databricks-sql-connector`, with `where`/`order_by`/`limit`/`offset` support and the same identifier-escaping and unsafe-`ORDER BY` rejection as `SnowflakeIngestor`; each call closes the SQL connection it opened unless one is already open (e.g. via the `with DatabricksIngestor(...)` context manager), which reuses and closes it exactly once instead of leaking a second connection per call + - `get_table_schema()`, `list_catalogs()`, `list_schemas()`, and `list_tables()` introspect Unity Catalog via `databricks-sdk`'s `WorkspaceClient`, validating both catalog and schema are resolved before calling the SDK; `get_table_lineage()` calls Unity Catalog's table-lineage REST API for upstream/downstream `Table --DEPENDS_ON--> Table` dependencies, plus an opt-in `include_column_lineage=True` that resolves per-column lineage via the column-lineage API + - `export_as_documents()` converts ingested rows into Semantica document dicts for KG construction, matching `SnowflakeIngestor.export_as_documents()`'s shape + - Registered as a lazy export in `semantica.ingest` (`DatabricksIngestor`, `DatabricksData`, `DatabricksConnector`) and as the `db-databricks` optional extra (`pip install "semantica[db-databricks]"`) in `pyproject.toml`, included in `db-all` + - New `docs/integrations/databricks.md` page modeled on `docs/integrations/snowflake.md`, plus a `DatabricksIngestor` section and table row in `docs/reference/ingest.md` and cross-links between the two integration pages + - 35 unit tests in `tests/test_databricks_ingestor.py` covering both auth methods, table/query ingestion, connection lifecycle (including reuse under the context manager), pagination, unsafe `ORDER BY` rejection, catalog/schema validation, schema/catalog/table listing, table and column lineage, document export, and the missing-dependency error path, closing #747 + - **SQLite Vector Store Backend (`sqlite-vec`)** (#726) by @Luffy2208 and @KaifAhmad1 - Added `SQLiteVecStore` (`semantica/vector_store/sqlite_vec_store.py`), a disk-backed local vector store using the `sqlite-vec` extension's `vec0` virtual tables, closing #240 - Supports Cosine and L2 distance metrics, dynamic JSON metadata filtering, read-only mode, and an in-memory (`:memory:`) mode diff --git a/docs/docs.json b/docs/docs.json index 95d11a5c..f52cdbd6 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -103,7 +103,8 @@ "pages": [ "integrations/agno", "integrations/docling", - "integrations/snowflake" + "integrations/snowflake", + "integrations/databricks" ] }, { diff --git a/docs/integrations/databricks.md b/docs/integrations/databricks.md new file mode 100644 index 00000000..f3ef7b65 --- /dev/null +++ b/docs/integrations/databricks.md @@ -0,0 +1,199 @@ +--- +title: "Databricks Integration" +description: "Ingest Unity Catalog metadata and Delta Lake tables from Databricks into Semantica's KG pipeline." +icon: "cloud" +--- + +> Extract Delta Lake tables and Unity Catalog metadata (schemas, lineage) from Databricks into Semantica with personal access token or OAuth M2M authentication. + + +## Installation + +```bash +# Install with Databricks support +pip install "semantica[db-databricks]" + +# Or install the connectors separately +pip install databricks-sdk databricks-sql-connector +``` + + +## Basic Usage + +```python +from semantica.ingest import DatabricksIngestor +import os + +ingestor = DatabricksIngestor( + host=os.getenv("DATABRICKS_HOST"), # e.g. https://adb-xxx.azuredatabricks.net + token=os.getenv("DATABRICKS_TOKEN"), + http_path=os.getenv("DATABRICKS_HTTP_PATH"), # SQL warehouse or cluster HTTP path + catalog=os.getenv("DATABRICKS_CATALOG", "main"), + schema=os.getenv("DATABRICKS_SCHEMA", "default"), +) + +data = ingestor.ingest_table("customers") +print(f"Retrieved {data.row_count} rows: columns: {data.columns}") +``` + + +Use environment variables (or a `.env` file with `python-dotenv`) to keep credentials out of source code. `DatabricksIngestor()` with no arguments reads from `DATABRICKS_*` environment variables automatically. + + + +## Authentication Methods + + + + ```python + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="dapi-xxxxxxxx", + http_path="/sql/1.0/warehouses/xxxxxxxx", + ) + ``` + + + ```python + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + client_id="your_service_principal_client_id", + client_secret="your_service_principal_client_secret", + http_path="/sql/1.0/warehouses/xxxxxxxx", + ) + ``` + Preferred for production: no long-lived personal token stored in config. + + + + +`http_path` identifies the SQL warehouse or all-purpose cluster used for query execution. Find it in the Databricks UI under **SQL Warehouses → Connection details**. Unity Catalog metadata calls (`list_catalogs`, `get_table_schema`, `get_table_lineage`, …) only need `host` and credentials — `http_path` is not required for those. + + + +## Querying + +### Ingest a table with filters + +```python +data = ingestor.ingest_table( + "customers", + catalog="main", + schema="default", + where="country = 'USA' AND created_date > '2024-01-01'", + order_by="created_date DESC", + limit=10000, +) +``` + +### Custom SQL + +```python +data = ingestor.ingest_query(""" + SELECT customer_id, SUM(amount) AS total_amount + FROM main.default.sales + WHERE date >= '2024-01-01' + GROUP BY customer_id +""") +``` + + +## Unity Catalog Metadata + +### Schema introspection + +```python +schema = ingestor.get_table_schema("customers") +for column in schema["columns"]: + print(f"{column['name']}: {column['type']}") +``` + +### Catalogs, schemas, and tables + +```python +catalogs = ingestor.list_catalogs() +schemas = ingestor.list_schemas(catalog="main") +tables = ingestor.list_tables(catalog="main", schema="default") +``` + +### Table and column lineage + +```python +lineage = ingestor.get_table_lineage("customers", catalog="main", schema="default") +print(lineage["upstream"]) # tables that feed into `customers` +print(lineage["downstream"]) # tables derived from `customers` +``` + +Use `get_table_lineage` to build `Table --DEPENDS_ON--> Table` edges in the knowledge graph directly from Unity Catalog's lineage tracking, without re-deriving lineage from query logs. + + +Pass `include_column_lineage=True` to also resolve per-column upstream/downstream references (one extra Unity Catalog request per column, so it's opt-in): + +```python +lineage = ingestor.get_table_lineage( + "customers", catalog="main", schema="default", include_column_lineage=True, +) +print(lineage["columns"]["email"]) +# {"upstream": ["main.default.raw_customers.email_address"], "downstream": []} +``` + + + + +## Export as Semantica Documents + +```python +documents = ingestor.export_as_documents( + data, + id_field="customer_id", + text_fields=["name", "email", "notes"], +) +print(f"Created {len(documents)} documents for processing") +``` + + +## Batch Processing Large Tables + +```python +PAGE_SIZE = 5000 +for page in range(total_pages): + data = ingestor.ingest_table( + "large_table", + limit=PAGE_SIZE, + offset=page * PAGE_SIZE, + ) + process_batch(data) +``` + +Or use the built-in `batch_size` parameter: + +```python +data = ingestor.ingest_query( + "SELECT * FROM main.default.large_table", + batch_size=5000, +) +``` + + +## Troubleshooting + +```python +from semantica.ingest import DatabricksConnector + +connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="dapi-xxxxxxxx", + http_path="/sql/1.0/warehouses/xxxxxxxx", +) +if not connector.test_connection(): + print("Connection failed: check host, http_path, and credentials") +``` + + +## See Also + +- [Ingest Module](../reference/ingest) — Full DatabricksIngestor and all other ingestors. +- [Snowflake Integration](snowflake) — Companion connector for a Snowflake + Databricks hybrid estate. +- [Pipeline](../reference/pipeline) — Use Databricks ingestion as a pipeline step. +- [Installation](../installation) — All optional dependency extras. +- [Knowledge Graph](../reference/kg) — Build a KG from ingested Databricks data. diff --git a/docs/integrations/snowflake.md b/docs/integrations/snowflake.md index b5cac3c8..bb3b94f3 100644 --- a/docs/integrations/snowflake.md +++ b/docs/integrations/snowflake.md @@ -172,6 +172,7 @@ if not connector.test_connection(): ## See Also - [Ingest Module](../reference/ingest) — Full SnowflakeIngestor and all other ingestors. +- [Databricks Integration](databricks) — Companion connector for a Snowflake + Databricks hybrid estate. - [Pipeline](../reference/pipeline) — Use Snowflake ingestion as a pipeline step. - [Installation](../installation) — All optional dependency extras. - [Knowledge Graph](../reference/kg) — Build a KG from ingested Snowflake data. diff --git a/docs/reference/ingest.md b/docs/reference/ingest.md index fdc3d092..fdd62df3 100644 --- a/docs/reference/ingest.md +++ b/docs/reference/ingest.md @@ -27,6 +27,7 @@ icon: "database" | `RepoIngestor` | Git repositories: source files, commit history, and metadata | | `DBIngestor` | SQL databases via SQLAlchemy: tables, views, and custom queries | | `SnowflakeIngestor` | Snowflake data warehouse queries and table exports | +| `DatabricksIngestor` | Databricks Unity Catalog metadata, Delta table queries, and lineage | | `ParquetIngestor` | Apache Parquet files and partitioned datasets with column selection | | `XMLIngestor` | XXE-safe XML parsing with optional XSD schema validation | | `EmailIngestor` | IMAP/POP3 email ingestion with attachment extraction | @@ -440,6 +441,24 @@ result = ingest("ontology.ttl") # -> {"ontology": OntologyData} result = ingestor.ingest_query("SELECT * FROM documents") result = ingestor.ingest_table("documents") ``` + + ### DatabricksIngestor + + ```python + from semantica.ingest import DatabricksIngestor + import os + + ingestor = DatabricksIngestor( + host=os.getenv("DATABRICKS_HOST"), + token=os.getenv("DATABRICKS_TOKEN"), + http_path=os.getenv("DATABRICKS_HTTP_PATH"), + catalog="main", + schema="default", + ) + result = ingestor.ingest_query("SELECT * FROM documents") + result = ingestor.ingest_table("documents") + lineage = ingestor.get_table_lineage("documents") + ``` ### StreamIngestor @@ -628,4 +647,5 @@ result = ingest_file("source_path", method="my_format") - [Parse](parse) — Parse raw sources into structured text and tables. - [Pipeline](pipeline) — Orchestrate ingest as the first pipeline step. - [Snowflake Integration](../integrations/snowflake) — Snowflake-specific setup and authentication guide. +- [Databricks Integration](../integrations/databricks) — Databricks Unity Catalog setup, authentication, and lineage guide. - [Provenance](provenance) — Track lineage from ingest through to inference. diff --git a/pyproject.toml b/pyproject.toml index 036bcc04..25900c24 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -121,12 +121,13 @@ shacl = ["pyshacl>=0.25.0"] # ---- Database Connectors ---- db-snowflake = ["snowflake-connector-python>=4.6.0", "cryptography>=49.0.0"] +db-databricks = ["databricks-sdk>=0.60.0", "databricks-sql-connector>=4.0.0"] db-arrow = ["pyarrow>=24.0.0"] ingest-parquet = ["pyarrow>=24.0.0"] ingest-arrow = ["pyarrow>=24.0.0"] db-all = [ - "semantica[db-snowflake,db-arrow]" + "semantica[db-snowflake,db-databricks,db-arrow]" ] # ---- Embedding / Models ---- diff --git a/semantica/ingest/__init__.py b/semantica/ingest/__init__.py index ef823ba9..dd63ff6f 100644 --- a/semantica/ingest/__init__.py +++ b/semantica/ingest/__init__.py @@ -218,6 +218,10 @@ _LAZY_EXPORTS: Dict[str, Tuple[str, str]] = { "SnowflakeIngestor": (".snowflake_ingestor", "SnowflakeIngestor"), "SnowflakeData": (".snowflake_ingestor", "SnowflakeData"), "SnowflakeConnector": (".snowflake_ingestor", "SnowflakeConnector"), + # Databricks ingestion + "DatabricksIngestor": (".databricks_ingestor", "DatabricksIngestor"), + "DatabricksData": (".databricks_ingestor", "DatabricksData"), + "DatabricksConnector": (".databricks_ingestor", "DatabricksConnector"), # Parquet ingestion "ParquetIngestor": (".parquet_ingestor", "ParquetIngestor"), "ParquetData": (".parquet_ingestor", "ParquetData"), @@ -341,6 +345,10 @@ __all__ = [ "SnowflakeIngestor", "SnowflakeData", "SnowflakeConnector", + # Databricks ingestion + "DatabricksIngestor", + "DatabricksData", + "DatabricksConnector", # Parquet ingestion "ParquetIngestor", "ParquetData", diff --git a/semantica/ingest/databricks_ingestor.py b/semantica/ingest/databricks_ingestor.py new file mode 100644 index 00000000..3a82ed33 --- /dev/null +++ b/semantica/ingest/databricks_ingestor.py @@ -0,0 +1,1070 @@ +""" +Databricks Ingestion Module + +This module provides comprehensive Databricks ingestion capabilities for the +Semantica framework, enabling data extraction from Databricks lakehouses +(Unity Catalog + Delta Lake) into the knowledge graph pipeline. + +Key Features: + - Native Databricks SQL connection using databricks-sql-connector + - Unity Catalog metadata via databricks-sdk (catalogs, schemas, tables, columns) + - Multiple authentication methods (personal access token, OAuth M2M) + - Query execution with streaming and pagination + - Table and schema introspection + - Table-level and (optional) column-level lineage via Unity Catalog's lineage API + - Progress tracking and error handling + - Connection management with proper cleanup + +Main Classes: + - DatabricksIngestor: Main Databricks ingestion class + - DatabricksConnector: Databricks connection handler + - DatabricksData: Data representation for Databricks ingestion + +Example Usage: + >>> from semantica.ingest import DatabricksIngestor + >>> ingestor = DatabricksIngestor( + ... host="https://adb-xxx.azuredatabricks.net", + ... token="dapi-xxxxxxxx", + ... http_path="/sql/1.0/warehouses/xxxxxxxx", + ... catalog="main", + ... schema="default", + ... ) + >>> data = ingestor.ingest_table("customers", limit=10000) + >>> query_data = ingestor.ingest_query("SELECT * FROM sales WHERE date > '2024-01-01'") + >>> documents = ingestor.export_as_documents(data) + +Author: Semantica Contributors +License: MIT +""" + +import os +import re +from dataclasses import dataclass, field +from datetime import datetime +from typing import Any, Dict, List, Optional + +from ..utils.exceptions import ProcessingError, ValidationError +from ..utils.logging import get_logger +from ..utils.progress_tracker import get_progress_tracker + +try: + from databricks import sql as databricks_sql + from databricks.sdk import WorkspaceClient + from databricks.sdk.core import Config, oauth_service_principal + + DATABRICKS_AVAILABLE = True +except (ImportError, OSError): + databricks_sql = None + WorkspaceClient = None + Config = None + oauth_service_principal = None + DATABRICKS_AVAILABLE = False + + +@dataclass +class DatabricksData: + """Databricks data representation.""" + + data: List[Dict[str, Any]] + row_count: int + columns: List[str] + table_name: Optional[str] = None + query: Optional[str] = None + catalog: Optional[str] = None + schema: Optional[str] = None + metadata: Dict[str, Any] = field(default_factory=dict) + ingested_at: datetime = field(default_factory=datetime.now) + + +class DatabricksConnector: + """ + Databricks connection management. + + This class manages connections to a Databricks workspace and provides + support for multiple authentication methods. + + Supported Authentication Methods: + - Personal Access Token (PAT) + - OAuth M2M (service principal client id / secret) + + Example Usage: + >>> connector = DatabricksConnector( + ... host="https://adb-xxx.azuredatabricks.net", + ... token="dapi-xxxxxxxx", + ... http_path="/sql/1.0/warehouses/xxxxxxxx", + ... ) + >>> conn = connector.connect() + >>> connector.disconnect() + """ + + def __init__( + self, + host: Optional[str] = None, + token: Optional[str] = None, + http_path: Optional[str] = None, + catalog: Optional[str] = None, + schema: Optional[str] = None, + client_id: Optional[str] = None, + client_secret: Optional[str] = None, + **config, + ): + """ + Initialize Databricks connector. + + Args: + host: Databricks workspace URL (e.g. 'https://adb-xxx.azuredatabricks.net') + token: Personal access token (for PAT authentication) + http_path: HTTP path of a SQL warehouse or cluster + (e.g. '/sql/1.0/warehouses/xxxxxxxx') + catalog: Default Unity Catalog catalog to use + schema: Default schema to use + client_id: OAuth M2M service principal client ID + client_secret: OAuth M2M service principal client secret + **config: Additional Databricks connection configuration + """ + if not DATABRICKS_AVAILABLE: + raise ImportError( + "databricks-sdk and databricks-sql-connector are required for " + "DatabricksConnector. Install them with: " + "pip install databricks-sdk databricks-sql-connector" + ) + + self.logger = get_logger("databricks_connector") + + # Get configuration from environment variables if not provided + self.host = host or os.getenv("DATABRICKS_HOST") + self.token = token or os.getenv("DATABRICKS_TOKEN") + self.http_path = http_path or os.getenv("DATABRICKS_HTTP_PATH") + self.catalog = catalog or os.getenv("DATABRICKS_CATALOG", "main") + self.schema = schema or os.getenv("DATABRICKS_SCHEMA", "default") + self.client_id = client_id or os.getenv("DATABRICKS_CLIENT_ID") + self.client_secret = client_secret or os.getenv("DATABRICKS_CLIENT_SECRET") + + # Validate required parameters + if not self.host: + raise ValidationError( + "Databricks host is required. " + "Provide via 'host' parameter or DATABRICKS_HOST environment variable." + ) + + if not (self.client_id and self.client_secret) and not self.token: + raise ValidationError( + "Databricks authentication is required. " + "Provide either 'token' (or DATABRICKS_TOKEN) for personal access " + "token authentication, or 'client_id'/'client_secret' (or " + "DATABRICKS_CLIENT_ID/DATABRICKS_CLIENT_SECRET) for OAuth M2M " + "authentication." + ) + + self.config = config + self.connection = None + self._workspace_client: Optional["WorkspaceClient"] = None + + self.logger.debug( + f"Databricks connector initialized: host={self.host}, " + f"catalog={self.catalog}, schema={self.schema}" + ) + + def connect(self): + """ + Establish a Databricks SQL connection for table/query ingestion. + + If a connection is already open (e.g. because the ingestor is being + used as a context manager), it is reused instead of opening a second + one, which would otherwise be left unclosed. + + Authentication: + - Personal access token: pass ``access_token`` to ``sql.connect``. + - OAuth M2M (service principal): ``sql.connect`` does **not** accept + ``client_id``/``client_secret`` directly. The correct mechanism is + ``credentials_provider``, a callable that returns a header-factory + produced by ``databricks.sdk.core.oauth_service_principal``. + + Returns: + Connection: databricks-sql-connector connection object + + Raises: + ProcessingError: If connection fails + ValidationError: If 'http_path' is not configured + """ + if self.connection is not None: + return self.connection + + if not self.http_path: + raise ValidationError( + "Databricks 'http_path' is required for SQL ingestion. " + "Provide via 'http_path' parameter or DATABRICKS_HTTP_PATH " + "environment variable (the HTTP path of a SQL warehouse or cluster)." + ) + + try: + conn_params: Dict[str, Any] = { + "server_hostname": self._hostname(), + "http_path": self.http_path, + } + + if self.client_id and self.client_secret: + # databricks-sql-connector ≥2.5 requires OAuth M2M to be wired + # through a credentials_provider callable; passing client_id / + # client_secret as plain kwargs is silently ignored and causes + # the connector to fall back to an interactive browser flow. + _client_id = self.client_id + _client_secret = self.client_secret + _host = self.host + + def _m2m_credentials_provider(): + cfg = Config( + host=_host, + client_id=_client_id, + client_secret=_client_secret, + ) + return oauth_service_principal(cfg) + + conn_params["credentials_provider"] = _m2m_credentials_provider + else: + conn_params["access_token"] = self.token + + conn_params.update(self.config) + + self.connection = databricks_sql.connect(**conn_params) + + self.logger.info(f"Connected to Databricks: {self.host}") + + return self.connection + + except Exception as e: + self.logger.error(f"Failed to connect to Databricks: {e}") + raise ProcessingError(f"Failed to connect to Databricks: {e}") from e + + def _hostname(self) -> str: + """Strip the scheme from the configured host for the SQL connector.""" + return re.sub(r"^https?://", "", self.host).rstrip("/") + + def get_workspace_client(self) -> "WorkspaceClient": + """ + Get (and lazily create) a Unity Catalog WorkspaceClient. + + Returns: + WorkspaceClient: databricks-sdk workspace client + + Raises: + ProcessingError: If client creation fails + """ + if self._workspace_client is not None: + return self._workspace_client + + try: + if self.client_id and self.client_secret: + self._workspace_client = WorkspaceClient( + host=self.host, + client_id=self.client_id, + client_secret=self.client_secret, + ) + else: + self._workspace_client = WorkspaceClient( + host=self.host, token=self.token + ) + return self._workspace_client + except Exception as e: + self.logger.error(f"Failed to create Databricks workspace client: {e}") + raise ProcessingError( + f"Failed to create Databricks workspace client: {e}" + ) from e + + def disconnect(self): + """Close Databricks SQL connection.""" + if self.connection: + try: + self.connection.close() + self.connection = None + self.logger.info("Disconnected from Databricks") + except Exception as e: + self.logger.warning(f"Error during disconnect: {e}") + + def test_connection(self) -> bool: + """ + Test the Databricks SQL connection without keeping it open. + + Returns: + bool: True if connection successful, False otherwise + """ + try: + conn = self.connect() + cursor = conn.cursor() + cursor.execute("SELECT 1") + cursor.fetchone() + cursor.close() + conn.close() + self.connection = None + return True + except Exception as e: + self.logger.debug(f"Connection test failed: {e}") + return False + + +class DatabricksIngestor: + """ + Databricks ingestion handler. + + This class provides comprehensive Databricks ingestion capabilities, + connecting to Databricks SQL warehouses/clusters for table and query + ingestion, and to Unity Catalog for metadata and lineage introspection. + + Features: + - Table ingestion with pagination + - Query execution with streaming + - Unity Catalog schema introspection + - Table-level lineage + - Multiple authentication methods + - Progress tracking and error handling + - Connection management with proper cleanup + + Example Usage: + >>> ingestor = DatabricksIngestor( + ... host="https://adb-xxx.azuredatabricks.net", + ... token="dapi-xxxxxxxx", + ... http_path="/sql/1.0/warehouses/xxxxxxxx", + ... catalog="main", + ... ) + >>> data = ingestor.ingest_table("customers", limit=10000) + >>> query_data = ingestor.ingest_query("SELECT * FROM sales") + """ + + def __init__( + self, + host: Optional[str] = None, + token: Optional[str] = None, + http_path: Optional[str] = None, + catalog: Optional[str] = None, + schema: Optional[str] = None, + client_id: Optional[str] = None, + client_secret: Optional[str] = None, + config: Optional[Dict[str, Any]] = None, + **kwargs, + ): + """ + Initialize Databricks ingestor. + + Args: + host: Databricks workspace URL + token: Personal access token + http_path: HTTP path of a SQL warehouse or cluster + catalog: Default catalog (default: 'main') + schema: Default schema (default: 'default') + client_id: OAuth M2M service principal client ID + client_secret: OAuth M2M service principal client secret + config: Optional configuration dictionary + **kwargs: Additional configuration parameters + """ + self.logger = get_logger("databricks_ingestor") + self.config = config or {} + self.config.update(kwargs) + + # Initialize connector + self.connector = DatabricksConnector( + host=host, + token=token, + http_path=http_path, + catalog=catalog, + schema=schema, + client_id=client_id, + client_secret=client_secret, + **self.config, + ) + + # Initialize progress tracker + self.progress_tracker = get_progress_tracker() + # Ensure progress tracker is enabled + if not self.progress_tracker.enabled: + self.progress_tracker.enabled = True + + self.logger.debug("Databricks ingestor initialized") + + def _escape_identifier(self, identifier: str) -> str: + """Escape a SQL identifier by wrapping in backticks and doubling internal backticks. + + Args: + identifier: The identifier to escape + + Returns: + Properly escaped identifier safe for SQL interpolation + """ + escaped = identifier.replace("`", "``") + return f"`{escaped}`" + + def _full_table_name( + self, table_name: str, catalog: Optional[str], schema: Optional[str] + ) -> str: + """Build a fully-qualified 'catalog.schema.table' identifier, omitting + any component (catalog and/or schema) that wasn't provided.""" + parts = [part for part in (catalog, schema, table_name) if part] + return ".".join(parts) + + def _escaped_table_ref( + self, table_name: str, catalog: Optional[str], schema: Optional[str] + ) -> str: + """Build a fully-qualified, identifier-escaped table reference for SQL, + omitting any component (catalog and/or schema) that wasn't provided.""" + parts = [ + self._escape_identifier(part) + for part in (catalog, schema, table_name) + if part + ] + return ".".join(parts) + + def ingest_table( + self, + table_name: str, + catalog: Optional[str] = None, + schema: Optional[str] = None, + limit: Optional[int] = None, + offset: Optional[int] = None, + where: Optional[str] = None, + order_by: Optional[str] = None, + **options, + ) -> DatabricksData: + """ + Ingest data from a Databricks (Delta) table. + + This method retrieves data from a Unity Catalog table with optional + filtering, pagination, and ordering. + + Args: + table_name: Name of the table to ingest + catalog: Catalog name (uses default if not provided) + schema: Schema name (uses default if not provided) + limit: Maximum number of rows to retrieve (optional) + offset: Row offset for pagination (optional) + where: WHERE clause for filtering (optional, must be trusted SQL) + order_by: ORDER BY clause for sorting (optional, must be trusted SQL) + **options: Additional query options + + Warning: + The 'where' and 'order_by' parameters accept raw SQL and must be + trusted input from the caller. Do not pass untrusted user input. + + Returns: + DatabricksData: Ingested data object containing: + - data: List of row dictionaries + - row_count: Number of rows retrieved + - columns: List of column names + - table_name: Table name + - catalog: Catalog name + - schema: Schema name + + Raises: + ProcessingError: If table ingestion fails + """ + catalog = catalog or self.connector.catalog + schema = schema or self.connector.schema + + tracking_id = self.progress_tracker.start_tracking( + file=f"{catalog}.{schema}.{table_name}", + module="ingest", + submodule="DatabricksIngestor", + message=f"Table: {catalog}.{schema}.{table_name}", + ) + + try: + already_connected = self.connector.connection is not None + conn = self.connector.connect() + try: + table_ref = self._escaped_table_ref(table_name, catalog, schema) + + query = f"SELECT * FROM {table_ref}" + + if where: + if ";" in where: + raise ValueError( + "Invalid WHERE clause: semicolons not permitted." + ) + query += f" WHERE {where}" + + if order_by: + _SAFE_ORDER_RE = re.compile( + r"^[A-Za-z_][A-Za-z0-9_]*(\s+(ASC|DESC))?" + r"(\s*,\s*[A-Za-z_][A-Za-z0-9_]*(\s+(ASC|DESC))?)*$", + re.IGNORECASE, + ) + if not _SAFE_ORDER_RE.match(order_by.strip()): + raise ValueError(f"Invalid ORDER BY clause: '{order_by}'") + query += f" ORDER BY {order_by}" + + if limit is not None: + query += f" LIMIT {int(limit)}" + + if offset is not None: + query += f" OFFSET {int(offset)}" + + self.logger.debug(f"Executing query: {query}") + + self.progress_tracker.update_tracking( + tracking_id, message="Executing query..." + ) + + cursor = conn.cursor() + cursor.execute(query) + + self.progress_tracker.update_tracking( + tracking_id, message="Fetching results..." + ) + + columns = ( + [desc[0] for desc in cursor.description] + if cursor.description + else [] + ) + rows = cursor.fetchall() + row_dicts = [dict(zip(columns, row)) for row in rows] + + cursor.close() + + data = self._convert_rows(row_dicts) + + self.progress_tracker.stop_tracking( + tracking_id, + status="completed", + message=f"Ingested {len(data)} rows", + ) + + self.logger.info(f"Table ingestion completed: {len(data)} row(s)") + + return DatabricksData( + data=data, + row_count=len(data), + columns=columns, + table_name=table_name, + catalog=catalog, + schema=schema, + metadata={"query": query}, + ) + finally: + if not already_connected: + self.connector.disconnect() + + except Exception as e: + self.progress_tracker.stop_tracking( + tracking_id, status="failed", message=str(e) + ) + self.logger.error(f"Failed to ingest table: {e}") + raise ProcessingError(f"Failed to ingest table: {e}") from e + + def ingest_query( + self, + query: str, + parameters: Optional[Dict[str, Any]] = None, + batch_size: Optional[int] = None, + **options, + ) -> DatabricksData: + """ + Execute a Databricks SQL query and ingest results. + + This method executes a SQL query and returns the results with + optional batching for large result sets. + + Args: + query: SQL query to execute + parameters: Query parameters for parameterized queries (optional) + batch_size: Batch size for result fetching (optional) + **options: Additional query options + + Returns: + DatabricksData: Query results as DatabricksData object + + Raises: + ProcessingError: If query execution fails + + Example: + >>> data = ingestor.ingest_query( + ... "SELECT * FROM sales WHERE date > :date", + ... parameters={"date": "2024-01-01"} + ... ) + """ + tracking_id = self.progress_tracker.start_tracking( + file="query", + module="ingest", + submodule="DatabricksIngestor", + message="Executing query...", + ) + + try: + already_connected = self.connector.connection is not None + conn = self.connector.connect() + try: + self.logger.debug(f"Executing query: {query[:100]}...") + + cursor = conn.cursor() + + if parameters: + cursor.execute(query, parameters) + else: + cursor.execute(query) + + self.progress_tracker.update_tracking( + tracking_id, message="Fetching results..." + ) + + columns = ( + [desc[0] for desc in cursor.description] + if cursor.description + else [] + ) + + if batch_size: + all_rows = [] + while True: + rows = cursor.fetchmany(batch_size) + if not rows: + break + all_rows.extend(rows) + self.progress_tracker.update_tracking( + tracking_id, message=f"Fetched {len(all_rows)} rows..." + ) + else: + all_rows = cursor.fetchall() + + row_dicts = [dict(zip(columns, row)) for row in all_rows] + + cursor.close() + + data = self._convert_rows(row_dicts) + + self.progress_tracker.stop_tracking( + tracking_id, + status="completed", + message=f"Query completed: {len(data)} rows", + ) + + self.logger.info(f"Query execution completed: {len(data)} row(s)") + + return DatabricksData( + data=data, + row_count=len(data), + columns=columns, + query=query, + metadata={"parameters": parameters} if parameters else {}, + ) + finally: + if not already_connected: + self.connector.disconnect() + + except Exception as e: + self.progress_tracker.stop_tracking( + tracking_id, status="failed", message=str(e) + ) + self.logger.error(f"Failed to execute query: {e}") + raise ProcessingError(f"Failed to execute query: {e}") from e + + def get_table_schema( + self, + table_name: str, + catalog: Optional[str] = None, + schema: Optional[str] = None, + ) -> Dict[str, Any]: + """ + Get schema information for a Unity Catalog table. + + Args: + table_name: Name of the table + catalog: Catalog name (uses default if not provided) + schema: Schema name (uses default if not provided) + + Returns: + dict: Table schema information containing: + - columns: List of column dictionaries with name, type, nullable + - primary_keys: List of primary key column names (if any) + + Raises: + ProcessingError: If schema retrieval fails + """ + try: + catalog = catalog or self.connector.catalog + schema = schema or self.connector.schema + + if not catalog: + raise ValidationError( + "Catalog name is required for schema introspection. " + "Provide via 'catalog' parameter or set default catalog in connector." + ) + + if not schema: + raise ValidationError( + "Schema name is required for schema introspection. " + "Provide via 'schema' parameter or set default schema in connector." + ) + + client = self.connector.get_workspace_client() + full_name = self._full_table_name(table_name, catalog, schema) + + table_info = client.tables.get(full_name=full_name) + + column_info = [] + primary_keys: List[str] = [] + for col in getattr(table_info, "columns", None) or []: + column_info.append( + { + "name": col.name, + "type": getattr(col, "type_text", None) + or getattr(col, "type_name", None), + "nullable": getattr(col, "nullable", True), + "comment": getattr(col, "comment", None), + } + ) + + self.logger.debug( + f"Retrieved schema for {full_name}: {len(column_info)} columns" + ) + + return {"columns": column_info, "primary_keys": primary_keys} + + except Exception as e: + self.logger.error(f"Failed to get table schema: {e}") + raise ProcessingError(f"Failed to get table schema: {e}") from e + + def list_catalogs(self) -> List[str]: + """ + List all catalogs visible in Unity Catalog. + + Returns: + list: List of catalog names + + Raises: + ProcessingError: If listing fails + """ + try: + client = self.connector.get_workspace_client() + catalogs = [c.name for c in client.catalogs.list()] + + self.logger.debug(f"Found {len(catalogs)} catalogs") + + return catalogs + + except Exception as e: + self.logger.error(f"Failed to list catalogs: {e}") + raise ProcessingError(f"Failed to list catalogs: {e}") from e + + def list_schemas(self, catalog: Optional[str] = None) -> List[str]: + """ + List all schemas in a catalog. + + Args: + catalog: Catalog name (uses default if not provided) + + Returns: + list: List of schema names + + Raises: + ProcessingError: If listing fails + """ + try: + catalog = catalog or self.connector.catalog + + if not catalog: + raise ValidationError( + "Catalog name is required for listing schemas. " + "Provide via 'catalog' parameter or set default catalog in connector." + ) + + client = self.connector.get_workspace_client() + schemas = [s.name for s in client.schemas.list(catalog_name=catalog)] + + self.logger.debug(f"Found {len(schemas)} schemas in {catalog}") + + return schemas + + except Exception as e: + self.logger.error(f"Failed to list schemas: {e}") + raise ProcessingError(f"Failed to list schemas: {e}") from e + + def list_tables( + self, catalog: Optional[str] = None, schema: Optional[str] = None + ) -> List[str]: + """ + List all tables in a Unity Catalog catalog/schema. + + Args: + catalog: Catalog name (uses default if not provided) + schema: Schema name (uses default if not provided) + + Returns: + list: List of table names + + Raises: + ProcessingError: If listing fails + """ + try: + catalog = catalog or self.connector.catalog + schema = schema or self.connector.schema + + if not catalog: + raise ValidationError( + "Catalog name is required for listing tables. " + "Provide via 'catalog' parameter or set default catalog in connector." + ) + + if not schema: + raise ValidationError( + "Schema name is required for listing tables. " + "Provide via 'schema' parameter or set default schema in connector." + ) + + client = self.connector.get_workspace_client() + tables = [ + t.name + for t in client.tables.list(catalog_name=catalog, schema_name=schema) + ] + + self.logger.debug(f"Found {len(tables)} tables in {catalog}.{schema}") + + return tables + + except Exception as e: + self.logger.error(f"Failed to list tables: {e}") + raise ProcessingError(f"Failed to list tables: {e}") from e + + def get_table_lineage( + self, + table_name: str, + catalog: Optional[str] = None, + schema: Optional[str] = None, + include_column_lineage: bool = False, + ) -> Dict[str, Any]: + """ + Get table-level (and optionally column-level) lineage from Unity Catalog. + + Args: + table_name: Name of the table + catalog: Catalog name (uses default if not provided) + schema: Schema name (uses default if not provided) + include_column_lineage: If True, also fetch per-column upstream/ + downstream lineage via Unity Catalog's column-lineage API. This + issues one additional request per column in the table, so it is + slower than table-level lineage alone and is opt-in. + + Returns: + dict: Lineage information containing: + - upstream: List of fully-qualified upstream table names + - downstream: List of fully-qualified downstream table names + - columns: (only when include_column_lineage=True) dict mapping + each column name to {"upstream": [...], "downstream": [...]} + fully-qualified column references + + Raises: + ProcessingError: If lineage retrieval fails + """ + try: + catalog = catalog or self.connector.catalog + schema = schema or self.connector.schema + + if not catalog: + raise ValidationError( + "Catalog name is required for lineage retrieval. " + "Provide via 'catalog' parameter or set default catalog in connector." + ) + + if not schema: + raise ValidationError( + "Schema name is required for lineage retrieval. " + "Provide via 'schema' parameter or set default schema in connector." + ) + + full_name = self._full_table_name(table_name, catalog, schema) + + client = self.connector.get_workspace_client() + + response = client.api_client.do( + "GET", + "/api/2.0/lineage-tracking/table-lineage", + query={"table_name": full_name, "include_entity_lineage": False}, + ) + + upstream = [ + item.get("tableInfo", {}).get("name") + for item in response.get("upstreams", []) or [] + if item.get("tableInfo") + ] + downstream = [ + item.get("tableInfo", {}).get("name") + for item in response.get("downstreams", []) or [] + if item.get("tableInfo") + ] + + self.logger.debug( + f"Retrieved lineage for {full_name}: " + f"{len(upstream)} upstream, {len(downstream)} downstream" + ) + + result: Dict[str, Any] = {"upstream": upstream, "downstream": downstream} + + if include_column_lineage: + result["columns"] = self._get_column_lineage( + client, full_name, table_name, catalog, schema + ) + + return result + + except Exception as e: + self.logger.error(f"Failed to get table lineage: {e}") + raise ProcessingError(f"Failed to get table lineage: {e}") from e + + def _get_column_lineage( + self, + client: "WorkspaceClient", + full_name: str, + table_name: str, + catalog: str, + schema: str, + ) -> Dict[str, Dict[str, List[str]]]: + """Fetch per-column lineage for every column in a table. + + Unity Catalog's column-lineage API resolves one column at a time, so + this issues one lineage-tracking request per column in the table's + schema. Columns that fail to resolve lineage (no permissions, no + tracked lineage, etc.) are skipped rather than failing the whole call. + """ + schema_info = self.get_table_schema(table_name, catalog=catalog, schema=schema) + columns_lineage: Dict[str, Dict[str, List[str]]] = {} + + for column in schema_info["columns"]: + column_name = column["name"] + try: + response = client.api_client.do( + "GET", + "/api/2.0/lineage-tracking/column-lineage", + query={"table_name": full_name, "column_name": column_name}, + ) + except Exception as e: + self.logger.debug( + f"Skipping column lineage for {full_name}.{column_name}: {e}" + ) + continue + + columns_lineage[column_name] = { + "upstream": [ + ref + for ref in ( + self._format_column_ref(item) + for item in response.get("upstream_cols", []) or [] + ) + if ref + ], + "downstream": [ + ref + for ref in ( + self._format_column_ref(item) + for item in response.get("downstream_cols", []) or [] + ) + if ref + ], + } + + return columns_lineage + + @staticmethod + def _format_column_ref(item: Dict[str, Any]) -> Optional[str]: + """Format a Unity Catalog column-lineage entry as 'catalog.schema.table.column'.""" + table = item.get("table_name") or item.get("tableName") + column = item.get("name") or item.get("column_name") or item.get("columnName") + if not table or not column: + return None + catalog_name = item.get("catalog_name") or item.get("catalogName") + schema_name = item.get("schema_name") or item.get("schemaName") + parts = [part for part in (catalog_name, schema_name, table, column) if part] + return ".".join(parts) + + def export_as_documents( + self, + data: DatabricksData, + id_field: str = "id", + text_fields: Optional[List[str]] = None, + ) -> List[Dict[str, Any]]: + """ + Convert Databricks data to document format for Semantica processing. + + Args: + data: DatabricksData object to convert + id_field: Field to use as document ID (default: 'id') + text_fields: List of fields to combine into document text (optional) + + Returns: + list: List of document dictionaries + + Example: + >>> data = ingestor.ingest_table("customers") + >>> documents = ingestor.export_as_documents(data, text_fields=["name", "notes"]) + """ + documents = [] + + for idx, row in enumerate(data.data): + doc = { + "id": str(row.get(id_field, idx)), + "metadata": { + "source": "databricks", + "table": data.table_name, + "catalog": data.catalog, + "schema": data.schema, + }, + } + + if text_fields: + text_parts = [] + for field_name in text_fields: + if field_name in row and row[field_name] is not None: + text_parts.append(str(row[field_name])) + doc["text"] = " ".join(text_parts) + else: + text_parts = [] + for key, value in row.items(): + if isinstance(value, str): + text_parts.append(value) + doc["text"] = " ".join(text_parts) + + doc["metadata"]["row_data"] = row + + documents.append(doc) + + self.logger.debug(f"Exported {len(documents)} documents") + + return documents + + def close(self): + """Close Databricks connection.""" + self.connector.disconnect() + + def _convert_rows(self, rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Convert Databricks rows to JSON-serializable format. + + Args: + rows: List of row dictionaries from Databricks + + Returns: + list: Converted row dictionaries + """ + converted = [] + + for row in rows: + converted_row = {} + for key, value in row.items(): + if isinstance(value, datetime): + converted_row[key] = value.isoformat() + elif isinstance(value, bytes): + try: + converted_row[key] = value.decode("utf-8") + except UnicodeDecodeError: + converted_row[key] = str(value) + else: + converted_row[key] = value + + converted.append(converted_row) + + return converted + + def __enter__(self): + """Context manager entry.""" + self.connector.connect() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """Context manager exit.""" + self.close() diff --git a/tests/test_databricks_ingestor.py b/tests/test_databricks_ingestor.py new file mode 100644 index 00000000..44473afb --- /dev/null +++ b/tests/test_databricks_ingestor.py @@ -0,0 +1,912 @@ +""" +Unit tests for Databricks Ingestor + +This test module uses mocks to test Databricks ingestion functionality +without requiring a live Databricks workspace. +""" + +import os +from datetime import datetime +from unittest.mock import MagicMock, Mock, patch + +import pytest + +# Test if databricks-sdk / databricks-sql-connector are available +try: + import databricks.sdk # noqa: F401 + from databricks import sql # noqa: F401 + + DATABRICKS_LIBS_AVAILABLE = True +except ImportError: + DATABRICKS_LIBS_AVAILABLE = False + + +@pytest.fixture(autouse=True) +def mock_databricks_if_needed(): + """Mock databricks modules if not installed.""" + if not DATABRICKS_LIBS_AVAILABLE: + with patch.dict( + "sys.modules", + { + "databricks": MagicMock(), + "databricks.sql": MagicMock(), + "databricks.sdk": MagicMock(), + }, + ): + yield + else: + yield + + +@pytest.fixture +def mock_databricks_connection(): + """Create a mock Databricks SQL connection.""" + mock_conn = Mock() + mock_cursor = Mock() + + mock_cursor.execute = Mock() + mock_cursor.fetchall = Mock(return_value=[]) + mock_cursor.fetchone = Mock(return_value=[1]) + mock_cursor.fetchmany = Mock(return_value=[]) + mock_cursor.description = [("id", None), ("name", None), ("value", None)] + mock_cursor.close = Mock() + + mock_conn.cursor = Mock(return_value=mock_cursor) + mock_conn.close = Mock() + + return mock_conn, mock_cursor + + +class TestDatabricksConnector: + """Test DatabricksConnector class.""" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_init_with_token(self, mock_sql): + """Test connector initialization with personal access token authentication.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + ) + + assert connector.host == "https://adb-xxx.azuredatabricks.net" + assert connector.token == "test_token" + assert connector.http_path == "/sql/1.0/warehouses/xxxx" + assert connector.catalog == "TEST_CATALOG" + assert connector.schema == "default" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_init_from_env(self, mock_sql): + """Test connector initialization from environment variables.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + with patch.dict( + os.environ, + { + "DATABRICKS_HOST": "https://env-host.azuredatabricks.net", + "DATABRICKS_TOKEN": "env_token", + "DATABRICKS_HTTP_PATH": "/sql/1.0/warehouses/env", + }, + ): + connector = DatabricksConnector() + + assert connector.host == "https://env-host.azuredatabricks.net" + assert connector.token == "env_token" + assert connector.http_path == "/sql/1.0/warehouses/env" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_init_missing_host(self, mock_sql): + """Test connector initialization fails without host.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + from semantica.utils.exceptions import ValidationError + + with pytest.raises(ValidationError, match="Databricks host is required"): + DatabricksConnector(token="test_token") + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_init_missing_auth(self, mock_sql): + """Test connector initialization fails without any authentication method.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + from semantica.utils.exceptions import ValidationError + + with pytest.raises(ValidationError, match="Databricks authentication is required"): + DatabricksConnector(host="https://adb-xxx.azuredatabricks.net") + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_connect_token_auth(self, mock_sql, mock_databricks_connection): + """Test connection with personal access token authentication.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + conn = connector.connect() + + assert conn == mock_conn + mock_sql.connect.assert_called_once() + + call_kwargs = mock_sql.connect.call_args[1] + assert call_kwargs["server_hostname"] == "adb-xxx.azuredatabricks.net" + assert call_kwargs["http_path"] == "/sql/1.0/warehouses/xxxx" + assert call_kwargs["access_token"] == "test_token" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.oauth_service_principal") + @patch("semantica.ingest.databricks_ingestor.Config") + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_connect_oauth_m2m( + self, mock_sql, mock_config, mock_oauth_sp, mock_databricks_connection + ): + """Test connection with OAuth M2M authentication uses credentials_provider. + + databricks-sql-connector does NOT accept client_id/client_secret as + direct kwargs to sql.connect(); the correct mechanism is a + credentials_provider callable wrapping oauth_service_principal(). + """ + from semantica.ingest.databricks_ingestor import DatabricksConnector + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + mock_oauth_sp.return_value = {"Authorization": "Bearer fake-token"} + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + http_path="/sql/1.0/warehouses/xxxx", + client_id="test_client_id", + client_secret="test_client_secret", + ) + + connector.connect() + + call_kwargs = mock_sql.connect.call_args[1] + + # Must use credentials_provider, not bare client_id/client_secret + assert "credentials_provider" in call_kwargs + assert callable(call_kwargs["credentials_provider"]) + assert "client_id" not in call_kwargs + assert "client_secret" not in call_kwargs + assert "access_token" not in call_kwargs + + # Invoke the provider to verify it wires Config + oauth_service_principal + call_kwargs["credentials_provider"]() + mock_config.assert_called_once_with( + host="https://adb-xxx.azuredatabricks.net", + client_id="test_client_id", + client_secret="test_client_secret", + ) + mock_oauth_sp.assert_called_once_with(mock_config.return_value) + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_connect_missing_http_path(self, mock_sql): + """Test connection fails without http_path.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + from semantica.utils.exceptions import ValidationError + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + ) + + with pytest.raises(ValidationError, match="http_path"): + connector.connect() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_disconnect(self, mock_sql, mock_databricks_connection): + """Test connection disconnect.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + connector.connect() + connector.disconnect() + + mock_conn.close.assert_called_once() + assert connector.connection is None + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_test_connection_success(self, mock_sql, mock_databricks_connection): + """Test successful connection test.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + result = connector.test_connection() + + assert result is True + mock_cursor.execute.assert_called_with("SELECT 1") + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_connector_test_connection_failure(self, mock_sql): + """Test connection test failure.""" + from semantica.ingest.databricks_ingestor import DatabricksConnector + + mock_sql.connect = Mock(side_effect=Exception("Connection failed")) + + connector = DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + result = connector.test_connection() + + assert result is False + + +class TestDatabricksIngestor: + """Test DatabricksIngestor class.""" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingestor_init(self, mock_sql): + """Test ingestor initialization.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + assert ingestor.connector is not None + assert ingestor.connector.host == "https://adb-xxx.azuredatabricks.net" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_basic(self, mock_sql, mock_databricks_connection): + """Test basic table ingestion.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + mock_cursor.fetchall = Mock( + return_value=[ + (1, "Alice", 100), + (2, "Bob", 200), + ] + ) + mock_cursor.description = [("id", None), ("name", None), ("value", None)] + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + schema="TEST_SCHEMA", + ) + + data = ingestor.ingest_table("customers") + + assert data.row_count == 2 + assert data.table_name == "customers" + assert data.catalog == "TEST_CATALOG" + assert data.schema == "TEST_SCHEMA" + assert len(data.columns) == 3 + assert "id" in data.columns + assert data.data[0]["name"] == "Alice" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_closes_connection(self, mock_sql, mock_databricks_connection): + """Test that ingest_table() closes the SQL connection after use instead of leaking it.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + ingestor.ingest_table("customers") + + mock_conn.close.assert_called_once() + assert ingestor.connector.connection is None + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_catalog_only(self, mock_sql, mock_databricks_connection): + """Test table ingestion still qualifies the reference when only catalog is provided.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + # Force a missing schema past the connector's "default" fallback. + ingestor.connector.schema = None + + ingestor.ingest_table("customers", catalog="main") + + executed_query = mock_cursor.execute.call_args[0][0] + assert "`main`.`customers`" in executed_query + assert "None" not in executed_query + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_with_limit(self, mock_sql, mock_databricks_connection): + """Test table ingestion with limit.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + ingestor.ingest_table("customers", limit=100) + + executed_query = mock_cursor.execute.call_args[0][0] + assert "LIMIT 100" in executed_query + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_with_where(self, mock_sql, mock_databricks_connection): + """Test table ingestion with WHERE clause.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + ingestor.ingest_table("customers", where="value > 100") + + executed_query = mock_cursor.execute.call_args[0][0] + assert "WHERE value > 100" in executed_query + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_table_rejects_unsafe_order_by(self, mock_sql, mock_databricks_connection): + """Test table ingestion rejects unsafe ORDER BY clauses.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + from semantica.utils.exceptions import ProcessingError + + with pytest.raises(ProcessingError): + ingestor.ingest_table("customers", order_by="value; DROP TABLE customers") + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_query_basic(self, mock_sql, mock_databricks_connection): + """Test basic query execution.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + mock_cursor.fetchall = Mock(return_value=[(1000,)]) + mock_cursor.description = [("total", None)] + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + query = "SELECT SUM(value) AS total FROM sales" + data = ingestor.ingest_query(query) + + assert data.row_count == 1 + assert data.query == query + assert data.data[0]["total"] == 1000 + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_query_closes_connection(self, mock_sql, mock_databricks_connection): + """Test that ingest_query() closes the SQL connection after use instead of leaking it.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + ingestor.ingest_query("SELECT * FROM sales") + + mock_conn.close.assert_called_once() + assert ingestor.connector.connection is None + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_ingest_query_with_batching(self, mock_sql, mock_databricks_connection): + """Test query execution with batch fetching.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + batch1 = [(1,), (2,)] + batch2 = [(3,)] + mock_cursor.fetchmany = Mock(side_effect=[batch1, batch2, []]) + mock_cursor.description = [("id", None)] + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + data = ingestor.ingest_query("SELECT * FROM customers", batch_size=2) + + assert data.row_count == 3 + assert len(data.data) == 3 + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_get_table_schema(self, mock_ws_client_cls, mock_sql): + """Test getting table schema information.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_column_1 = Mock(name="id") + mock_column_1.name = "id" + mock_column_1.type_text = "BIGINT" + mock_column_1.nullable = False + mock_column_1.comment = None + + mock_column_2 = Mock(name="name") + mock_column_2.name = "name" + mock_column_2.type_text = "STRING" + mock_column_2.nullable = True + mock_column_2.comment = None + + mock_table_info = Mock() + mock_table_info.columns = [mock_column_1, mock_column_2] + + mock_ws_client = Mock() + mock_ws_client.tables.get = Mock(return_value=mock_table_info) + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + schema="TEST_SCHEMA", + ) + + schema = ingestor.get_table_schema("customers") + + assert len(schema["columns"]) == 2 + assert schema["columns"][0]["name"] == "id" + assert schema["columns"][0]["type"] == "BIGINT" + assert schema["columns"][0]["nullable"] is False + mock_ws_client.tables.get.assert_called_once_with( + full_name="TEST_CATALOG.TEST_SCHEMA.customers" + ) + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_get_table_schema_requires_schema(self, mock_ws_client_cls, mock_sql): + """Test that get_table_schema() raises instead of calling the SDK with a missing schema.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + from semantica.utils.exceptions import ProcessingError + + mock_ws_client = Mock() + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + ) + ingestor.connector.schema = None + + with pytest.raises(ProcessingError, match="Schema name is required"): + ingestor.get_table_schema("customers") + + mock_ws_client.tables.get.assert_not_called() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_list_catalogs(self, mock_ws_client_cls, mock_sql): + """Test listing catalogs.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_catalog_1 = Mock() + mock_catalog_1.name = "main" + mock_catalog_2 = Mock() + mock_catalog_2.name = "samples" + + mock_ws_client = Mock() + mock_ws_client.catalogs.list = Mock(return_value=[mock_catalog_1, mock_catalog_2]) + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + catalogs = ingestor.list_catalogs() + + assert len(catalogs) == 2 + assert "main" in catalogs + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_list_tables(self, mock_ws_client_cls, mock_sql): + """Test listing tables in a catalog/schema.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_table_1 = Mock() + mock_table_1.name = "customers" + mock_table_2 = Mock() + mock_table_2.name = "orders" + + mock_ws_client = Mock() + mock_ws_client.tables.list = Mock(return_value=[mock_table_1, mock_table_2]) + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + schema="TEST_SCHEMA", + ) + + tables = ingestor.list_tables() + + assert len(tables) == 2 + assert "customers" in tables + mock_ws_client.tables.list.assert_called_once_with( + catalog_name="TEST_CATALOG", schema_name="TEST_SCHEMA" + ) + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_list_tables_requires_schema(self, mock_ws_client_cls, mock_sql): + """Test that list_tables() raises instead of calling the SDK with schema_name=None.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + from semantica.utils.exceptions import ProcessingError + + mock_ws_client = Mock() + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="TEST_CATALOG", + ) + # Force a missing schema past the connector's "default" fallback. + ingestor.connector.schema = None + + with pytest.raises(ProcessingError, match="Schema name is required"): + ingestor.list_tables() + + mock_ws_client.tables.list.assert_not_called() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_get_table_lineage(self, mock_ws_client_cls, mock_sql): + """Test getting table lineage.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_ws_client = Mock() + mock_ws_client.api_client.do = Mock( + return_value={ + "upstreams": [{"tableInfo": {"name": "main.default.raw_customers"}}], + "downstreams": [{"tableInfo": {"name": "main.default.customer_summary"}}], + } + ) + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="main", + schema="default", + ) + + lineage = ingestor.get_table_lineage("customers") + + assert lineage["upstream"] == ["main.default.raw_customers"] + assert lineage["downstream"] == ["main.default.customer_summary"] + assert "columns" not in lineage + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_get_table_lineage_with_column_lineage(self, mock_ws_client_cls, mock_sql): + """Test that include_column_lineage=True fetches per-column lineage.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_column_1 = Mock() + mock_column_1.name = "id" + mock_column_1.type_text = "BIGINT" + mock_column_1.nullable = False + mock_column_1.comment = None + + mock_column_2 = Mock() + mock_column_2.name = "name" + mock_column_2.type_text = "STRING" + mock_column_2.nullable = True + mock_column_2.comment = None + + mock_table_info = Mock() + mock_table_info.columns = [mock_column_1, mock_column_2] + + def do_side_effect(method, path, query=None): + if path.endswith("table-lineage"): + return { + "upstreams": [{"tableInfo": {"name": "main.default.raw_customers"}}], + "downstreams": [], + } + if path.endswith("column-lineage"): + if query["column_name"] == "id": + return { + "upstream_cols": [ + { + "catalog_name": "main", + "schema_name": "default", + "table_name": "raw_customers", + "name": "customer_id", + } + ], + "downstream_cols": [], + } + return {"upstream_cols": [], "downstream_cols": []} + raise AssertionError(f"unexpected path: {path}") + + mock_ws_client = Mock() + mock_ws_client.tables.get = Mock(return_value=mock_table_info) + mock_ws_client.api_client.do = Mock(side_effect=do_side_effect) + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="main", + schema="default", + ) + + lineage = ingestor.get_table_lineage("customers", include_column_lineage=True) + + assert lineage["upstream"] == ["main.default.raw_customers"] + assert lineage["columns"]["id"]["upstream"] == ["main.default.raw_customers.customer_id"] + assert lineage["columns"]["id"]["downstream"] == [] + assert lineage["columns"]["name"]["upstream"] == [] + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + @patch("semantica.ingest.databricks_ingestor.WorkspaceClient") + def test_get_table_lineage_requires_schema(self, mock_ws_client_cls, mock_sql): + """Test that get_table_lineage() validates schema before calling the REST API.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + from semantica.utils.exceptions import ProcessingError + + mock_ws_client = Mock() + mock_ws_client_cls.return_value = mock_ws_client + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + catalog="main", + ) + ingestor.connector.schema = None + + with pytest.raises(ProcessingError, match="Schema name is required"): + ingestor.get_table_lineage("customers") + + mock_ws_client.api_client.do.assert_not_called() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_export_as_documents(self, mock_sql): + """Test exporting data as documents.""" + from semantica.ingest.databricks_ingestor import DatabricksData, DatabricksIngestor + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + data = DatabricksData( + data=[ + {"id": 1, "name": "Alice", "description": "Engineer"}, + {"id": 2, "name": "Bob", "description": "Designer"}, + ], + row_count=2, + columns=["id", "name", "description"], + table_name="employees", + catalog="main", + schema="default", + ) + + documents = ingestor.export_as_documents( + data, id_field="id", text_fields=["name", "description"] + ) + + assert len(documents) == 2 + assert documents[0]["id"] == "1" + assert documents[0]["text"] == "Alice Engineer" + assert documents[0]["metadata"]["source"] == "databricks" + assert documents[0]["metadata"]["table"] == "employees" + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_context_manager(self, mock_sql, mock_databricks_connection): + """Test using ingestor as context manager.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, _ = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + with DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) as ingestor: + assert ingestor.connector.connection == mock_conn + + mock_conn.close.assert_called() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_context_manager_reuses_connection_across_calls( + self, mock_sql, mock_databricks_connection + ): + """Test that ingest_table()/ingest_query() reuse (not leak) the connection + opened by __enter__, and only close it once on __exit__.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + mock_conn, mock_cursor = mock_databricks_connection + mock_sql.connect = Mock(return_value=mock_conn) + + with DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) as ingestor: + ingestor.ingest_table("customers") + # The connection opened by __enter__ must still be the live one: + # ingest_table() should not have closed it out from under the + # context manager. + assert ingestor.connector.connection == mock_conn + mock_conn.close.assert_not_called() + + ingestor.ingest_query("SELECT * FROM sales") + assert ingestor.connector.connection == mock_conn + mock_conn.close.assert_not_called() + + # databricks_sql.connect() should only have been called once (by + # __enter__): both ingestion calls reused that same connection. + mock_sql.connect.assert_called_once() + mock_conn.close.assert_called_once() + + @patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", True) + @patch("semantica.ingest.databricks_ingestor.databricks_sql") + def test_convert_datetime(self, mock_sql): + """Test datetime conversion in _convert_rows.""" + from semantica.ingest.databricks_ingestor import DatabricksIngestor + + ingestor = DatabricksIngestor( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + http_path="/sql/1.0/warehouses/xxxx", + ) + + test_dt = datetime(2024, 1, 15, 10, 30, 0) + rows = [{"timestamp": test_dt, "value": 100}] + + converted = ingestor._convert_rows(rows) + + assert converted[0]["timestamp"] == "2024-01-15T10:30:00" + assert converted[0]["value"] == 100 + + def test_import_error_without_databricks(self): + """Test that proper error is raised when databricks libraries are not installed.""" + with patch("semantica.ingest.databricks_ingestor.DATABRICKS_AVAILABLE", False): + from semantica.ingest.databricks_ingestor import DatabricksConnector + + with pytest.raises(ImportError, match="databricks-sdk"): + DatabricksConnector( + host="https://adb-xxx.azuredatabricks.net", + token="test_token", + ) + + +class TestDatabricksData: + """Test DatabricksData dataclass.""" + + def test_databricks_data_creation(self): + """Test DatabricksData creation.""" + from semantica.ingest.databricks_ingestor import DatabricksData + + data = DatabricksData( + data=[{"col1": "val1"}], + row_count=1, + columns=["col1"], + table_name="test_table", + ) + + assert data.row_count == 1 + assert data.table_name == "test_table" + assert len(data.data) == 1 + assert isinstance(data.ingested_at, datetime) + + def test_databricks_data_with_metadata(self): + """Test DatabricksData with metadata.""" + from semantica.ingest.databricks_ingestor import DatabricksData + + metadata = {"custom_field": "value"} + data = DatabricksData( + data=[], + row_count=0, + columns=[], + metadata=metadata, + ) + + assert data.metadata["custom_field"] == "value"