Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
The table of contents is too big for display.
Diff view
Diff view
  •  
  •  
  •  
69 changes: 37 additions & 32 deletions sqlspec/adapters/adbc/data_dictionary.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,7 @@ def get_version(self, driver: "AdbcDriver") -> "VersionInfo | None":

try:
version_query_dialect = "mysql" if dialect == "mariadb" else dialect
version_value = driver.select_value_or_none(self._get_query(version_query_dialect, "version"))
version_value = driver.select_value_or_none(self._get_query(version_query_dialect, "version", "current"))
except Exception:
self._log_version_unavailable(dialect, "query_failed")
self.cache_version(driver_id, None)
Expand Down Expand Up @@ -450,18 +450,18 @@ def get_tables(self, driver: "AdbcDriver", schema: "str | None" = None) -> "list

if dialect == "bigquery":
tables_table, kcu_table, rc_table = format_bigquery_information_schema_tables(schema_name)
query_text = self._get_query_text(dialect, "tables_by_schema").format(
query_text = self._get_query_text(dialect, "tables", "by_schema").format(
tables_table=tables_table, kcu_table=kcu_table, rc_table=rc_table
)
return driver.select(query_text, schema_type=TableMetadata)

if dialect == "sqlite":
schema_prefix = f"{format_identifier(schema_name)}." if schema_name else ""
query_text = self._get_query_text(dialect, "tables_by_schema").format(schema_prefix=schema_prefix)
query_text = self._get_query_text(dialect, "tables", "by_schema").format(schema_prefix=schema_prefix)
return driver.select(query_text, schema_type=TableMetadata)

return driver.select(
self._get_query(dialect, "tables_by_schema"), schema_name=schema_name, schema_type=TableMetadata
self._get_query(dialect, "tables", "by_schema"), schema_name=schema_name, schema_type=TableMetadata
)

def get_columns(
Expand All @@ -484,30 +484,30 @@ def get_columns(
if dialect == "bigquery":
schema_prefix = format_bigquery_schema_prefix(schema_name)
if table is None:
query_text = self._get_query_text(dialect, "columns_by_schema").format(schema_prefix=schema_prefix)
query_text = self._get_query_text(dialect, "columns", "by_schema").format(schema_prefix=schema_prefix)
return driver.select(query_text, schema_name=schema_name, schema_type=ColumnMetadata)
query_text = self._get_query_text(dialect, "columns_by_table").format(schema_prefix=schema_prefix)
query_text = self._get_query_text(dialect, "columns", "by_table").format(schema_prefix=schema_prefix)
table_name = self._resolve_identifier(dialect, table)
return driver.select(query_text, table_name=table_name, schema_name=schema_name, schema_type=ColumnMetadata)

if dialect == "sqlite":
schema_prefix = f"{format_identifier(schema_name)}." if schema_name else ""
if table is None:
query_text = self._get_query_text(dialect, "columns_by_schema").format(schema_prefix=schema_prefix)
query_text = self._get_query_text(dialect, "columns", "by_schema").format(schema_prefix=schema_prefix)
return driver.select(query_text, schema_type=ColumnMetadata)
table_identifier = f"{schema_name}.{table}" if schema_name else table
query_text = self._get_query_text(dialect, "columns_by_table").format(
query_text = self._get_query_text(dialect, "columns", "by_table").format(
table_name=format_identifier(table_identifier)
)
return driver.select(query_text, schema_type=ColumnMetadata)

if table is None:
return driver.select(
self._get_query(dialect, "columns_by_schema"), schema_name=schema_name, schema_type=ColumnMetadata
self._get_query(dialect, "columns", "by_schema"), schema_name=schema_name, schema_type=ColumnMetadata
)
table_name = self._resolve_identifier(dialect, table)
return driver.select(
self._get_query(dialect, "columns_by_table"),
self._get_query(dialect, "columns", "by_table"),
schema_name=schema_name,
table_name=table_name,
schema_type=ColumnMetadata,
Expand Down Expand Up @@ -537,7 +537,7 @@ def get_indexes(

table_name = table
table_identifier = f"{schema_name}.{table_name}" if schema_name else table_name
index_list_sql = self._get_query_text(dialect, "indexes_by_table").format(
index_list_sql = self._get_query_text(dialect, "indexes", "by_table").format(
table_name=format_identifier(table_identifier)
)
index_list_rows = driver.select(index_list_sql)
Expand All @@ -547,7 +547,7 @@ def get_indexes(
if not index_name:
continue
index_identifier = f"{schema_name}.{index_name}" if schema_name else index_name
columns_sql = self._get_query_text(dialect, "index_columns_by_index").format(
columns_sql = self._get_query_text(dialect, "indexes", "columns_by_index").format(
index_name=format_identifier(index_identifier)
)
columns_rows = driver.select(columns_sql)
Expand All @@ -567,17 +567,17 @@ def get_indexes(
return index_metadata_list

if dialect == "duckdb":
query_name = "indexes_by_schema" if table is None else "indexes_by_table"
return driver.select(self._get_query(dialect, query_name), schema_type=IndexMetadata)
operation = "by_schema" if table is None else "by_table"
return driver.select(self._get_query(dialect, "indexes", operation), schema_type=IndexMetadata)

if table is None:
return driver.select(
self._get_query(dialect, "indexes_by_schema"), schema_name=schema_name, schema_type=IndexMetadata
self._get_query(dialect, "indexes", "by_schema"), schema_name=schema_name, schema_type=IndexMetadata
)

table_name = self._resolve_identifier(dialect, table)
return driver.select(
self._get_query(dialect, "indexes_by_table"),
self._get_query(dialect, "indexes", "by_table"),
schema_name=schema_name,
table_name=table_name,
schema_type=IndexMetadata,
Expand All @@ -603,11 +603,11 @@ def get_foreign_keys(
if dialect == "bigquery":
_, kcu_table, rc_table = format_bigquery_information_schema_tables(schema_name)
if table is None:
query_text = self._get_query_text(dialect, "foreign_keys_by_schema").format(
query_text = self._get_query_text(dialect, "foreign_keys", "by_schema").format(
kcu_table=kcu_table, rc_table=rc_table
)
return driver.select(query_text, schema_name=schema_name, schema_type=ForeignKeyMetadata)
query_text = self._get_query_text(dialect, "foreign_keys_by_table").format(
query_text = self._get_query_text(dialect, "foreign_keys", "by_table").format(
kcu_table=kcu_table, rc_table=rc_table
)
table_name = self._resolve_identifier(dialect, table)
Expand All @@ -618,23 +618,25 @@ def get_foreign_keys(
if dialect == "sqlite":
if table is None:
schema_prefix = f"{format_identifier(schema_name)}." if schema_name else ""
query_text = self._get_query_text(dialect, "foreign_keys_by_schema").format(schema_prefix=schema_prefix)
query_text = self._get_query_text(dialect, "foreign_keys", "by_schema").format(
schema_prefix=schema_prefix
)
return driver.select(query_text, schema_type=ForeignKeyMetadata)
table_label = table.replace("'", "''")
table_identifier = f"{schema_name}.{table}" if schema_name else table
query_text = self._get_query_text(dialect, "foreign_keys_by_table").format(
query_text = self._get_query_text(dialect, "foreign_keys", "by_table").format(
table_name=format_identifier(table_identifier), table_label=table_label
)
return driver.select(query_text, schema_type=ForeignKeyMetadata)

if table is None:
query_text_optional = self._get_query_text_or_none(dialect, "foreign_keys_by_schema")
query_text_optional = self._get_query_text_or_none(dialect, "foreign_keys", "by_schema")
if query_text_optional is not None:
return driver.select(query_text_optional, schema_name=schema_name, schema_type=ForeignKeyMetadata)

resolved_table_name = self._resolve_identifier(dialect, table) if table is not None else None
return driver.select(
self._get_query(dialect, "foreign_keys_by_table"),
self._get_query(dialect, "foreign_keys", "by_table"),
schema_name=schema_name,
table_name=resolved_table_name,
schema_type=ForeignKeyMetadata,
Expand Down Expand Up @@ -782,17 +784,20 @@ def _normalize_dialect(self, driver: "AdbcDriver") -> str:
dialect_value = str(driver.dialect)
return normalize_dialect_name(dialect_value)

def _get_query(self, dialect: str, name: str) -> "SQL":
def _get_query(self, dialect: str, domain: str, operation: str) -> "SQL":
loader = get_data_dictionary_loader()
return loader.get_query(dialect, name)
query = loader.get_domain_query(dialect, domain, operation)
if query.sql is None:
msg = f"No data-dictionary query found for {dialect}/{domain}/{operation}"
raise SQLFileNotFoundError(msg)
return query.sql

def _get_query_text(self, dialect: str, name: str) -> str:
loader = get_data_dictionary_loader()
return loader.get_query_text(dialect, name)
def _get_query_text(self, dialect: str, domain: str, operation: str) -> str:
return self._get_query(dialect, domain, operation).raw_sql

def _get_query_text_or_none(self, dialect: str, name: str) -> "str | None":
def _get_query_text_or_none(self, dialect: str, domain: str, operation: str) -> "str | None":
try:
return self._get_query_text(dialect, name)
return self._get_query_text(dialect, domain, operation)
except SQLFileNotFoundError:
return None

Expand Down Expand Up @@ -920,7 +925,7 @@ def _capability_for_domain(self, dialect: str, domain: str, probes: "dict[str, b
return self._transport_capability(domain)
if domain in {"schemas", "objects", "tables"} and (probes["objects"] or probes["table_types"]):
return self._transport_capability(domain)
if domain == "indexes" and self._has_query(dialect, "indexes_by_schema"):
if domain == "indexes" and self._has_query(dialect, "indexes", "by_schema"):
return MetadataCapability(
domain=domain,
support=MetadataSupport.SUPPORTED,
Expand All @@ -945,9 +950,9 @@ def _transport_capability(self, domain: str) -> "MetadataCapability":
warnings=(_ADBC_TRANSPORT_WARNING,),
)

def _has_query(self, dialect: str, name: str) -> bool:
def _has_query(self, dialect: str, domain: str, operation: str) -> bool:
try:
self._get_query(dialect, name)
self._get_query(dialect, domain, operation)
except SQLFileNotFoundError:
return False
return True
Expand Down
20 changes: 10 additions & 10 deletions sqlspec/adapters/aiomysql/data_dictionary.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ async def get_version(self, driver: "AiomysqlDriver") -> "VersionInfo | None":
if driver_id in self._version_fetch_attempted:
return self._version_cache.get(driver_id)

version_value = await driver.select_value_or_none(self.get_query("version"))
version_value = await driver.select_value_or_none(self.get_query("version", "current"))
if not version_value:
self._log_version_unavailable(type(self).dialect, "missing")
self.cache_version(driver_id, None)
Expand Down Expand Up @@ -153,7 +153,7 @@ async def get_ddl(
"""Get native SHOW CREATE output and replay-sensitive context for a table."""
_ = include_dependencies, prefer_native, redact
schema_name = self._resolve_metadata_schema(schema)
raw_version = await driver.select_value_or_none(self.get_query_text("version"))
raw_version = await driver.select_value_or_none(self.get_query_text("version", "current"))
sql_mode = await driver.select_value_or_none("SELECT @@sql_mode")
sql_quote_show_create = await driver.select_value_or_none("SELECT @@sql_quote_show_create")
statement = build_mysql_show_create_statement(object_name, schema_name, object_type)
Expand Down Expand Up @@ -210,7 +210,7 @@ async def get_tables(self, driver: "AiomysqlDriver", schema: "str | None" = None
schema_name = self._resolve_metadata_schema(schema)
self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables")
return await driver.select(
self.get_query("tables_by_schema"), schema_name=schema_name, schema_type=TableMetadata
self.get_query("tables", "by_schema"), schema_name=schema_name, schema_type=TableMetadata
)

async def get_columns(
Expand All @@ -221,12 +221,12 @@ async def get_columns(
if table is None:
self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="columns")
return await driver.select(
self.get_query("columns_by_schema"), schema_name=schema_name, schema_type=ColumnMetadata
self.get_query("columns", "by_schema"), schema_name=schema_name, schema_type=ColumnMetadata
)

self._log_table_describe(driver, schema_name=schema_name, table_name=table, operation="columns")
return await driver.select(
self.get_query("columns_by_table"), table_name=table, schema_name=schema_name, schema_type=ColumnMetadata
self.get_query("columns", "by_table"), table_name=table, schema_name=schema_name, schema_type=ColumnMetadata
)

async def get_indexes(
Expand All @@ -237,12 +237,12 @@ async def get_indexes(
if table is None:
self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="indexes")
return await driver.select(
self.get_query("indexes_by_schema"), schema_name=schema_name, schema_type=IndexMetadata
self.get_query("indexes", "by_schema"), schema_name=schema_name, schema_type=IndexMetadata
)

self._log_table_describe(driver, schema_name=schema_name, table_name=table, operation="indexes")
return await driver.select(
self.get_query("indexes_by_table"), table_name=table, schema_name=schema_name, schema_type=IndexMetadata
self.get_query("indexes", "by_table"), table_name=table, schema_name=schema_name, schema_type=IndexMetadata
)

async def get_foreign_keys(
Expand All @@ -253,12 +253,12 @@ async def get_foreign_keys(
if table is None:
self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="foreign_keys")
return await driver.select(
self.get_query("foreign_keys_by_schema"), schema_name=schema_name, schema_type=ForeignKeyMetadata
self.get_query("foreign_keys", "by_schema"), schema_name=schema_name, schema_type=ForeignKeyMetadata
)

self._log_table_describe(driver, schema_name=schema_name, table_name=table, operation="foreign_keys")
return await driver.select(
self.get_query("foreign_keys_by_table"),
self.get_query("foreign_keys", "by_table"),
table_name=table,
schema_name=schema_name,
schema_type=ForeignKeyMetadata,
Expand All @@ -270,7 +270,7 @@ def _resolve_metadata_schema(self, schema: "str | None") -> "str | None":

async def _get_engine_version(self, driver: "AiomysqlDriver") -> "MySQLEngineVersion | None":
try:
version_value = await driver.select_value_or_none(self.get_query_text("version"))
version_value = await driver.select_value_or_none(self.get_query_text("version", "current"))
except Exception:
return None
if version_value is None:
Expand Down
6 changes: 3 additions & 3 deletions sqlspec/adapters/aiosqlite/data_dictionary.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ async def get_version(self, driver: "AiosqliteDriver") -> "VersionInfo | None":
return self._version_cache.get(driver_id)
# Not cached, fetch from database

version_value = await driver.select_value_or_none(self.get_query("version"))
version_value = await driver.select_value_or_none(self.get_query("version", "current"))
if not version_value:
self._log_version_unavailable(type(self).dialect, "missing")
self.cache_version(driver_id, None)
Expand Down Expand Up @@ -158,11 +158,11 @@ async def get_foreign_keys(
schema_name = self.resolve_schema(schema)
if table is None:
self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="foreign_keys")
query_text = self._get_domain_query_text("constraints", "foreign_keys_by_schema")
query_text = self._get_domain_query_text("foreign_keys", "by_schema")
return await driver.select(query_text, schema_name=schema_name, schema_type=ForeignKeyMetadata)

self._log_table_describe(driver, schema_name=schema_name, table_name=table, operation="foreign_keys")
query_text = self._get_domain_query_text("constraints", "foreign_keys_by_table")
query_text = self._get_domain_query_text("foreign_keys", "by_table")
return await driver.select(
query_text, table_name=table, schema_name=schema_name, schema_type=ForeignKeyMetadata
)
Expand Down
Loading
Loading