Skip to content
Open
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
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion ibm_db_sa/ibm_db.py
Original file line number Diff line number Diff line change
Expand Up @@ -297,7 +297,15 @@ def create_connect_args(
@log_entry_exit
def _get_default_schema_name(self, connection):
logger.debug("Fetching current schema from DB2")
schema = connection.connection.get_current_schema()
# Ask the server. ibm_db_dbi's get_current_schema() falls back to the
# user argument passed to connect(), which is empty when credentials
# are supplied in the DSN, so it can return '' for a valid session.
if hasattr(connection, "exec_driver_sql"):
result = connection.exec_driver_sql("VALUES CURRENT SCHEMA")
else: # SQLAlchemy < 1.4
result = connection.execute("VALUES CURRENT SCHEMA")
schema = result.scalar()
schema = schema.strip() if schema else schema
logger.debug("Current schema returned: %s", schema)
normalized_schema_name = self.normalize_name(schema)
logger.debug("Normalized schema: %s", normalized_schema_name)
Expand Down
63 changes: 63 additions & 0 deletions test/test_default_schema.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
"""Default schema detection for the ibm_db dialect."""

from sqlalchemy import Column, Integer, Table, inspect
from sqlalchemy.testing import fixtures
from sqlalchemy.testing.assertions import eq_

from ibm_db_sa.ibm_db import DB2Dialect_ibm_db


class _Result:
def __init__(self, value):
self.value = value

def scalar(self):
return self.value


class _DBAPIConnection:
# ibm_db_dbi returns the user passed to connect(), which is '' when the
# credentials are part of the DSN.
def get_current_schema(self):
return ""


class _Connection:
connection = _DBAPIConnection()

def __init__(self):
self.statements = []

def exec_driver_sql(self, statement):
self.statements.append(statement)
return _Result("DB2INST1 ")


class TestDefaultSchemaName(fixtures.TestBase):
def test_uses_server_current_schema(self):
connection = _Connection()
eq_(DB2Dialect_ibm_db()._get_default_schema_name(connection), "db2inst1")
eq_(connection.statements, ["VALUES CURRENT SCHEMA"])


class TestDefaultSchemaReflection(fixtures.TestBase):
__only_on__ = "ibm_db_sa+ibm_db_sa"
__backend__ = True

def test_default_schema_matches_server(self, connection):
current = connection.exec_driver_sql("VALUES CURRENT SCHEMA").scalar()
expected = connection.dialect.normalize_name(current.strip())
eq_(connection.dialect.default_schema_name, expected)
eq_(inspect(connection).default_schema_name, expected)

def test_reflection_without_schema(self, metadata, connection):
Table(
"default_schema_reflect",
metadata,
Column("id", Integer, primary_key=True, autoincrement=False),
Column("amount", Integer),
).create(connection)
inspector = inspect(connection)
eq_("default_schema_reflect" in inspector.get_table_names(), True)
columns = inspector.get_columns("default_schema_reflect")
eq_([c["name"] for c in columns], ["id", "amount"])