Apply black
This commit is contained in:
+24
-11
@@ -5,17 +5,24 @@ from dataset.table import Table
|
||||
from dataset.util import row_type
|
||||
|
||||
# shut up useless SA warning:
|
||||
warnings.filterwarnings("ignore", "Unicode type received non-unicode bind param value.")
|
||||
warnings.filterwarnings(
|
||||
'ignore', 'Unicode type received non-unicode bind param value.')
|
||||
warnings.filterwarnings(
|
||||
'ignore', 'Skipping unsupported ALTER for creation of implicit constraint')
|
||||
"ignore", "Skipping unsupported ALTER for creation of implicit constraint"
|
||||
)
|
||||
|
||||
__all__ = ['Database', 'Table', 'freeze', 'connect']
|
||||
__version__ = '1.3.2'
|
||||
__all__ = ["Database", "Table", "freeze", "connect"]
|
||||
__version__ = "1.3.2"
|
||||
|
||||
|
||||
def connect(url=None, schema=None, reflect_metadata=True, engine_kwargs=None,
|
||||
reflect_views=True, ensure_schema=True, row_type=row_type):
|
||||
def connect(
|
||||
url=None,
|
||||
schema=None,
|
||||
reflect_metadata=True,
|
||||
engine_kwargs=None,
|
||||
reflect_views=True,
|
||||
ensure_schema=True,
|
||||
row_type=row_type,
|
||||
):
|
||||
""" Opens a new connection to a database.
|
||||
|
||||
*url* can be any valid `SQLAlchemy engine URL`_. If *url* is not defined
|
||||
@@ -40,8 +47,14 @@ def connect(url=None, schema=None, reflect_metadata=True, engine_kwargs=None,
|
||||
.. _DB connection timeout: http://docs.sqlalchemy.org/en/latest/core/pooling.html#setting-pool-recycle
|
||||
"""
|
||||
if url is None:
|
||||
url = os.environ.get('DATABASE_URL', 'sqlite://')
|
||||
url = os.environ.get("DATABASE_URL", "sqlite://")
|
||||
|
||||
return Database(url, schema=schema, reflect_metadata=reflect_metadata,
|
||||
engine_kwargs=engine_kwargs, reflect_views=reflect_views,
|
||||
ensure_schema=ensure_schema, row_type=row_type)
|
||||
return Database(
|
||||
url,
|
||||
schema=schema,
|
||||
reflect_metadata=reflect_metadata,
|
||||
engine_kwargs=engine_kwargs,
|
||||
reflect_views=reflect_views,
|
||||
ensure_schema=ensure_schema,
|
||||
row_type=row_type,
|
||||
)
|
||||
|
||||
+29
-18
@@ -22,9 +22,16 @@ log = logging.getLogger(__name__)
|
||||
class Database(object):
|
||||
"""A database object represents a SQL database with multiple tables."""
|
||||
|
||||
def __init__(self, url, schema=None, reflect_metadata=True,
|
||||
engine_kwargs=None, reflect_views=True,
|
||||
ensure_schema=True, row_type=row_type):
|
||||
def __init__(
|
||||
self,
|
||||
url,
|
||||
schema=None,
|
||||
reflect_metadata=True,
|
||||
engine_kwargs=None,
|
||||
reflect_views=True,
|
||||
ensure_schema=True,
|
||||
row_type=row_type,
|
||||
):
|
||||
"""Configure and connect to the database."""
|
||||
if engine_kwargs is None:
|
||||
engine_kwargs = {}
|
||||
@@ -41,13 +48,13 @@ class Database(object):
|
||||
if len(parsed_url.query):
|
||||
query = parse_qs(parsed_url.query)
|
||||
if schema is None:
|
||||
schema_qs = query.get('schema', query.get('searchpath', []))
|
||||
schema_qs = query.get("schema", query.get("searchpath", []))
|
||||
if len(schema_qs):
|
||||
schema = schema_qs.pop()
|
||||
|
||||
self.schema = schema
|
||||
self.engine = create_engine(url, **engine_kwargs)
|
||||
self.is_postgres = self.engine.dialect.name == 'postgresql'
|
||||
self.is_postgres = self.engine.dialect.name == "postgresql"
|
||||
self.types = Types(is_postgres=self.is_postgres)
|
||||
self.url = url
|
||||
self.row_type = row_type
|
||||
@@ -57,7 +64,7 @@ class Database(object):
|
||||
@property
|
||||
def executable(self):
|
||||
"""Connection against which statements will be executed."""
|
||||
if not hasattr(self.local, 'conn'):
|
||||
if not hasattr(self.local, "conn"):
|
||||
self.local.conn = self.engine.connect()
|
||||
return self.local.conn
|
||||
|
||||
@@ -80,7 +87,7 @@ class Database(object):
|
||||
@property
|
||||
def in_transaction(self):
|
||||
"""Check if this database is in a transactional context."""
|
||||
if not hasattr(self.local, 'tx'):
|
||||
if not hasattr(self.local, "tx"):
|
||||
return False
|
||||
return len(self.local.tx) > 0
|
||||
|
||||
@@ -94,7 +101,7 @@ class Database(object):
|
||||
|
||||
No data will be written until the transaction has been committed.
|
||||
"""
|
||||
if not hasattr(self.local, 'tx'):
|
||||
if not hasattr(self.local, "tx"):
|
||||
self.local.tx = []
|
||||
self.local.tx.append(self.executable.begin())
|
||||
|
||||
@@ -103,7 +110,7 @@ class Database(object):
|
||||
|
||||
Make all statements executed since the transaction was begun permanent.
|
||||
"""
|
||||
if hasattr(self.local, 'tx') and self.local.tx:
|
||||
if hasattr(self.local, "tx") and self.local.tx:
|
||||
tx = self.local.tx.pop()
|
||||
tx.commit()
|
||||
self._flush_tables()
|
||||
@@ -113,7 +120,7 @@ class Database(object):
|
||||
|
||||
Discard all statements executed since the transaction was begun.
|
||||
"""
|
||||
if hasattr(self.local, 'tx') and self.local.tx:
|
||||
if hasattr(self.local, "tx") and self.local.tx:
|
||||
tx = self.local.tx.pop()
|
||||
tx.rollback()
|
||||
self._flush_tables()
|
||||
@@ -190,15 +197,19 @@ class Database(object):
|
||||
table5 = db.create_table('population5',
|
||||
primary_id=False)
|
||||
"""
|
||||
assert not isinstance(primary_type, str), \
|
||||
'Text-based primary_type support is dropped, use db.types.'
|
||||
assert not isinstance(
|
||||
primary_type, str
|
||||
), "Text-based primary_type support is dropped, use db.types."
|
||||
table_name = normalize_table_name(table_name)
|
||||
with self.lock:
|
||||
if table_name not in self._tables:
|
||||
self._tables[table_name] = Table(self, table_name,
|
||||
primary_id=primary_id,
|
||||
primary_type=primary_type,
|
||||
auto_create=True)
|
||||
self._tables[table_name] = Table(
|
||||
self,
|
||||
table_name,
|
||||
primary_id=primary_id,
|
||||
primary_type=primary_type,
|
||||
auto_create=True,
|
||||
)
|
||||
return self._tables.get(table_name)
|
||||
|
||||
def load_table(self, table_name):
|
||||
@@ -265,7 +276,7 @@ class Database(object):
|
||||
"""
|
||||
if isinstance(query, str):
|
||||
query = text(query)
|
||||
_step = kwargs.pop('_step', QUERY_STEP)
|
||||
_step = kwargs.pop("_step", QUERY_STEP)
|
||||
if _step is False or _step == 0:
|
||||
_step = None
|
||||
rp = self.executable.execute(query, *args, **kwargs)
|
||||
@@ -273,4 +284,4 @@ class Database(object):
|
||||
|
||||
def __repr__(self):
|
||||
"""Text representation contains the URL."""
|
||||
return '<Database(%s)>' % safe_url(self.url)
|
||||
return "<Database(%s)>" % safe_url(self.url)
|
||||
|
||||
+75
-76
@@ -22,20 +22,27 @@ log = logging.getLogger(__name__)
|
||||
|
||||
class Table(object):
|
||||
"""Represents a table in a database and exposes common operations."""
|
||||
PRIMARY_DEFAULT = 'id'
|
||||
|
||||
def __init__(self, database, table_name, primary_id=None,
|
||||
primary_type=None, auto_create=False):
|
||||
PRIMARY_DEFAULT = "id"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
database,
|
||||
table_name,
|
||||
primary_id=None,
|
||||
primary_type=None,
|
||||
auto_create=False,
|
||||
):
|
||||
"""Initialise the table from database schema."""
|
||||
self.db = database
|
||||
self.name = normalize_table_name(table_name)
|
||||
self._table = None
|
||||
self._columns = None
|
||||
self._indexes = []
|
||||
self._primary_id = primary_id if primary_id is not None \
|
||||
else self.PRIMARY_DEFAULT
|
||||
self._primary_type = primary_type if primary_type is not None \
|
||||
else Types.integer
|
||||
self._primary_id = (
|
||||
primary_id if primary_id is not None else self.PRIMARY_DEFAULT
|
||||
)
|
||||
self._primary_type = primary_type if primary_type is not None else Types.integer
|
||||
self._auto_create = auto_create
|
||||
|
||||
@property
|
||||
@@ -206,8 +213,7 @@ class Table(object):
|
||||
if return_count:
|
||||
return self.count(clause)
|
||||
|
||||
def update_many(self, rows, keys, chunk_size=1000, ensure=None,
|
||||
types=None):
|
||||
def update_many(self, rows, keys, chunk_size=1000, ensure=None, types=None):
|
||||
"""Update many rows in the table at a time.
|
||||
|
||||
This is significantly faster than updating them one by one. Per default
|
||||
@@ -229,16 +235,14 @@ class Table(object):
|
||||
|
||||
# bindparam requires names to not conflict (cannot be "id" for id)
|
||||
for key in keys:
|
||||
row['_%s' % key] = row[key]
|
||||
row["_%s" % key] = row[key]
|
||||
|
||||
# Update when chunk_size is fulfilled or this is the last row
|
||||
if len(chunk) == chunk_size or index == len(rows) - 1:
|
||||
cl = [self.table.c[k] == bindparam('_%s' % k) for k in keys]
|
||||
cl = [self.table.c[k] == bindparam("_%s" % k) for k in keys]
|
||||
stmt = self.table.update(
|
||||
whereclause=and_(*cl),
|
||||
values={
|
||||
col: bindparam(col, required=False) for col in columns
|
||||
}
|
||||
values={col: bindparam(col, required=False) for col in columns},
|
||||
)
|
||||
self.db.executable.execute(stmt, chunk)
|
||||
chunk = []
|
||||
@@ -261,8 +265,7 @@ class Table(object):
|
||||
return self.insert(row, ensure=False)
|
||||
return True
|
||||
|
||||
def upsert_many(self, rows, keys, chunk_size=1000, ensure=None,
|
||||
types=None):
|
||||
def upsert_many(self, rows, keys, chunk_size=1000, ensure=None, types=None):
|
||||
"""
|
||||
Sorts multiple input rows into upserts and inserts. Inserts are passed
|
||||
to insert_many and upserts are updated.
|
||||
@@ -311,19 +314,20 @@ class Table(object):
|
||||
with self.db.lock:
|
||||
self._flush_metadata()
|
||||
try:
|
||||
self._table = SQLATable(self.name,
|
||||
self.db.metadata,
|
||||
schema=self.db.schema,
|
||||
autoload=True)
|
||||
self._table = SQLATable(
|
||||
self.name, self.db.metadata, schema=self.db.schema, autoload=True
|
||||
)
|
||||
except NoSuchTableError:
|
||||
self._table = None
|
||||
|
||||
def _threading_warn(self):
|
||||
if self.db.in_transaction and threading.active_count() > 1:
|
||||
warnings.warn("Changing the database schema inside a transaction "
|
||||
"in a multi-threaded environment is likely to lead "
|
||||
"to race conditions and synchronization issues.",
|
||||
RuntimeWarning)
|
||||
warnings.warn(
|
||||
"Changing the database schema inside a transaction "
|
||||
"in a multi-threaded environment is likely to lead "
|
||||
"to race conditions and synchronization issues.",
|
||||
RuntimeWarning,
|
||||
)
|
||||
|
||||
def _sync_table(self, columns):
|
||||
"""Lazy load, create or adapt the table structure in the database."""
|
||||
@@ -338,18 +342,21 @@ class Table(object):
|
||||
# Keep the lock scope small because this is run very often.
|
||||
with self.db.lock:
|
||||
self._threading_warn()
|
||||
self._table = SQLATable(self.name,
|
||||
self.db.metadata,
|
||||
schema=self.db.schema)
|
||||
self._table = SQLATable(
|
||||
self.name, self.db.metadata, schema=self.db.schema
|
||||
)
|
||||
if self._primary_id is not False:
|
||||
# This can go wrong on DBMS like MySQL and SQLite where
|
||||
# tables cannot have no columns.
|
||||
primary_id = self._primary_id
|
||||
primary_type = self._primary_type
|
||||
increment = primary_type in [Types.integer, Types.bigint]
|
||||
column = Column(primary_id, primary_type,
|
||||
primary_key=True,
|
||||
autoincrement=increment)
|
||||
column = Column(
|
||||
primary_id,
|
||||
primary_type,
|
||||
primary_key=True,
|
||||
autoincrement=increment,
|
||||
)
|
||||
self._table.append_column(column)
|
||||
for column in columns:
|
||||
if not column.name == self._primary_id:
|
||||
@@ -361,9 +368,7 @@ class Table(object):
|
||||
self._threading_warn()
|
||||
for column in columns:
|
||||
if not self.has_column(column.name):
|
||||
self.db.op.add_column(self.name,
|
||||
column,
|
||||
self.db.schema)
|
||||
self.db.op.add_column(self.name, column, self.db.schema)
|
||||
self._reflect_table()
|
||||
|
||||
def _sync_columns(self, row, ensure, types=None):
|
||||
@@ -397,31 +402,31 @@ class Table(object):
|
||||
return ensure
|
||||
|
||||
def _generate_clause(self, column, op, value):
|
||||
if op in ('like',):
|
||||
if op in ("like",):
|
||||
return self.table.c[column].like(value)
|
||||
if op in ('ilike',):
|
||||
if op in ("ilike",):
|
||||
return self.table.c[column].ilike(value)
|
||||
if op in ('>', 'gt'):
|
||||
if op in (">", "gt"):
|
||||
return self.table.c[column] > value
|
||||
if op in ('<', 'lt'):
|
||||
if op in ("<", "lt"):
|
||||
return self.table.c[column] < value
|
||||
if op in ('>=', 'gte'):
|
||||
if op in (">=", "gte"):
|
||||
return self.table.c[column] >= value
|
||||
if op in ('<=', 'lte'):
|
||||
if op in ("<=", "lte"):
|
||||
return self.table.c[column] <= value
|
||||
if op in ('=', '==', 'is'):
|
||||
if op in ("=", "==", "is"):
|
||||
return self.table.c[column] == value
|
||||
if op in ('!=', '<>', 'not'):
|
||||
if op in ("!=", "<>", "not"):
|
||||
return self.table.c[column] != value
|
||||
if op in ('in'):
|
||||
if op in ("in"):
|
||||
return self.table.c[column].in_(value)
|
||||
if op in ('between', '..'):
|
||||
if op in ("between", ".."):
|
||||
start, end = value
|
||||
return self.table.c[column].between(start, end)
|
||||
if op in ('startswith',):
|
||||
return self.table.c[column].like('%' + value)
|
||||
if op in ('endswith',):
|
||||
return self.table.c[column].like(value + '%')
|
||||
if op in ("startswith",):
|
||||
return self.table.c[column].like("%" + value)
|
||||
if op in ("endswith",):
|
||||
return self.table.c[column].like(value + "%")
|
||||
return false()
|
||||
|
||||
def _args_to_clause(self, args, clauses=()):
|
||||
@@ -431,12 +436,12 @@ class Table(object):
|
||||
if not self.has_column(column):
|
||||
clauses.append(false())
|
||||
elif isinstance(value, (list, tuple, set)):
|
||||
clauses.append(self._generate_clause(column, 'in', value))
|
||||
clauses.append(self._generate_clause(column, "in", value))
|
||||
elif isinstance(value, dict):
|
||||
for op, op_value in value.items():
|
||||
clauses.append(self._generate_clause(column, op, op_value))
|
||||
else:
|
||||
clauses.append(self._generate_clause(column, '=', value))
|
||||
clauses.append(self._generate_clause(column, "=", value))
|
||||
return and_(*clauses)
|
||||
|
||||
def _args_to_order_by(self, order_by):
|
||||
@@ -444,11 +449,11 @@ class Table(object):
|
||||
for ordering in ensure_list(order_by):
|
||||
if ordering is None:
|
||||
continue
|
||||
column = ordering.lstrip('-')
|
||||
column = ordering.lstrip("-")
|
||||
column = self._get_column_name(column)
|
||||
if not self.has_column(column):
|
||||
continue
|
||||
if ordering.startswith('-'):
|
||||
if ordering.startswith("-"):
|
||||
orderings.append(self.table.c[column].desc())
|
||||
else:
|
||||
orderings.append(self.table.c[column].asc())
|
||||
@@ -501,7 +506,7 @@ class Table(object):
|
||||
::
|
||||
table.drop_column('created_at')
|
||||
"""
|
||||
if self.db.engine.dialect.name == 'sqlite':
|
||||
if self.db.engine.dialect.name == "sqlite":
|
||||
raise RuntimeError("SQLite does not support dropping columns.")
|
||||
name = self._get_column_name(name)
|
||||
with self.db.lock:
|
||||
@@ -510,11 +515,7 @@ class Table(object):
|
||||
return
|
||||
|
||||
self._threading_warn()
|
||||
self.db.op.drop_column(
|
||||
self.table.name,
|
||||
name,
|
||||
self.table.schema
|
||||
)
|
||||
self.db.op.drop_column(self.table.name, name, self.table.schema)
|
||||
self._reflect_table()
|
||||
|
||||
def drop(self):
|
||||
@@ -542,7 +543,7 @@ class Table(object):
|
||||
return False
|
||||
indexes = self.db.inspect.get_indexes(self.name, schema=self.db.schema)
|
||||
for index in indexes:
|
||||
if columns == set(index.get('column_names', [])):
|
||||
if columns == set(index.get("column_names", [])):
|
||||
self._indexes.append(columns)
|
||||
return True
|
||||
return False
|
||||
@@ -600,19 +601,17 @@ class Table(object):
|
||||
if not self.exists:
|
||||
return iter([])
|
||||
|
||||
_limit = kwargs.pop('_limit', None)
|
||||
_offset = kwargs.pop('_offset', 0)
|
||||
order_by = kwargs.pop('order_by', None)
|
||||
_streamed = kwargs.pop('_streamed', False)
|
||||
_step = kwargs.pop('_step', QUERY_STEP)
|
||||
_limit = kwargs.pop("_limit", None)
|
||||
_offset = kwargs.pop("_offset", 0)
|
||||
order_by = kwargs.pop("order_by", None)
|
||||
_streamed = kwargs.pop("_streamed", False)
|
||||
_step = kwargs.pop("_step", QUERY_STEP)
|
||||
if _step is False or _step == 0:
|
||||
_step = None
|
||||
|
||||
order_by = self._args_to_order_by(order_by)
|
||||
args = self._args_to_clause(kwargs, clauses=_clauses)
|
||||
query = self.table.select(whereclause=args,
|
||||
limit=_limit,
|
||||
offset=_offset)
|
||||
query = self.table.select(whereclause=args, limit=_limit, offset=_offset)
|
||||
if len(order_by):
|
||||
query = query.order_by(*order_by)
|
||||
|
||||
@@ -621,9 +620,7 @@ class Table(object):
|
||||
conn = self.db.engine.connect()
|
||||
conn = conn.execution_options(stream_results=True)
|
||||
|
||||
return ResultIter(conn.execute(query),
|
||||
row_type=self.db.row_type,
|
||||
step=_step)
|
||||
return ResultIter(conn.execute(query), row_type=self.db.row_type, step=_step)
|
||||
|
||||
def find_one(self, *args, **kwargs):
|
||||
"""Get a single result from the table.
|
||||
@@ -637,8 +634,8 @@ class Table(object):
|
||||
if not self.exists:
|
||||
return None
|
||||
|
||||
kwargs['_limit'] = 1
|
||||
kwargs['_step'] = None
|
||||
kwargs["_limit"] = 1
|
||||
kwargs["_step"] = None
|
||||
resiter = self.find(*args, **kwargs)
|
||||
try:
|
||||
for row in resiter:
|
||||
@@ -692,10 +689,12 @@ class Table(object):
|
||||
if not len(columns):
|
||||
return iter([])
|
||||
|
||||
q = expression.select(columns,
|
||||
distinct=True,
|
||||
whereclause=clause,
|
||||
order_by=[c.asc() for c in columns])
|
||||
q = expression.select(
|
||||
columns,
|
||||
distinct=True,
|
||||
whereclause=clause,
|
||||
order_by=[c.asc() for c in columns],
|
||||
)
|
||||
return self.db.query(q)
|
||||
|
||||
# Legacy methods for running find queries.
|
||||
@@ -715,4 +714,4 @@ class Table(object):
|
||||
|
||||
def __repr__(self):
|
||||
"""Get table representation."""
|
||||
return '<Table(%s)>' % self.table.name
|
||||
return "<Table(%s)>" % self.table.name
|
||||
|
||||
@@ -8,6 +8,7 @@ from sqlalchemy.types import TypeEngine
|
||||
|
||||
class Types(object):
|
||||
"""A holder class for easy access to SQLAlchemy type names."""
|
||||
|
||||
integer = Integer
|
||||
string = Unicode
|
||||
text = UnicodeText
|
||||
|
||||
+11
-11
@@ -62,17 +62,17 @@ class ResultIter(object):
|
||||
def normalize_column_name(name):
|
||||
"""Check if a string is a reasonable thing to use as a column name."""
|
||||
if not isinstance(name, str):
|
||||
raise ValueError('%r is not a valid column name.' % name)
|
||||
raise ValueError("%r is not a valid column name." % name)
|
||||
|
||||
# limit to 63 characters
|
||||
name = name.strip()[:63]
|
||||
# column names can be 63 *bytes* max in postgresql
|
||||
if isinstance(name, str):
|
||||
while len(name.encode('utf-8')) >= 64:
|
||||
name = name[:len(name) - 1]
|
||||
while len(name.encode("utf-8")) >= 64:
|
||||
name = name[: len(name) - 1]
|
||||
|
||||
if not len(name) or '.' in name or '-' in name:
|
||||
raise ValueError('%r is not a valid column name.' % name)
|
||||
if not len(name) or "." in name or "-" in name:
|
||||
raise ValueError("%r is not a valid column name." % name)
|
||||
return name
|
||||
|
||||
|
||||
@@ -80,7 +80,7 @@ def normalize_column_key(name):
|
||||
"""Return a comparable column name."""
|
||||
if name is None or not isinstance(name, str):
|
||||
return None
|
||||
return name.upper().strip().replace(' ', '')
|
||||
return name.upper().strip().replace(" ", "")
|
||||
|
||||
|
||||
def normalize_table_name(name):
|
||||
@@ -97,16 +97,16 @@ def safe_url(url):
|
||||
"""Remove password from printed connection URLs."""
|
||||
parsed = urlparse(url)
|
||||
if parsed.password is not None:
|
||||
pwd = ':%s@' % parsed.password
|
||||
url = url.replace(pwd, ':*****@')
|
||||
pwd = ":%s@" % parsed.password
|
||||
url = url.replace(pwd, ":*****@")
|
||||
return url
|
||||
|
||||
|
||||
def index_name(table, columns):
|
||||
"""Generate an artificial index name."""
|
||||
sig = '||'.join(columns)
|
||||
key = sha1(sig.encode('utf-8')).hexdigest()[:16]
|
||||
return 'ix_%s_%s' % (table, key)
|
||||
sig = "||".join(columns)
|
||||
key = sha1(sig.encode("utf-8")).hexdigest()[:16]
|
||||
return "ix_%s_%s" % (table, key)
|
||||
|
||||
|
||||
def pad_chunk_columns(chunk, columns):
|
||||
|
||||
Reference in New Issue
Block a user