upgraded PyMysql to 0.7.9, thanks niphlod
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
'''
|
||||
PyMySQL: A pure-Python drop-in replacement for MySQLdb.
|
||||
"""
|
||||
PyMySQL: A pure-Python MySQL client library.
|
||||
|
||||
Copyright (c) 2010 PyMySQL contributors
|
||||
Copyright (c) 2010-2016 PyMySQL contributors
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
@@ -20,40 +20,32 @@ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
|
||||
THE SOFTWARE.
|
||||
|
||||
'''
|
||||
|
||||
VERSION = (0, 5, None)
|
||||
|
||||
from constants import FIELD_TYPE
|
||||
from converters import escape_dict, escape_sequence, escape_string
|
||||
from err import Warning, Error, InterfaceError, DataError, \
|
||||
DatabaseError, OperationalError, IntegrityError, InternalError, \
|
||||
NotSupportedError, ProgrammingError, MySQLError
|
||||
from times import Date, Time, Timestamp, \
|
||||
DateFromTicks, TimeFromTicks, TimestampFromTicks
|
||||
|
||||
"""
|
||||
import sys
|
||||
|
||||
try:
|
||||
frozenset
|
||||
except NameError:
|
||||
from sets import ImmutableSet as frozenset
|
||||
try:
|
||||
from sets import BaseSet as set
|
||||
except ImportError:
|
||||
from sets import Set as set
|
||||
from ._compat import PY2
|
||||
from .constants import FIELD_TYPE
|
||||
from .converters import escape_dict, escape_sequence, escape_string
|
||||
from .err import (
|
||||
Warning, Error, InterfaceError, DataError,
|
||||
DatabaseError, OperationalError, IntegrityError, InternalError,
|
||||
NotSupportedError, ProgrammingError, MySQLError)
|
||||
from .times import (
|
||||
Date, Time, Timestamp,
|
||||
DateFromTicks, TimeFromTicks, TimestampFromTicks)
|
||||
|
||||
|
||||
VERSION = (0, 7, 9, None)
|
||||
threadsafety = 1
|
||||
apilevel = "2.0"
|
||||
paramstyle = "format"
|
||||
paramstyle = "pyformat"
|
||||
|
||||
|
||||
class DBAPISet(frozenset):
|
||||
|
||||
|
||||
def __ne__(self, other):
|
||||
if isinstance(other, set):
|
||||
return super(DBAPISet, self).__ne__(self, other)
|
||||
return frozenset.__ne__(self, other)
|
||||
else:
|
||||
return other not in self
|
||||
|
||||
@@ -80,40 +72,52 @@ TIMESTAMP = DBAPISet([FIELD_TYPE.TIMESTAMP, FIELD_TYPE.DATETIME])
|
||||
DATETIME = TIMESTAMP
|
||||
ROWID = DBAPISet()
|
||||
|
||||
|
||||
def Binary(x):
|
||||
"""Return x as a binary type."""
|
||||
return str(x)
|
||||
if PY2:
|
||||
return bytearray(x)
|
||||
else:
|
||||
return bytes(x)
|
||||
|
||||
|
||||
def Connect(*args, **kwargs):
|
||||
"""
|
||||
Connect to the database; see connections.Connection.__init__() for
|
||||
more information.
|
||||
"""
|
||||
from connections import Connection
|
||||
from .connections import Connection
|
||||
return Connection(*args, **kwargs)
|
||||
|
||||
from pymysql import connections as _orig_conn
|
||||
if _orig_conn.Connection.__init__.__doc__ is not None:
|
||||
Connect.__doc__ = _orig_conn.Connection.__init__.__doc__
|
||||
del _orig_conn
|
||||
|
||||
|
||||
def get_client_info(): # for MySQLdb compatibility
|
||||
return '%s.%s.%s' % VERSION
|
||||
return '.'.join(map(str, VERSION))
|
||||
|
||||
connect = Connection = Connect
|
||||
|
||||
# we include a doctored version_info here for MySQLdb compatibility
|
||||
version_info = (1,2,2,"final",0)
|
||||
version_info = (1,2,6,"final",0)
|
||||
|
||||
NULL = "NULL"
|
||||
|
||||
__version__ = get_client_info()
|
||||
|
||||
def thread_safe():
|
||||
return True # match MySQLdb.thread_safe()
|
||||
return True # match MySQLdb.thread_safe()
|
||||
|
||||
def install_as_MySQLdb():
|
||||
"""
|
||||
After this function is called, any application that imports MySQLdb or
|
||||
_mysql will unwittingly actually use
|
||||
_mysql will unwittingly actually use
|
||||
"""
|
||||
sys.modules["MySQLdb"] = sys.modules["_mysql"] = sys.modules["pymysql"]
|
||||
|
||||
|
||||
__all__ = [
|
||||
'BINARY', 'Binary', 'Connect', 'Connection', 'DATE', 'Date',
|
||||
'Time', 'Timestamp', 'DateFromTicks', 'TimeFromTicks', 'TimestampFromTicks',
|
||||
@@ -126,6 +130,5 @@ __all__ = [
|
||||
'paramstyle', 'threadsafety', 'version_info',
|
||||
|
||||
"install_as_MySQLdb",
|
||||
|
||||
"NULL","__version__",
|
||||
]
|
||||
"NULL", "__version__",
|
||||
]
|
||||
|
||||
Executable
+21
@@ -0,0 +1,21 @@
|
||||
import sys
|
||||
|
||||
PY2 = sys.version_info[0] == 2
|
||||
PYPY = hasattr(sys, 'pypy_translation_info')
|
||||
JYTHON = sys.platform.startswith('java')
|
||||
IRONPYTHON = sys.platform == 'cli'
|
||||
CPYTHON = not PYPY and not JYTHON and not IRONPYTHON
|
||||
|
||||
if PY2:
|
||||
import __builtin__
|
||||
range_type = xrange
|
||||
text_type = unicode
|
||||
long_type = long
|
||||
str_type = basestring
|
||||
unichr = __builtin__.unichr
|
||||
else:
|
||||
range_type = range
|
||||
text_type = str
|
||||
long_type = int
|
||||
str_type = str
|
||||
unichr = chr
|
||||
Executable
+134
@@ -0,0 +1,134 @@
|
||||
"""
|
||||
SocketIO imported from socket module in Python 3.
|
||||
|
||||
Copyright (c) 2001-2013 Python Software Foundation; All Rights Reserved.
|
||||
"""
|
||||
|
||||
from socket import *
|
||||
import io
|
||||
import errno
|
||||
|
||||
__all__ = ['SocketIO']
|
||||
|
||||
EINTR = errno.EINTR
|
||||
_blocking_errnos = (errno.EAGAIN, errno.EWOULDBLOCK)
|
||||
|
||||
class SocketIO(io.RawIOBase):
|
||||
|
||||
"""Raw I/O implementation for stream sockets.
|
||||
|
||||
This class supports the makefile() method on sockets. It provides
|
||||
the raw I/O interface on top of a socket object.
|
||||
"""
|
||||
|
||||
# One might wonder why not let FileIO do the job instead. There are two
|
||||
# main reasons why FileIO is not adapted:
|
||||
# - it wouldn't work under Windows (where you can't used read() and
|
||||
# write() on a socket handle)
|
||||
# - it wouldn't work with socket timeouts (FileIO would ignore the
|
||||
# timeout and consider the socket non-blocking)
|
||||
|
||||
# XXX More docs
|
||||
|
||||
def __init__(self, sock, mode):
|
||||
if mode not in ("r", "w", "rw", "rb", "wb", "rwb"):
|
||||
raise ValueError("invalid mode: %r" % mode)
|
||||
io.RawIOBase.__init__(self)
|
||||
self._sock = sock
|
||||
if "b" not in mode:
|
||||
mode += "b"
|
||||
self._mode = mode
|
||||
self._reading = "r" in mode
|
||||
self._writing = "w" in mode
|
||||
self._timeout_occurred = False
|
||||
|
||||
def readinto(self, b):
|
||||
"""Read up to len(b) bytes into the writable buffer *b* and return
|
||||
the number of bytes read. If the socket is non-blocking and no bytes
|
||||
are available, None is returned.
|
||||
|
||||
If *b* is non-empty, a 0 return value indicates that the connection
|
||||
was shutdown at the other end.
|
||||
"""
|
||||
self._checkClosed()
|
||||
self._checkReadable()
|
||||
if self._timeout_occurred:
|
||||
raise IOError("cannot read from timed out object")
|
||||
while True:
|
||||
try:
|
||||
return self._sock.recv_into(b)
|
||||
except timeout:
|
||||
self._timeout_occurred = True
|
||||
raise
|
||||
except error as e:
|
||||
n = e.args[0]
|
||||
if n == EINTR:
|
||||
continue
|
||||
if n in _blocking_errnos:
|
||||
return None
|
||||
raise
|
||||
|
||||
def write(self, b):
|
||||
"""Write the given bytes or bytearray object *b* to the socket
|
||||
and return the number of bytes written. This can be less than
|
||||
len(b) if not all data could be written. If the socket is
|
||||
non-blocking and no bytes could be written None is returned.
|
||||
"""
|
||||
self._checkClosed()
|
||||
self._checkWritable()
|
||||
try:
|
||||
return self._sock.send(b)
|
||||
except error as e:
|
||||
# XXX what about EINTR?
|
||||
if e.args[0] in _blocking_errnos:
|
||||
return None
|
||||
raise
|
||||
|
||||
def readable(self):
|
||||
"""True if the SocketIO is open for reading.
|
||||
"""
|
||||
if self.closed:
|
||||
raise ValueError("I/O operation on closed socket.")
|
||||
return self._reading
|
||||
|
||||
def writable(self):
|
||||
"""True if the SocketIO is open for writing.
|
||||
"""
|
||||
if self.closed:
|
||||
raise ValueError("I/O operation on closed socket.")
|
||||
return self._writing
|
||||
|
||||
def seekable(self):
|
||||
"""True if the SocketIO is open for seeking.
|
||||
"""
|
||||
if self.closed:
|
||||
raise ValueError("I/O operation on closed socket.")
|
||||
return super().seekable()
|
||||
|
||||
def fileno(self):
|
||||
"""Return the file descriptor of the underlying socket.
|
||||
"""
|
||||
self._checkClosed()
|
||||
return self._sock.fileno()
|
||||
|
||||
@property
|
||||
def name(self):
|
||||
if not self.closed:
|
||||
return self.fileno()
|
||||
else:
|
||||
return -1
|
||||
|
||||
@property
|
||||
def mode(self):
|
||||
return self._mode
|
||||
|
||||
def close(self):
|
||||
"""Close the SocketIO object. This doesn't close the underlying
|
||||
socket, except if all references to it have disappeared.
|
||||
"""
|
||||
if self.closed:
|
||||
return
|
||||
io.RawIOBase.close(self)
|
||||
self._sock._decref_socketios()
|
||||
self._sock = None
|
||||
|
||||
@@ -5,11 +5,28 @@ MBLENGTH = {
|
||||
91:2
|
||||
}
|
||||
|
||||
class Charset:
|
||||
|
||||
class Charset(object):
|
||||
def __init__(self, id, name, collation, is_default):
|
||||
self.id, self.name, self.collation = id, name, collation
|
||||
self.is_default = is_default == 'Yes'
|
||||
|
||||
def __repr__(self):
|
||||
return "Charset(id=%s, name=%r, collation=%r)" % (
|
||||
self.id, self.name, self.collation)
|
||||
|
||||
@property
|
||||
def encoding(self):
|
||||
name = self.name
|
||||
if name == 'utf8mb4':
|
||||
return 'utf8'
|
||||
return name
|
||||
|
||||
@property
|
||||
def is_binary(self):
|
||||
return self.id == 63
|
||||
|
||||
|
||||
class Charsets:
|
||||
def __init__(self):
|
||||
self._by_id = {}
|
||||
@@ -21,6 +38,7 @@ class Charsets:
|
||||
return self._by_id[id]
|
||||
|
||||
def by_name(self, name):
|
||||
name = name.lower()
|
||||
for c in self._by_id.values():
|
||||
if c.name == name and c.is_default:
|
||||
return c
|
||||
@@ -92,13 +110,11 @@ _charsets.add(Charset(52, 'cp1251', 'cp1251_general_cs', ''))
|
||||
_charsets.add(Charset(53, 'macroman', 'macroman_bin', ''))
|
||||
_charsets.add(Charset(54, 'utf16', 'utf16_general_ci', 'Yes'))
|
||||
_charsets.add(Charset(55, 'utf16', 'utf16_bin', ''))
|
||||
_charsets.add(Charset(56, 'utf16le', 'utf16le_general_ci', 'Yes'))
|
||||
_charsets.add(Charset(57, 'cp1256', 'cp1256_general_ci', 'Yes'))
|
||||
_charsets.add(Charset(58, 'cp1257', 'cp1257_bin', ''))
|
||||
_charsets.add(Charset(59, 'cp1257', 'cp1257_general_ci', 'Yes'))
|
||||
_charsets.add(Charset(60, 'utf32', 'utf32_general_ci', 'Yes'))
|
||||
_charsets.add(Charset(61, 'utf32', 'utf32_bin', ''))
|
||||
_charsets.add(Charset(62, 'utf16le', 'utf16le_bin', ''))
|
||||
_charsets.add(Charset(63, 'binary', 'binary', 'Yes'))
|
||||
_charsets.add(Charset(64, 'armscii8', 'armscii8_bin', ''))
|
||||
_charsets.add(Charset(65, 'ascii', 'ascii_bin', ''))
|
||||
@@ -155,10 +171,6 @@ _charsets.add(Charset(117, 'utf16', 'utf16_persian_ci', ''))
|
||||
_charsets.add(Charset(118, 'utf16', 'utf16_esperanto_ci', ''))
|
||||
_charsets.add(Charset(119, 'utf16', 'utf16_hungarian_ci', ''))
|
||||
_charsets.add(Charset(120, 'utf16', 'utf16_sinhala_ci', ''))
|
||||
_charsets.add(Charset(121, 'utf16', 'utf16_german2_ci', ''))
|
||||
_charsets.add(Charset(122, 'utf16', 'utf16_croatian_ci', ''))
|
||||
_charsets.add(Charset(123, 'utf16', 'utf16_unicode_520_ci', ''))
|
||||
_charsets.add(Charset(124, 'utf16', 'utf16_vietnamese_ci', ''))
|
||||
_charsets.add(Charset(128, 'ucs2', 'ucs2_unicode_ci', ''))
|
||||
_charsets.add(Charset(129, 'ucs2', 'ucs2_icelandic_ci', ''))
|
||||
_charsets.add(Charset(130, 'ucs2', 'ucs2_latvian_ci', ''))
|
||||
@@ -179,10 +191,6 @@ _charsets.add(Charset(144, 'ucs2', 'ucs2_persian_ci', ''))
|
||||
_charsets.add(Charset(145, 'ucs2', 'ucs2_esperanto_ci', ''))
|
||||
_charsets.add(Charset(146, 'ucs2', 'ucs2_hungarian_ci', ''))
|
||||
_charsets.add(Charset(147, 'ucs2', 'ucs2_sinhala_ci', ''))
|
||||
_charsets.add(Charset(148, 'ucs2', 'ucs2_german2_ci', ''))
|
||||
_charsets.add(Charset(149, 'ucs2', 'ucs2_croatian_ci', ''))
|
||||
_charsets.add(Charset(150, 'ucs2', 'ucs2_unicode_520_ci', ''))
|
||||
_charsets.add(Charset(151, 'ucs2', 'ucs2_vietnamese_ci', ''))
|
||||
_charsets.add(Charset(159, 'ucs2', 'ucs2_general_mysql500_ci', ''))
|
||||
_charsets.add(Charset(160, 'utf32', 'utf32_unicode_ci', ''))
|
||||
_charsets.add(Charset(161, 'utf32', 'utf32_icelandic_ci', ''))
|
||||
@@ -204,10 +212,6 @@ _charsets.add(Charset(176, 'utf32', 'utf32_persian_ci', ''))
|
||||
_charsets.add(Charset(177, 'utf32', 'utf32_esperanto_ci', ''))
|
||||
_charsets.add(Charset(178, 'utf32', 'utf32_hungarian_ci', ''))
|
||||
_charsets.add(Charset(179, 'utf32', 'utf32_sinhala_ci', ''))
|
||||
_charsets.add(Charset(180, 'utf32', 'utf32_german2_ci', ''))
|
||||
_charsets.add(Charset(181, 'utf32', 'utf32_croatian_ci', ''))
|
||||
_charsets.add(Charset(182, 'utf32', 'utf32_unicode_520_ci', ''))
|
||||
_charsets.add(Charset(183, 'utf32', 'utf32_vietnamese_ci', ''))
|
||||
_charsets.add(Charset(192, 'utf8', 'utf8_unicode_ci', ''))
|
||||
_charsets.add(Charset(193, 'utf8', 'utf8_icelandic_ci', ''))
|
||||
_charsets.add(Charset(194, 'utf8', 'utf8_latvian_ci', ''))
|
||||
@@ -228,10 +232,6 @@ _charsets.add(Charset(208, 'utf8', 'utf8_persian_ci', ''))
|
||||
_charsets.add(Charset(209, 'utf8', 'utf8_esperanto_ci', ''))
|
||||
_charsets.add(Charset(210, 'utf8', 'utf8_hungarian_ci', ''))
|
||||
_charsets.add(Charset(211, 'utf8', 'utf8_sinhala_ci', ''))
|
||||
_charsets.add(Charset(212, 'utf8', 'utf8_german2_ci', ''))
|
||||
_charsets.add(Charset(213, 'utf8', 'utf8_croatian_ci', ''))
|
||||
_charsets.add(Charset(214, 'utf8', 'utf8_unicode_520_ci', ''))
|
||||
_charsets.add(Charset(215, 'utf8', 'utf8_vietnamese_ci', ''))
|
||||
_charsets.add(Charset(223, 'utf8', 'utf8_general_mysql500_ci', ''))
|
||||
_charsets.add(Charset(224, 'utf8mb4', 'utf8mb4_unicode_ci', ''))
|
||||
_charsets.add(Charset(225, 'utf8mb4', 'utf8mb4_icelandic_ci', ''))
|
||||
@@ -258,9 +258,13 @@ _charsets.add(Charset(245, 'utf8mb4', 'utf8mb4_croatian_ci', ''))
|
||||
_charsets.add(Charset(246, 'utf8mb4', 'utf8mb4_unicode_520_ci', ''))
|
||||
_charsets.add(Charset(247, 'utf8mb4', 'utf8mb4_vietnamese_ci', ''))
|
||||
|
||||
def charset_by_name(name):
|
||||
return _charsets.by_name(name)
|
||||
|
||||
def charset_by_id(id):
|
||||
return _charsets.by_id(id)
|
||||
charset_by_name = _charsets.by_name
|
||||
charset_by_id = _charsets.by_id
|
||||
|
||||
|
||||
def charset_to_encoding(name):
|
||||
"""Convert MySQL's charset name to Python's codec name"""
|
||||
if name == 'utf8mb4':
|
||||
return 'utf8'
|
||||
return name
|
||||
|
||||
+1077
-647
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
||||
|
||||
# https://dev.mysql.com/doc/internals/en/capability-flags.html#packet-Protocol::CapabilityFlags
|
||||
LONG_PASSWORD = 1
|
||||
FOUND_ROWS = 1 << 1
|
||||
LONG_FLAG = 1 << 2
|
||||
@@ -12,9 +12,20 @@ PROTOCOL_41 = 1 << 9
|
||||
INTERACTIVE = 1 << 10
|
||||
SSL = 1 << 11
|
||||
IGNORE_SIGPIPE = 1 << 12
|
||||
TRANSACTIONS = 1 << 13
|
||||
TRANSACTIONS = 1 << 13
|
||||
SECURE_CONNECTION = 1 << 15
|
||||
MULTI_STATEMENTS = 1 << 16
|
||||
MULTI_RESULTS = 1 << 17
|
||||
CAPABILITIES = LONG_PASSWORD|LONG_FLAG|TRANSACTIONS| \
|
||||
PROTOCOL_41|SECURE_CONNECTION
|
||||
PS_MULTI_RESULTS = 1 << 18
|
||||
PLUGIN_AUTH = 1 << 19
|
||||
PLUGIN_AUTH_LENENC_CLIENT_DATA = 1 << 21
|
||||
CAPABILITIES = (
|
||||
LONG_PASSWORD | LONG_FLAG | PROTOCOL_41 | TRANSACTIONS
|
||||
| SECURE_CONNECTION | MULTI_STATEMENTS | MULTI_RESULTS
|
||||
| PLUGIN_AUTH | PLUGIN_AUTH_LENENC_CLIENT_DATA)
|
||||
|
||||
# Not done yet
|
||||
CONNECT_ATTRS = 1 << 20
|
||||
HANDLE_EXPIRED_PASSWORDS = 1 << 22
|
||||
SESSION_TRACK = 1 << 23
|
||||
DEPRECATE_EOF = 1 << 24
|
||||
|
||||
@@ -21,3 +21,13 @@ COM_BINLOG_DUMP = 0x12
|
||||
COM_TABLE_DUMP = 0x13
|
||||
COM_CONNECT_OUT = 0x14
|
||||
COM_REGISTER_SLAVE = 0x15
|
||||
COM_STMT_PREPARE = 0x16
|
||||
COM_STMT_EXECUTE = 0x17
|
||||
COM_STMT_SEND_LONG_DATA = 0x18
|
||||
COM_STMT_CLOSE = 0x19
|
||||
COM_STMT_RESET = 0x1a
|
||||
COM_SET_OPTION = 0x1b
|
||||
COM_STMT_FETCH = 0x1c
|
||||
COM_DAEMON = 0x1d
|
||||
COM_BINLOG_DUMP_GTID = 0x1e
|
||||
COM_END = 0x1f
|
||||
|
||||
Executable
+68
@@ -0,0 +1,68 @@
|
||||
# flake8: noqa
|
||||
# errmsg.h
|
||||
CR_ERROR_FIRST = 2000
|
||||
CR_UNKNOWN_ERROR = 2000
|
||||
CR_SOCKET_CREATE_ERROR = 2001
|
||||
CR_CONNECTION_ERROR = 2002
|
||||
CR_CONN_HOST_ERROR = 2003
|
||||
CR_IPSOCK_ERROR = 2004
|
||||
CR_UNKNOWN_HOST = 2005
|
||||
CR_SERVER_GONE_ERROR = 2006
|
||||
CR_VERSION_ERROR = 2007
|
||||
CR_OUT_OF_MEMORY = 2008
|
||||
CR_WRONG_HOST_INFO = 2009
|
||||
CR_LOCALHOST_CONNECTION = 2010
|
||||
CR_TCP_CONNECTION = 2011
|
||||
CR_SERVER_HANDSHAKE_ERR = 2012
|
||||
CR_SERVER_LOST = 2013
|
||||
CR_COMMANDS_OUT_OF_SYNC = 2014
|
||||
CR_NAMEDPIPE_CONNECTION = 2015
|
||||
CR_NAMEDPIPEWAIT_ERROR = 2016
|
||||
CR_NAMEDPIPEOPEN_ERROR = 2017
|
||||
CR_NAMEDPIPESETSTATE_ERROR = 2018
|
||||
CR_CANT_READ_CHARSET = 2019
|
||||
CR_NET_PACKET_TOO_LARGE = 2020
|
||||
CR_EMBEDDED_CONNECTION = 2021
|
||||
CR_PROBE_SLAVE_STATUS = 2022
|
||||
CR_PROBE_SLAVE_HOSTS = 2023
|
||||
CR_PROBE_SLAVE_CONNECT = 2024
|
||||
CR_PROBE_MASTER_CONNECT = 2025
|
||||
CR_SSL_CONNECTION_ERROR = 2026
|
||||
CR_MALFORMED_PACKET = 2027
|
||||
CR_WRONG_LICENSE = 2028
|
||||
|
||||
CR_NULL_POINTER = 2029
|
||||
CR_NO_PREPARE_STMT = 2030
|
||||
CR_PARAMS_NOT_BOUND = 2031
|
||||
CR_DATA_TRUNCATED = 2032
|
||||
CR_NO_PARAMETERS_EXISTS = 2033
|
||||
CR_INVALID_PARAMETER_NO = 2034
|
||||
CR_INVALID_BUFFER_USE = 2035
|
||||
CR_UNSUPPORTED_PARAM_TYPE = 2036
|
||||
|
||||
CR_SHARED_MEMORY_CONNECTION = 2037
|
||||
CR_SHARED_MEMORY_CONNECT_REQUEST_ERROR = 2038
|
||||
CR_SHARED_MEMORY_CONNECT_ANSWER_ERROR = 2039
|
||||
CR_SHARED_MEMORY_CONNECT_FILE_MAP_ERROR = 2040
|
||||
CR_SHARED_MEMORY_CONNECT_MAP_ERROR = 2041
|
||||
CR_SHARED_MEMORY_FILE_MAP_ERROR = 2042
|
||||
CR_SHARED_MEMORY_MAP_ERROR = 2043
|
||||
CR_SHARED_MEMORY_EVENT_ERROR = 2044
|
||||
CR_SHARED_MEMORY_CONNECT_ABANDONED_ERROR = 2045
|
||||
CR_SHARED_MEMORY_CONNECT_SET_ERROR = 2046
|
||||
CR_CONN_UNKNOW_PROTOCOL = 2047
|
||||
CR_INVALID_CONN_HANDLE = 2048
|
||||
CR_SECURE_AUTH = 2049
|
||||
CR_FETCH_CANCELED = 2050
|
||||
CR_NO_DATA = 2051
|
||||
CR_NO_STMT_METADATA = 2052
|
||||
CR_NO_RESULT_SET = 2053
|
||||
CR_NOT_IMPLEMENTED = 2054
|
||||
CR_SERVER_LOST_EXTENDED = 2055
|
||||
CR_STMT_CLOSED = 2056
|
||||
CR_NEW_STMT_METADATA = 2057
|
||||
CR_ALREADY_CONNECTED = 2058
|
||||
CR_AUTH_PLUGIN_CANNOT_LOAD = 2059
|
||||
CR_DUPLICATE_CONNECTION_ATTR = 2060
|
||||
CR_AUTH_PLUGIN_ERR = 2061
|
||||
CR_ERROR_LAST = 2061
|
||||
@@ -17,6 +17,7 @@ YEAR = 13
|
||||
NEWDATE = 14
|
||||
VARCHAR = 15
|
||||
BIT = 16
|
||||
JSON = 245
|
||||
NEWDECIMAL = 246
|
||||
ENUM = 247
|
||||
SET = 248
|
||||
|
||||
@@ -9,4 +9,3 @@ SERVER_STATUS_LAST_ROW_SENT = 128
|
||||
SERVER_STATUS_DB_DROPPED = 256
|
||||
SERVER_STATUS_NO_BACKSLASH_ESCAPES = 512
|
||||
SERVER_STATUS_METADATA_CHANGED = 1024
|
||||
|
||||
|
||||
+241
-178
@@ -1,106 +1,162 @@
|
||||
import re
|
||||
from ._compat import PY2, text_type, long_type, JYTHON, IRONPYTHON, unichr
|
||||
|
||||
import datetime
|
||||
from decimal import Decimal
|
||||
import re
|
||||
import time
|
||||
import sys
|
||||
|
||||
from constants import FIELD_TYPE, FLAG
|
||||
from charset import charset_by_id
|
||||
from .constants import FIELD_TYPE, FLAG
|
||||
from .charset import charset_by_id, charset_to_encoding
|
||||
|
||||
PYTHON3 = sys.version_info[0] > 2
|
||||
|
||||
try:
|
||||
set
|
||||
except NameError:
|
||||
try:
|
||||
from sets import BaseSet as set
|
||||
except ImportError:
|
||||
from sets import Set as set
|
||||
def escape_item(val, charset, mapping=None):
|
||||
if mapping is None:
|
||||
mapping = encoders
|
||||
encoder = mapping.get(type(val))
|
||||
|
||||
ESCAPE_REGEX = re.compile(r"[\0\n\r\032\'\"\\]")
|
||||
ESCAPE_MAP = {'\0': '\\0', '\n': '\\n', '\r': '\\r', '\032': '\\Z',
|
||||
'\'': '\\\'', '"': '\\"', '\\': '\\\\'}
|
||||
# Fallback to default when no encoder found
|
||||
if not encoder:
|
||||
try:
|
||||
encoder = mapping[text_type]
|
||||
except KeyError:
|
||||
raise TypeError("no default type converter defined")
|
||||
|
||||
def escape_item(val, charset):
|
||||
if type(val) in [tuple, list, set]:
|
||||
return escape_sequence(val, charset)
|
||||
if type(val) is dict:
|
||||
return escape_dict(val, charset)
|
||||
if PYTHON3 and hasattr(val, "decode") and not isinstance(val, unicode):
|
||||
# deal with py3k bytes
|
||||
val = val.decode(charset)
|
||||
encoder = encoders[type(val)]
|
||||
val = encoder(val)
|
||||
if type(val) in [str, int]:
|
||||
return val
|
||||
val = val.encode(charset)
|
||||
if encoder in (escape_dict, escape_sequence):
|
||||
val = encoder(val, charset, mapping)
|
||||
else:
|
||||
val = encoder(val, mapping)
|
||||
return val
|
||||
|
||||
def escape_dict(val, charset):
|
||||
def escape_dict(val, charset, mapping=None):
|
||||
n = {}
|
||||
for k, v in val.items():
|
||||
quoted = escape_item(v, charset)
|
||||
quoted = escape_item(v, charset, mapping)
|
||||
n[k] = quoted
|
||||
return n
|
||||
|
||||
def escape_sequence(val, charset):
|
||||
def escape_sequence(val, charset, mapping=None):
|
||||
n = []
|
||||
for item in val:
|
||||
quoted = escape_item(item, charset)
|
||||
quoted = escape_item(item, charset, mapping)
|
||||
n.append(quoted)
|
||||
return "(" + ",".join(n) + ")"
|
||||
|
||||
def escape_set(val, charset):
|
||||
val = map(lambda x: escape_item(x, charset), val)
|
||||
return ','.join(val)
|
||||
def escape_set(val, charset, mapping=None):
|
||||
return ','.join([escape_item(x, charset, mapping) for x in val])
|
||||
|
||||
def escape_bool(value):
|
||||
def escape_bool(value, mapping=None):
|
||||
return str(int(value))
|
||||
|
||||
def escape_object(value):
|
||||
def escape_object(value, mapping=None):
|
||||
return str(value)
|
||||
|
||||
def escape_int(value):
|
||||
return value
|
||||
def escape_int(value, mapping=None):
|
||||
return str(value)
|
||||
|
||||
escape_long = escape_object
|
||||
|
||||
def escape_float(value):
|
||||
def escape_float(value, mapping=None):
|
||||
return ('%.15g' % value)
|
||||
|
||||
def escape_string(value):
|
||||
return ("'%s'" % ESCAPE_REGEX.sub(
|
||||
lambda match: ESCAPE_MAP.get(match.group(0)), value))
|
||||
_escape_table = [unichr(x) for x in range(128)]
|
||||
_escape_table[0] = u'\\0'
|
||||
_escape_table[ord('\\')] = u'\\\\'
|
||||
_escape_table[ord('\n')] = u'\\n'
|
||||
_escape_table[ord('\r')] = u'\\r'
|
||||
_escape_table[ord('\032')] = u'\\Z'
|
||||
_escape_table[ord('"')] = u'\\"'
|
||||
_escape_table[ord("'")] = u"\\'"
|
||||
|
||||
def escape_unicode(value):
|
||||
return escape_string(value)
|
||||
def _escape_unicode(value, mapping=None):
|
||||
"""escapes *value* without adding quote.
|
||||
|
||||
def escape_None(value):
|
||||
Value should be unicode
|
||||
"""
|
||||
return value.translate(_escape_table)
|
||||
|
||||
if PY2:
|
||||
def escape_string(value, mapping=None):
|
||||
"""escape_string escapes *value* but not surround it with quotes.
|
||||
|
||||
Value should be bytes or unicode.
|
||||
"""
|
||||
if isinstance(value, unicode):
|
||||
return _escape_unicode(value)
|
||||
assert isinstance(value, (bytes, bytearray))
|
||||
value = value.replace('\\', '\\\\')
|
||||
value = value.replace('\0', '\\0')
|
||||
value = value.replace('\n', '\\n')
|
||||
value = value.replace('\r', '\\r')
|
||||
value = value.replace('\032', '\\Z')
|
||||
value = value.replace("'", "\\'")
|
||||
value = value.replace('"', '\\"')
|
||||
return value
|
||||
|
||||
def escape_bytes(value, mapping=None):
|
||||
assert isinstance(value, (bytes, bytearray))
|
||||
return b"_binary'%s'" % escape_string(value)
|
||||
else:
|
||||
escape_string = _escape_unicode
|
||||
|
||||
# On Python ~3.5, str.decode('ascii', 'surrogateescape') is slow.
|
||||
# (fixed in Python 3.6, http://bugs.python.org/issue24870)
|
||||
# Workaround is str.decode('latin1') then translate 0x80-0xff into 0udc80-0udcff.
|
||||
# We can escape special chars and surrogateescape at once.
|
||||
_escape_bytes_table = _escape_table + [chr(i) for i in range(0xdc80, 0xdd00)]
|
||||
|
||||
def escape_bytes(value, mapping=None):
|
||||
return "_binary'%s'" % value.decode('latin1').translate(_escape_bytes_table)
|
||||
|
||||
|
||||
def escape_unicode(value, mapping=None):
|
||||
return u"'%s'" % _escape_unicode(value)
|
||||
|
||||
def escape_str(value, mapping=None):
|
||||
return "'%s'" % escape_string(str(value), mapping)
|
||||
|
||||
def escape_None(value, mapping=None):
|
||||
return 'NULL'
|
||||
|
||||
def escape_timedelta(obj):
|
||||
def escape_timedelta(obj, mapping=None):
|
||||
seconds = int(obj.seconds) % 60
|
||||
minutes = int(obj.seconds // 60) % 60
|
||||
hours = int(obj.seconds // 3600) % 24 + int(obj.days) * 24
|
||||
return escape_string('%02d:%02d:%02d' % (hours, minutes, seconds))
|
||||
if obj.microseconds:
|
||||
fmt = "'{0:02d}:{1:02d}:{2:02d}.{3:06d}'"
|
||||
else:
|
||||
fmt = "'{0:02d}:{1:02d}:{2:02d}'"
|
||||
return fmt.format(hours, minutes, seconds, obj.microseconds)
|
||||
|
||||
def escape_time(obj):
|
||||
s = "%02d:%02d:%02d" % (int(obj.hour), int(obj.minute),
|
||||
int(obj.second))
|
||||
def escape_time(obj, mapping=None):
|
||||
if obj.microsecond:
|
||||
s += ".%f" % obj.microsecond
|
||||
fmt = "'{0.hour:02}:{0.minute:02}:{0.second:02}.{0.microsecond:06}'"
|
||||
else:
|
||||
fmt = "'{0.hour:02}:{0.minute:02}:{0.second:02}'"
|
||||
return fmt.format(obj)
|
||||
|
||||
return escape_string(s)
|
||||
def escape_datetime(obj, mapping=None):
|
||||
if obj.microsecond:
|
||||
fmt = "'{0.year:04}-{0.month:02}-{0.day:02} {0.hour:02}:{0.minute:02}:{0.second:02}.{0.microsecond:06}'"
|
||||
else:
|
||||
fmt = "'{0.year:04}-{0.month:02}-{0.day:02} {0.hour:02}:{0.minute:02}:{0.second:02}'"
|
||||
return fmt.format(obj)
|
||||
|
||||
def escape_datetime(obj):
|
||||
return escape_string(obj.strftime("%Y-%m-%d %H:%M:%S"))
|
||||
def escape_date(obj, mapping=None):
|
||||
fmt = "'{0.year:04}-{0.month:02}-{0.day:02}'"
|
||||
return fmt.format(obj)
|
||||
|
||||
def escape_date(obj):
|
||||
return escape_string(obj.strftime("%Y-%m-%d"))
|
||||
|
||||
def escape_struct_time(obj):
|
||||
def escape_struct_time(obj, mapping=None):
|
||||
return escape_datetime(datetime.datetime(*obj[:6]))
|
||||
|
||||
def convert_datetime(connection, field, obj):
|
||||
def _convert_second_fraction(s):
|
||||
if not s:
|
||||
return 0
|
||||
# Pad zeros to ensure the fraction length in microseconds
|
||||
s = s.ljust(6, '0')
|
||||
return int(s[:6])
|
||||
|
||||
DATETIME_RE = re.compile(r"(\d{1,4})-(\d{1,2})-(\d{1,2})[T ](\d{1,2}):(\d{1,2}):(\d{1,2})(?:.(\d{1,6}))?")
|
||||
|
||||
|
||||
def convert_datetime(obj):
|
||||
"""Returns a DATETIME or TIMESTAMP column value as a datetime object:
|
||||
|
||||
>>> datetime_or_None('2007-02-25 23:06:20')
|
||||
@@ -116,22 +172,24 @@ def convert_datetime(connection, field, obj):
|
||||
True
|
||||
|
||||
"""
|
||||
if not isinstance(obj, unicode):
|
||||
obj = obj.decode(connection.charset)
|
||||
if ' ' in obj:
|
||||
sep = ' '
|
||||
elif 'T' in obj:
|
||||
sep = 'T'
|
||||
else:
|
||||
return convert_date(connection, field, obj)
|
||||
if not PY2 and isinstance(obj, (bytes, bytearray)):
|
||||
obj = obj.decode('ascii')
|
||||
|
||||
m = DATETIME_RE.match(obj)
|
||||
if not m:
|
||||
return convert_date(obj)
|
||||
|
||||
try:
|
||||
ymd, hms = obj.split(sep, 1)
|
||||
return datetime.datetime(*[ int(x) for x in ymd.split('-')+hms.split(':') ])
|
||||
groups = list(m.groups())
|
||||
groups[-1] = _convert_second_fraction(groups[-1])
|
||||
return datetime.datetime(*[ int(x) for x in groups ])
|
||||
except ValueError:
|
||||
return convert_date(connection, field, obj)
|
||||
return convert_date(obj)
|
||||
|
||||
def convert_timedelta(connection, field, obj):
|
||||
TIMEDELTA_RE = re.compile(r"(-)?(\d{1,3}):(\d{1,2}):(\d{1,2})(?:.(\d{1,6}))?")
|
||||
|
||||
|
||||
def convert_timedelta(obj):
|
||||
"""Returns a TIME column as a timedelta object:
|
||||
|
||||
>>> timedelta_or_None('25:06:17')
|
||||
@@ -148,25 +206,33 @@ def convert_timedelta(connection, field, obj):
|
||||
can accept values as (+|-)DD HH:MM:SS. The latter format will not
|
||||
be parsed correctly by this function.
|
||||
"""
|
||||
if not PY2 and isinstance(obj, (bytes, bytearray)):
|
||||
obj = obj.decode('ascii')
|
||||
|
||||
m = TIMEDELTA_RE.match(obj)
|
||||
if not m:
|
||||
return None
|
||||
|
||||
try:
|
||||
microseconds = 0
|
||||
if not isinstance(obj, unicode):
|
||||
obj = obj.decode(connection.charset)
|
||||
if "." in obj:
|
||||
(obj, tail) = obj.split('.')
|
||||
microseconds = int(tail)
|
||||
hours, minutes, seconds = obj.split(':')
|
||||
groups = list(m.groups())
|
||||
groups[-1] = _convert_second_fraction(groups[-1])
|
||||
negate = -1 if groups[0] else 1
|
||||
hours, minutes, seconds, microseconds = groups[1:]
|
||||
|
||||
tdelta = datetime.timedelta(
|
||||
hours = int(hours),
|
||||
minutes = int(minutes),
|
||||
seconds = int(seconds),
|
||||
microseconds = microseconds
|
||||
)
|
||||
microseconds = int(microseconds)
|
||||
) * negate
|
||||
return tdelta
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def convert_time(connection, field, obj):
|
||||
TIME_RE = re.compile(r"(\d{1,2}):(\d{1,2}):(\d{1,2})(?:.(\d{1,6}))?")
|
||||
|
||||
|
||||
def convert_time(obj):
|
||||
"""Returns a TIME column as a time object:
|
||||
|
||||
>>> time_or_None('15:06:17')
|
||||
@@ -188,18 +254,24 @@ def convert_time(connection, field, obj):
|
||||
to be treated as time-of-day and not a time offset, then you can
|
||||
use set this function as the converter for FIELD_TYPE.TIME.
|
||||
"""
|
||||
if not PY2 and isinstance(obj, (bytes, bytearray)):
|
||||
obj = obj.decode('ascii')
|
||||
|
||||
m = TIME_RE.match(obj)
|
||||
if not m:
|
||||
return None
|
||||
|
||||
try:
|
||||
microseconds = 0
|
||||
if "." in obj:
|
||||
(obj, tail) = obj.split('.')
|
||||
microseconds = int(tail)
|
||||
hours, minutes, seconds = obj.split(':')
|
||||
groups = list(m.groups())
|
||||
groups[-1] = _convert_second_fraction(groups[-1])
|
||||
hours, minutes, seconds, microseconds = groups
|
||||
return datetime.time(hour=int(hours), minute=int(minutes),
|
||||
second=int(seconds), microsecond=microseconds)
|
||||
second=int(seconds), microsecond=int(microseconds))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def convert_date(connection, field, obj):
|
||||
|
||||
def convert_date(obj):
|
||||
"""Returns a DATE column as a date object:
|
||||
|
||||
>>> date_or_None('2007-02-26')
|
||||
@@ -213,14 +285,15 @@ def convert_date(connection, field, obj):
|
||||
True
|
||||
|
||||
"""
|
||||
if not PY2 and isinstance(obj, (bytes, bytearray)):
|
||||
obj = obj.decode('ascii')
|
||||
try:
|
||||
if not isinstance(obj, unicode):
|
||||
obj = obj.decode(connection.charset)
|
||||
return datetime.date(*[ int(x) for x in obj.split('-', 2) ])
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def convert_mysql_timestamp(connection, field, timestamp):
|
||||
|
||||
def convert_mysql_timestamp(timestamp):
|
||||
"""Convert a MySQL TIMESTAMP to a Timestamp object.
|
||||
|
||||
MySQL >= 4.1 returns TIMESTAMP in the same format as DATETIME:
|
||||
@@ -241,11 +314,10 @@ def convert_mysql_timestamp(connection, field, timestamp):
|
||||
True
|
||||
|
||||
"""
|
||||
if not isinstance(timestamp, unicode):
|
||||
timestamp = timestamp.decode(connection.charset)
|
||||
|
||||
if not PY2 and isinstance(timestamp, (bytes, bytearray)):
|
||||
timestamp = timestamp.decode('ascii')
|
||||
if timestamp[4] == '-':
|
||||
return convert_datetime(connection, field, timestamp)
|
||||
return convert_datetime(timestamp)
|
||||
timestamp += "0"*(14-len(timestamp)) # padding
|
||||
year, month, day, hour, minute, second = \
|
||||
int(timestamp[:4]), int(timestamp[4:6]), int(timestamp[6:8]), \
|
||||
@@ -256,101 +328,92 @@ def convert_mysql_timestamp(connection, field, timestamp):
|
||||
return None
|
||||
|
||||
def convert_set(s):
|
||||
if isinstance(s, (bytes, bytearray)):
|
||||
return set(s.split(b","))
|
||||
return set(s.split(","))
|
||||
|
||||
def convert_bit(connection, field, b):
|
||||
#b = "\x00" * (8 - len(b)) + b # pad w/ zeroes
|
||||
#return struct.unpack(">Q", b)[0]
|
||||
#
|
||||
# the snippet above is right, but MySQLdb doesn't process bits,
|
||||
# so we shouldn't either
|
||||
return b
|
||||
|
||||
def through(x):
|
||||
return x
|
||||
|
||||
|
||||
#def convert_bit(b):
|
||||
# b = "\x00" * (8 - len(b)) + b # pad w/ zeroes
|
||||
# return struct.unpack(">Q", b)[0]
|
||||
#
|
||||
# the snippet above is right, but MySQLdb doesn't process bits,
|
||||
# so we shouldn't either
|
||||
convert_bit = through
|
||||
|
||||
|
||||
def convert_characters(connection, field, data):
|
||||
field_charset = charset_by_id(field.charsetnr).name
|
||||
encoding = charset_to_encoding(field_charset)
|
||||
if field.flags & FLAG.SET:
|
||||
return convert_set(data.decode(field_charset))
|
||||
return convert_set(data.decode(encoding))
|
||||
if field.flags & FLAG.BINARY:
|
||||
return data
|
||||
|
||||
if connection.use_unicode:
|
||||
data = data.decode(field_charset)
|
||||
data = data.decode(encoding)
|
||||
elif connection.charset != field_charset:
|
||||
data = data.decode(field_charset)
|
||||
data = data.encode(connection.charset)
|
||||
data = data.decode(encoding)
|
||||
data = data.encode(connection.encoding)
|
||||
return data
|
||||
|
||||
def convert_int(connection, field, data):
|
||||
return int(data)
|
||||
|
||||
def convert_long(connection, field, data):
|
||||
return long(data)
|
||||
|
||||
def convert_float(connection, field, data):
|
||||
return float(data)
|
||||
|
||||
encoders = {
|
||||
bool: escape_bool,
|
||||
int: escape_int,
|
||||
long: escape_long,
|
||||
float: escape_float,
|
||||
str: escape_string,
|
||||
unicode: escape_unicode,
|
||||
tuple: escape_sequence,
|
||||
list:escape_sequence,
|
||||
set:escape_sequence,
|
||||
dict:escape_dict,
|
||||
type(None):escape_None,
|
||||
datetime.date: escape_date,
|
||||
datetime.datetime : escape_datetime,
|
||||
datetime.timedelta : escape_timedelta,
|
||||
datetime.time : escape_time,
|
||||
time.struct_time : escape_struct_time,
|
||||
}
|
||||
bool: escape_bool,
|
||||
int: escape_int,
|
||||
long_type: escape_int,
|
||||
float: escape_float,
|
||||
str: escape_str,
|
||||
text_type: escape_unicode,
|
||||
tuple: escape_sequence,
|
||||
list: escape_sequence,
|
||||
set: escape_sequence,
|
||||
frozenset: escape_sequence,
|
||||
dict: escape_dict,
|
||||
bytearray: escape_bytes,
|
||||
type(None): escape_None,
|
||||
datetime.date: escape_date,
|
||||
datetime.datetime: escape_datetime,
|
||||
datetime.timedelta: escape_timedelta,
|
||||
datetime.time: escape_time,
|
||||
time.struct_time: escape_struct_time,
|
||||
Decimal: escape_object,
|
||||
}
|
||||
|
||||
if not PY2 or JYTHON or IRONPYTHON:
|
||||
encoders[bytes] = escape_bytes
|
||||
|
||||
decoders = {
|
||||
FIELD_TYPE.BIT: convert_bit,
|
||||
FIELD_TYPE.TINY: convert_int,
|
||||
FIELD_TYPE.SHORT: convert_int,
|
||||
FIELD_TYPE.LONG: convert_long,
|
||||
FIELD_TYPE.FLOAT: convert_float,
|
||||
FIELD_TYPE.DOUBLE: convert_float,
|
||||
FIELD_TYPE.DECIMAL: convert_float,
|
||||
FIELD_TYPE.NEWDECIMAL: convert_float,
|
||||
FIELD_TYPE.LONGLONG: convert_long,
|
||||
FIELD_TYPE.INT24: convert_int,
|
||||
FIELD_TYPE.YEAR: convert_int,
|
||||
FIELD_TYPE.TIMESTAMP: convert_mysql_timestamp,
|
||||
FIELD_TYPE.DATETIME: convert_datetime,
|
||||
FIELD_TYPE.TIME: convert_timedelta,
|
||||
FIELD_TYPE.DATE: convert_date,
|
||||
FIELD_TYPE.SET: convert_set,
|
||||
FIELD_TYPE.BLOB: convert_characters,
|
||||
FIELD_TYPE.TINY_BLOB: convert_characters,
|
||||
FIELD_TYPE.MEDIUM_BLOB: convert_characters,
|
||||
FIELD_TYPE.LONG_BLOB: convert_characters,
|
||||
FIELD_TYPE.STRING: convert_characters,
|
||||
FIELD_TYPE.VAR_STRING: convert_characters,
|
||||
FIELD_TYPE.VARCHAR: convert_characters,
|
||||
#FIELD_TYPE.BLOB: str,
|
||||
#FIELD_TYPE.STRING: str,
|
||||
#FIELD_TYPE.VAR_STRING: str,
|
||||
#FIELD_TYPE.VARCHAR: str
|
||||
}
|
||||
conversions = decoders # for MySQLdb compatibility
|
||||
FIELD_TYPE.BIT: convert_bit,
|
||||
FIELD_TYPE.TINY: int,
|
||||
FIELD_TYPE.SHORT: int,
|
||||
FIELD_TYPE.LONG: int,
|
||||
FIELD_TYPE.FLOAT: float,
|
||||
FIELD_TYPE.DOUBLE: float,
|
||||
FIELD_TYPE.LONGLONG: int,
|
||||
FIELD_TYPE.INT24: int,
|
||||
FIELD_TYPE.YEAR: int,
|
||||
FIELD_TYPE.TIMESTAMP: convert_mysql_timestamp,
|
||||
FIELD_TYPE.DATETIME: convert_datetime,
|
||||
FIELD_TYPE.TIME: convert_timedelta,
|
||||
FIELD_TYPE.DATE: convert_date,
|
||||
FIELD_TYPE.SET: convert_set,
|
||||
FIELD_TYPE.BLOB: through,
|
||||
FIELD_TYPE.TINY_BLOB: through,
|
||||
FIELD_TYPE.MEDIUM_BLOB: through,
|
||||
FIELD_TYPE.LONG_BLOB: through,
|
||||
FIELD_TYPE.STRING: through,
|
||||
FIELD_TYPE.VAR_STRING: through,
|
||||
FIELD_TYPE.VARCHAR: through,
|
||||
FIELD_TYPE.DECIMAL: Decimal,
|
||||
FIELD_TYPE.NEWDECIMAL: Decimal,
|
||||
}
|
||||
|
||||
try:
|
||||
# python version > 2.3
|
||||
from decimal import Decimal
|
||||
def convert_decimal(connection, field, data):
|
||||
data = data.decode(connection.charset)
|
||||
return Decimal(data)
|
||||
decoders[FIELD_TYPE.DECIMAL] = convert_decimal
|
||||
decoders[FIELD_TYPE.NEWDECIMAL] = convert_decimal
|
||||
|
||||
def escape_decimal(obj):
|
||||
return unicode(obj)
|
||||
encoders[Decimal] = escape_decimal
|
||||
|
||||
except ImportError:
|
||||
pass
|
||||
# for MySQLdb compatibility
|
||||
conversions = encoders.copy()
|
||||
conversions.update(decoders)
|
||||
Thing2Literal = escape_str
|
||||
|
||||
+304
-195
@@ -1,67 +1,82 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
import struct
|
||||
from __future__ import print_function, absolute_import
|
||||
from functools import partial
|
||||
import re
|
||||
import warnings
|
||||
|
||||
try:
|
||||
import cStringIO as StringIO
|
||||
except ImportError:
|
||||
import StringIO
|
||||
from ._compat import range_type, text_type, PY2
|
||||
from . import err
|
||||
|
||||
from err import Warning, Error, InterfaceError, DataError, \
|
||||
DatabaseError, OperationalError, IntegrityError, InternalError, \
|
||||
NotSupportedError, ProgrammingError
|
||||
|
||||
insert_values = re.compile(r'\svalues\s*(\(.+\))', re.IGNORECASE)
|
||||
#: Regular expression for :meth:`Cursor.executemany`.
|
||||
#: executemany only suports simple bulk insert.
|
||||
#: You can use it to load large dataset.
|
||||
RE_INSERT_VALUES = re.compile(
|
||||
r"\s*((?:INSERT|REPLACE)\s.+\sVALUES?\s+)" +
|
||||
r"(\(\s*(?:%s|%\(.+\)s)\s*(?:,\s*(?:%s|%\(.+\)s)\s*)*\))" +
|
||||
r"(\s*(?:ON DUPLICATE.*)?)\Z",
|
||||
re.IGNORECASE | re.DOTALL)
|
||||
|
||||
|
||||
class Cursor(object):
|
||||
'''
|
||||
"""
|
||||
This is the object you use to interact with the database.
|
||||
'''
|
||||
"""
|
||||
|
||||
#: Max stetement size which :meth:`executemany` generates.
|
||||
#:
|
||||
#: Max size of allowed statement is max_allowed_packet - packet_header_size.
|
||||
#: Default value of max_allowed_packet is 1048576.
|
||||
max_stmt_length = 1024000
|
||||
|
||||
_defer_warnings = False
|
||||
|
||||
def __init__(self, connection):
|
||||
'''
|
||||
"""
|
||||
Do not create an instance of a Cursor yourself. Call
|
||||
connections.Connection.cursor().
|
||||
'''
|
||||
from weakref import proxy
|
||||
self.connection = proxy(connection)
|
||||
"""
|
||||
self.connection = connection
|
||||
self.description = None
|
||||
self.rownumber = 0
|
||||
self.rowcount = -1
|
||||
self.arraysize = 1
|
||||
self._executed = None
|
||||
self.messages = []
|
||||
self.errorhandler = connection.errorhandler
|
||||
self._has_next = None
|
||||
self._rows = ()
|
||||
|
||||
def __del__(self):
|
||||
'''
|
||||
When this gets GC'd close it.
|
||||
'''
|
||||
self.close()
|
||||
self._result = None
|
||||
self._rows = None
|
||||
self._warnings_handled = False
|
||||
|
||||
def close(self):
|
||||
'''
|
||||
"""
|
||||
Closing a cursor just exhausts all remaining data.
|
||||
'''
|
||||
if not self.connection:
|
||||
"""
|
||||
conn = self.connection
|
||||
if conn is None:
|
||||
return
|
||||
try:
|
||||
while self.nextset():
|
||||
pass
|
||||
except:
|
||||
pass
|
||||
finally:
|
||||
self.connection = None
|
||||
|
||||
self.connection = None
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
del exc_info
|
||||
self.close()
|
||||
|
||||
def _get_db(self):
|
||||
if not self.connection:
|
||||
self.errorhandler(self, ProgrammingError, "cursor closed")
|
||||
raise err.ProgrammingError("Cursor closed")
|
||||
return self.connection
|
||||
|
||||
def _check_executed(self):
|
||||
if not self._executed:
|
||||
self.errorhandler(self, ProgrammingError, "execute() first")
|
||||
raise err.ProgrammingError("execute() first")
|
||||
|
||||
def _conv_row(self, row):
|
||||
return row
|
||||
|
||||
def setinputsizes(self, *args):
|
||||
"""Does nothing, required by DB API."""
|
||||
@@ -69,69 +84,152 @@ class Cursor(object):
|
||||
def setoutputsizes(self, *args):
|
||||
"""Does nothing, required by DB API."""
|
||||
|
||||
def nextset(self):
|
||||
''' Get the next query set '''
|
||||
if self._executed:
|
||||
self.fetchall()
|
||||
del self.messages[:]
|
||||
|
||||
if not self._has_next:
|
||||
def _nextset(self, unbuffered=False):
|
||||
"""Get the next query set"""
|
||||
conn = self._get_db()
|
||||
current_result = self._result
|
||||
# for unbuffered queries warnings are only available once whole result has been read
|
||||
if unbuffered:
|
||||
self._show_warnings()
|
||||
if current_result is None or current_result is not conn._result:
|
||||
return None
|
||||
connection = self._get_db()
|
||||
connection.next_result()
|
||||
if not current_result.has_next:
|
||||
return None
|
||||
conn.next_result(unbuffered=unbuffered)
|
||||
self._do_get_result()
|
||||
return True
|
||||
|
||||
def execute(self, query, args=None):
|
||||
''' Execute a query '''
|
||||
from sys import exc_info
|
||||
def nextset(self):
|
||||
return self._nextset(False)
|
||||
|
||||
def _ensure_bytes(self, x, encoding=None):
|
||||
if isinstance(x, text_type):
|
||||
x = x.encode(encoding)
|
||||
elif isinstance(x, (tuple, list)):
|
||||
x = type(x)(self._ensure_bytes(v, encoding=encoding) for v in x)
|
||||
return x
|
||||
|
||||
def _escape_args(self, args, conn):
|
||||
ensure_bytes = partial(self._ensure_bytes, encoding=conn.encoding)
|
||||
|
||||
if isinstance(args, (tuple, list)):
|
||||
if PY2:
|
||||
args = tuple(map(ensure_bytes, args))
|
||||
return tuple(conn.literal(arg) for arg in args)
|
||||
elif isinstance(args, dict):
|
||||
if PY2:
|
||||
args = dict((ensure_bytes(key), ensure_bytes(val)) for
|
||||
(key, val) in args.items())
|
||||
return dict((key, conn.literal(val)) for (key, val) in args.items())
|
||||
else:
|
||||
# If it's not a dictionary let's try escaping it anyways.
|
||||
# Worst case it will throw a Value error
|
||||
if PY2:
|
||||
args = ensure_bytes(args)
|
||||
return conn.escape(args)
|
||||
|
||||
def mogrify(self, query, args=None):
|
||||
"""
|
||||
Returns the exact string that is sent to the database by calling the
|
||||
execute() method.
|
||||
|
||||
This method follows the extension to the DB API 2.0 followed by Psycopg.
|
||||
"""
|
||||
conn = self._get_db()
|
||||
charset = conn.charset
|
||||
del self.messages[:]
|
||||
|
||||
# TODO: make sure that conn.escape is correct
|
||||
|
||||
if isinstance(query, unicode):
|
||||
query = query.encode(charset)
|
||||
if PY2: # Use bytes on Python 2 always
|
||||
query = self._ensure_bytes(query, encoding=conn.encoding)
|
||||
|
||||
if args is not None:
|
||||
if isinstance(args, tuple) or isinstance(args, list):
|
||||
escaped_args = tuple(conn.escape(arg) for arg in args)
|
||||
elif isinstance(args, dict):
|
||||
escaped_args = dict((key, conn.escape(val)) for (key, val) in args.items())
|
||||
else:
|
||||
#If it's not a dictionary let's try escaping it anyways.
|
||||
#Worst case it will throw a Value error
|
||||
escaped_args = conn.escape(args)
|
||||
query = query % self._escape_args(args, conn)
|
||||
|
||||
query = query % escaped_args
|
||||
return query
|
||||
|
||||
result = 0
|
||||
try:
|
||||
result = self._query(query)
|
||||
except:
|
||||
exc, value, tb = exc_info()
|
||||
del tb
|
||||
self.messages.append((exc,value))
|
||||
self.errorhandler(self, exc, value)
|
||||
def execute(self, query, args=None):
|
||||
"""Execute a query
|
||||
|
||||
:param str query: Query to execute.
|
||||
|
||||
:param args: parameters used with query. (optional)
|
||||
:type args: tuple, list or dict
|
||||
|
||||
:return: Number of affected rows
|
||||
:rtype: int
|
||||
|
||||
If args is a list or tuple, %s can be used as a placeholder in the query.
|
||||
If args is a dict, %(name)s can be used as a placeholder in the query.
|
||||
"""
|
||||
while self.nextset():
|
||||
pass
|
||||
|
||||
query = self.mogrify(query, args)
|
||||
|
||||
result = self._query(query)
|
||||
self._executed = query
|
||||
return result
|
||||
|
||||
def executemany(self, query, args):
|
||||
''' Run several data against one query '''
|
||||
del self.messages[:]
|
||||
#conn = self._get_db()
|
||||
# type: (str, list) -> int
|
||||
"""Run several data against one query
|
||||
|
||||
:param query: query to execute on server
|
||||
:param args: Sequence of sequences or mappings. It is used as parameter.
|
||||
:return: Number of rows affected, if any.
|
||||
|
||||
This method improves performance on multiple-row INSERT and
|
||||
REPLACE. Otherwise it is equivalent to looping over args with
|
||||
execute().
|
||||
"""
|
||||
if not args:
|
||||
return
|
||||
#charset = conn.charset
|
||||
#if isinstance(query, unicode):
|
||||
# query = query.encode(charset)
|
||||
|
||||
self.rowcount = sum([ self.execute(query, arg) for arg in args ])
|
||||
m = RE_INSERT_VALUES.match(query)
|
||||
if m:
|
||||
q_prefix = m.group(1) % ()
|
||||
q_values = m.group(2).rstrip()
|
||||
q_postfix = m.group(3) or ''
|
||||
assert q_values[0] == '(' and q_values[-1] == ')'
|
||||
return self._do_execute_many(q_prefix, q_values, q_postfix, args,
|
||||
self.max_stmt_length,
|
||||
self._get_db().encoding)
|
||||
|
||||
self.rowcount = sum(self.execute(query, arg) for arg in args)
|
||||
return self.rowcount
|
||||
|
||||
def _do_execute_many(self, prefix, values, postfix, args, max_stmt_length, encoding):
|
||||
conn = self._get_db()
|
||||
escape = self._escape_args
|
||||
if isinstance(prefix, text_type):
|
||||
prefix = prefix.encode(encoding)
|
||||
if PY2 and isinstance(values, text_type):
|
||||
values = values.encode(encoding)
|
||||
if isinstance(postfix, text_type):
|
||||
postfix = postfix.encode(encoding)
|
||||
sql = bytearray(prefix)
|
||||
args = iter(args)
|
||||
v = values % escape(next(args), conn)
|
||||
if isinstance(v, text_type):
|
||||
if PY2:
|
||||
v = v.encode(encoding)
|
||||
else:
|
||||
v = v.encode(encoding, 'surrogateescape')
|
||||
sql += v
|
||||
rows = 0
|
||||
for arg in args:
|
||||
v = values % escape(arg, conn)
|
||||
if isinstance(v, text_type):
|
||||
if PY2:
|
||||
v = v.encode(encoding)
|
||||
else:
|
||||
v = v.encode(encoding, 'surrogateescape')
|
||||
if len(sql) + len(v) + len(postfix) + 1 > max_stmt_length:
|
||||
rows += self.execute(sql + postfix)
|
||||
sql = bytearray(prefix)
|
||||
else:
|
||||
sql += b','
|
||||
sql += v
|
||||
rows += self.execute(sql + postfix)
|
||||
self.rowcount = rows
|
||||
return rows
|
||||
|
||||
def callproc(self, procname, args=()):
|
||||
"""Execute stored procedure procname with args
|
||||
@@ -164,23 +262,18 @@ class Cursor(object):
|
||||
conn = self._get_db()
|
||||
for index, arg in enumerate(args):
|
||||
q = "SET @_%s_%d=%s" % (procname, index, conn.escape(arg))
|
||||
if isinstance(q, unicode):
|
||||
q = q.encode(conn.charset)
|
||||
self._query(q)
|
||||
self.nextset()
|
||||
|
||||
q = "CALL %s(%s)" % (procname,
|
||||
','.join(['@_%s_%d' % (procname, i)
|
||||
for i in range(len(args))]))
|
||||
if isinstance(q, unicode):
|
||||
q = q.encode(conn.charset)
|
||||
for i in range_type(len(args))]))
|
||||
self._query(q)
|
||||
self._executed = q
|
||||
|
||||
return args
|
||||
|
||||
def fetchone(self):
|
||||
''' Fetch the next row '''
|
||||
"""Fetch the next row"""
|
||||
self._check_executed()
|
||||
if self._rows is None or self.rownumber >= len(self._rows):
|
||||
return None
|
||||
@@ -189,20 +282,20 @@ class Cursor(object):
|
||||
return result
|
||||
|
||||
def fetchmany(self, size=None):
|
||||
''' Fetch several rows '''
|
||||
"""Fetch several rows"""
|
||||
self._check_executed()
|
||||
if self._rows is None:
|
||||
return ()
|
||||
end = self.rownumber + (size or self.arraysize)
|
||||
result = self._rows[self.rownumber:end]
|
||||
if self._rows is None:
|
||||
return None
|
||||
self.rownumber = min(end, len(self._rows))
|
||||
return result
|
||||
|
||||
def fetchall(self):
|
||||
''' Fetch all the rows '''
|
||||
"""Fetch all the rows"""
|
||||
self._check_executed()
|
||||
if self._rows is None:
|
||||
return None
|
||||
return ()
|
||||
if self.rownumber:
|
||||
result = self._rows[self.rownumber:]
|
||||
else:
|
||||
@@ -217,11 +310,10 @@ class Cursor(object):
|
||||
elif mode == 'absolute':
|
||||
r = value
|
||||
else:
|
||||
self.errorhandler(self, ProgrammingError,
|
||||
"unknown scroll mode %s" % mode)
|
||||
raise err.ProgrammingError("unknown scroll mode %s" % mode)
|
||||
|
||||
if r < 0 or r >= len(self._rows):
|
||||
self.errorhandler(self, IndexError, "out of range")
|
||||
if not (0 <= r < len(self._rows)):
|
||||
raise IndexError("out of range")
|
||||
self.rownumber = r
|
||||
|
||||
def _query(self, q):
|
||||
@@ -233,92 +325,112 @@ class Cursor(object):
|
||||
|
||||
def _do_get_result(self):
|
||||
conn = self._get_db()
|
||||
self.rowcount = conn._result.affected_rows
|
||||
|
||||
self.rownumber = 0
|
||||
self.description = conn._result.description
|
||||
self.lastrowid = conn._result.insert_id
|
||||
self._rows = conn._result.rows
|
||||
self._has_next = conn._result.has_next
|
||||
self._result = result = conn._result
|
||||
|
||||
self.rowcount = result.affected_rows
|
||||
self.description = result.description
|
||||
self.lastrowid = result.insert_id
|
||||
self._rows = result.rows
|
||||
self._warnings_handled = False
|
||||
|
||||
if not self._defer_warnings:
|
||||
self._show_warnings()
|
||||
|
||||
def _show_warnings(self):
|
||||
if self._warnings_handled:
|
||||
return
|
||||
self._warnings_handled = True
|
||||
if self._result and (self._result.has_next or not self._result.warning_count):
|
||||
return
|
||||
ws = self._get_db().show_warnings()
|
||||
if ws is None:
|
||||
return
|
||||
for w in ws:
|
||||
msg = w[-1]
|
||||
if PY2:
|
||||
if isinstance(msg, unicode):
|
||||
msg = msg.encode('utf-8', 'replace')
|
||||
warnings.warn(err.Warning(*w[1:3]), stacklevel=4)
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.fetchone, None)
|
||||
|
||||
Warning = Warning
|
||||
Error = Error
|
||||
InterfaceError = InterfaceError
|
||||
DatabaseError = DatabaseError
|
||||
DataError = DataError
|
||||
OperationalError = OperationalError
|
||||
IntegrityError = IntegrityError
|
||||
InternalError = InternalError
|
||||
ProgrammingError = ProgrammingError
|
||||
NotSupportedError = NotSupportedError
|
||||
Warning = err.Warning
|
||||
Error = err.Error
|
||||
InterfaceError = err.InterfaceError
|
||||
DatabaseError = err.DatabaseError
|
||||
DataError = err.DataError
|
||||
OperationalError = err.OperationalError
|
||||
IntegrityError = err.IntegrityError
|
||||
InternalError = err.InternalError
|
||||
ProgrammingError = err.ProgrammingError
|
||||
NotSupportedError = err.NotSupportedError
|
||||
|
||||
class DictCursor(Cursor):
|
||||
|
||||
class DictCursorMixin(object):
|
||||
# You can override this to use OrderedDict or other dict-like types.
|
||||
dict_type = dict
|
||||
|
||||
def _do_get_result(self):
|
||||
super(DictCursorMixin, self)._do_get_result()
|
||||
fields = []
|
||||
if self.description:
|
||||
for f in self._result.fields:
|
||||
name = f.name
|
||||
if name in fields:
|
||||
name = f.table_name + '.' + name
|
||||
fields.append(name)
|
||||
self._fields = fields
|
||||
|
||||
if fields and self._rows:
|
||||
self._rows = [self._conv_row(r) for r in self._rows]
|
||||
|
||||
def _conv_row(self, row):
|
||||
if row is None:
|
||||
return None
|
||||
return self.dict_type(zip(self._fields, row))
|
||||
|
||||
|
||||
class DictCursor(DictCursorMixin, Cursor):
|
||||
"""A cursor which returns results as a dictionary"""
|
||||
|
||||
def execute(self, query, args=None):
|
||||
result = super(DictCursor, self).execute(query, args)
|
||||
if self.description:
|
||||
self._fields = [ field[0] for field in self.description ]
|
||||
return result
|
||||
|
||||
def fetchone(self):
|
||||
''' Fetch the next row '''
|
||||
self._check_executed()
|
||||
if self._rows is None or self.rownumber >= len(self._rows):
|
||||
return None
|
||||
result = dict(zip(self._fields, self._rows[self.rownumber]))
|
||||
self.rownumber += 1
|
||||
return result
|
||||
|
||||
def fetchmany(self, size=None):
|
||||
''' Fetch several rows '''
|
||||
self._check_executed()
|
||||
if self._rows is None:
|
||||
return None
|
||||
end = self.rownumber + (size or self.arraysize)
|
||||
result = [ dict(zip(self._fields, r)) for r in self._rows[self.rownumber:end] ]
|
||||
self.rownumber = min(end, len(self._rows))
|
||||
return tuple(result)
|
||||
|
||||
def fetchall(self):
|
||||
''' Fetch all the rows '''
|
||||
self._check_executed()
|
||||
if self._rows is None:
|
||||
return None
|
||||
if self.rownumber:
|
||||
result = [ dict(zip(self._fields, r)) for r in self._rows[self.rownumber:] ]
|
||||
else:
|
||||
result = [ dict(zip(self._fields, r)) for r in self._rows ]
|
||||
self.rownumber = len(self._rows)
|
||||
return tuple(result)
|
||||
|
||||
class SSCursor(Cursor):
|
||||
"""
|
||||
Unbuffered Cursor, mainly useful for queries that return a lot of data,
|
||||
or for connections to remote servers over a slow network.
|
||||
|
||||
|
||||
Instead of copying every row of data into a buffer, this will fetch
|
||||
rows as needed. The upside of this, is the client uses much less memory,
|
||||
and rows are returned much faster when traveling over a slow network,
|
||||
or if the result set is very big.
|
||||
|
||||
|
||||
There are limitations, though. The MySQL protocol doesn't support
|
||||
returning the total number of rows, so the only way to tell how many rows
|
||||
there are is to iterate over every row returned. Also, it currently isn't
|
||||
possible to scroll backwards, as only the current row is held in memory.
|
||||
"""
|
||||
|
||||
|
||||
_defer_warnings = True
|
||||
|
||||
def _conv_row(self, row):
|
||||
return row
|
||||
|
||||
def close(self):
|
||||
conn = self._get_db()
|
||||
conn._result._finish_unbuffered_query()
|
||||
|
||||
conn = self.connection
|
||||
if conn is None:
|
||||
return
|
||||
|
||||
if self._result is not None and self._result is conn._result:
|
||||
self._result._finish_unbuffered_query()
|
||||
|
||||
try:
|
||||
if self._has_next:
|
||||
while self.nextset(): pass
|
||||
except: pass
|
||||
while self.nextset():
|
||||
pass
|
||||
finally:
|
||||
self.connection = None
|
||||
|
||||
def _query(self, q):
|
||||
conn = self._get_db()
|
||||
@@ -326,38 +438,31 @@ class SSCursor(Cursor):
|
||||
conn.query(q, unbuffered=True)
|
||||
self._do_get_result()
|
||||
return self.rowcount
|
||||
|
||||
|
||||
def nextset(self):
|
||||
return self._nextset(unbuffered=True)
|
||||
|
||||
def read_next(self):
|
||||
""" Read next row """
|
||||
|
||||
conn = self._get_db()
|
||||
conn._result._read_rowdata_packet_unbuffered()
|
||||
return conn._result.rows
|
||||
|
||||
"""Read next row"""
|
||||
return self._conv_row(self._result._read_rowdata_packet_unbuffered())
|
||||
|
||||
def fetchone(self):
|
||||
""" Fetch next row """
|
||||
|
||||
"""Fetch next row"""
|
||||
self._check_executed()
|
||||
row = self.read_next()
|
||||
if row is None:
|
||||
self._show_warnings()
|
||||
return None
|
||||
self.rownumber += 1
|
||||
return row
|
||||
|
||||
|
||||
def fetchall(self):
|
||||
"""
|
||||
Fetch all, as per MySQLdb. Pretty useless for large queries, as
|
||||
it is buffered. See fetchall_unbuffered(), if you want an unbuffered
|
||||
generator version of this method.
|
||||
"""
|
||||
|
||||
rows = []
|
||||
while True:
|
||||
row = self.fetchone()
|
||||
if row is None:
|
||||
break
|
||||
rows.append(row)
|
||||
return tuple(rows)
|
||||
return list(self.fetchall_unbuffered())
|
||||
|
||||
def fetchall_unbuffered(self):
|
||||
"""
|
||||
@@ -365,46 +470,50 @@ class SSCursor(Cursor):
|
||||
however, it doesn't make sense to return everything in a list, as that
|
||||
would use ridiculous memory for large result sets.
|
||||
"""
|
||||
|
||||
row = self.fetchone()
|
||||
while row is not None:
|
||||
yield row
|
||||
row = self.fetchone()
|
||||
|
||||
return iter(self.fetchone, None)
|
||||
|
||||
def __iter__(self):
|
||||
return self.fetchall_unbuffered()
|
||||
|
||||
def fetchmany(self, size=None):
|
||||
""" Fetch many """
|
||||
|
||||
"""Fetch many"""
|
||||
self._check_executed()
|
||||
if size is None:
|
||||
size = self.arraysize
|
||||
|
||||
|
||||
rows = []
|
||||
for i in range(0, size):
|
||||
for i in range_type(size):
|
||||
row = self.read_next()
|
||||
if row is None:
|
||||
self._show_warnings()
|
||||
break
|
||||
rows.append(row)
|
||||
self.rownumber += 1
|
||||
return tuple(rows)
|
||||
|
||||
return rows
|
||||
|
||||
def scroll(self, value, mode='relative'):
|
||||
self._check_executed()
|
||||
if not mode == 'relative' and not mode == 'absolute':
|
||||
self.errorhandler(self, ProgrammingError,
|
||||
"unknown scroll mode %s" % mode)
|
||||
|
||||
|
||||
if mode == 'relative':
|
||||
if value < 0:
|
||||
self.errorhandler(self, NotSupportedError,
|
||||
"Backwards scrolling not supported by this cursor")
|
||||
|
||||
for i in range(0, value): self.read_next()
|
||||
raise err.NotSupportedError(
|
||||
"Backwards scrolling not supported by this cursor")
|
||||
|
||||
for _ in range_type(value):
|
||||
self.read_next()
|
||||
self.rownumber += value
|
||||
else:
|
||||
elif mode == 'absolute':
|
||||
if value < self.rownumber:
|
||||
self.errorhandler(self, NotSupportedError,
|
||||
raise err.NotSupportedError(
|
||||
"Backwards scrolling not supported by this cursor")
|
||||
|
||||
|
||||
end = value - self.rownumber
|
||||
for i in range(0, end): self.read_next()
|
||||
for _ in range_type(end):
|
||||
self.read_next()
|
||||
self.rownumber = value
|
||||
else:
|
||||
raise err.ProgrammingError("unknown scroll mode %s" % mode)
|
||||
|
||||
|
||||
class SSDictCursor(DictCursorMixin, SSCursor):
|
||||
"""An unbuffered cursor, which returns results as a dictionary"""
|
||||
|
||||
@@ -1,57 +1,39 @@
|
||||
import struct
|
||||
|
||||
from .constants import ER
|
||||
|
||||
try:
|
||||
StandardError, Warning
|
||||
except ImportError:
|
||||
try:
|
||||
from exceptions import StandardError, Warning
|
||||
except ImportError:
|
||||
import sys
|
||||
e = sys.modules['exceptions']
|
||||
StandardError = e.StandardError
|
||||
Warning = e.Warning
|
||||
|
||||
from constants import ER
|
||||
import sys
|
||||
|
||||
class MySQLError(StandardError):
|
||||
|
||||
class MySQLError(Exception):
|
||||
"""Exception related to operation with MySQL."""
|
||||
|
||||
|
||||
class Warning(Warning, MySQLError):
|
||||
|
||||
"""Exception raised for important warnings like data truncations
|
||||
while inserting, etc."""
|
||||
|
||||
class Error(MySQLError):
|
||||
|
||||
class Error(MySQLError):
|
||||
"""Exception that is the base class of all other error exceptions
|
||||
(not Warning)."""
|
||||
|
||||
|
||||
class InterfaceError(Error):
|
||||
|
||||
"""Exception raised for errors that are related to the database
|
||||
interface rather than the database itself."""
|
||||
|
||||
|
||||
class DatabaseError(Error):
|
||||
|
||||
"""Exception raised for errors that are related to the
|
||||
database."""
|
||||
|
||||
|
||||
class DataError(DatabaseError):
|
||||
|
||||
"""Exception raised for errors that are due to problems with the
|
||||
processed data like division by zero, numeric value out of range,
|
||||
etc."""
|
||||
|
||||
|
||||
class OperationalError(DatabaseError):
|
||||
|
||||
"""Exception raised for errors that are related to the database's
|
||||
operation and not necessarily under the control of the programmer,
|
||||
e.g. an unexpected disconnect occurs, the data source name is not
|
||||
@@ -60,28 +42,24 @@ class OperationalError(DatabaseError):
|
||||
|
||||
|
||||
class IntegrityError(DatabaseError):
|
||||
|
||||
"""Exception raised when the relational integrity of the database
|
||||
is affected, e.g. a foreign key check fails, duplicate key,
|
||||
etc."""
|
||||
|
||||
|
||||
class InternalError(DatabaseError):
|
||||
|
||||
"""Exception raised when the database encounters an internal
|
||||
error, e.g. the cursor is not valid anymore, the transaction is
|
||||
out of sync, etc."""
|
||||
|
||||
|
||||
class ProgrammingError(DatabaseError):
|
||||
|
||||
"""Exception raised for programming errors, e.g. table not found
|
||||
or already exists, syntax error in the SQL statement, wrong number
|
||||
of parameters specified, etc."""
|
||||
|
||||
|
||||
class NotSupportedError(DatabaseError):
|
||||
|
||||
"""Exception raised in case a method or database API was used
|
||||
which is not supported by the database, e.g. requesting a
|
||||
.rollback() on a connection that does not support transaction or
|
||||
@@ -90,10 +68,12 @@ class NotSupportedError(DatabaseError):
|
||||
|
||||
error_map = {}
|
||||
|
||||
|
||||
def _map_error(exc, *errors):
|
||||
for error in errors:
|
||||
error_map[error] = exc
|
||||
|
||||
|
||||
_map_error(ProgrammingError, ER.DB_CREATE_EXISTS, ER.SYNTAX_ERROR,
|
||||
ER.PARSE_ERROR, ER.NO_SUCH_TABLE, ER.WRONG_DB_NAME,
|
||||
ER.WRONG_TABLE_NAME, ER.FIELD_SPECIFIED_TWICE,
|
||||
@@ -104,44 +84,24 @@ _map_error(DataError, ER.WARN_DATA_TRUNCATED, ER.WARN_NULL_TO_NOTNULL,
|
||||
ER.DATA_TOO_LONG, ER.DATETIME_FUNCTION_OVERFLOW)
|
||||
_map_error(IntegrityError, ER.DUP_ENTRY, ER.NO_REFERENCED_ROW,
|
||||
ER.NO_REFERENCED_ROW_2, ER.ROW_IS_REFERENCED, ER.ROW_IS_REFERENCED_2,
|
||||
ER.CANNOT_ADD_FOREIGN)
|
||||
ER.CANNOT_ADD_FOREIGN, ER.BAD_NULL_ERROR)
|
||||
_map_error(NotSupportedError, ER.WARNING_NOT_COMPLETE_ROLLBACK,
|
||||
ER.NOT_SUPPORTED_YET, ER.FEATURE_DISABLED, ER.UNKNOWN_STORAGE_ENGINE)
|
||||
_map_error(OperationalError, ER.DBACCESS_DENIED_ERROR, ER.ACCESS_DENIED_ERROR,
|
||||
ER.TABLEACCESS_DENIED_ERROR, ER.COLUMNACCESS_DENIED_ERROR)
|
||||
_map_error(OperationalError, ER.DBACCESS_DENIED_ERROR, ER.ACCESS_DENIED_ERROR,
|
||||
ER.CON_COUNT_ERROR, ER.TABLEACCESS_DENIED_ERROR,
|
||||
ER.COLUMNACCESS_DENIED_ERROR)
|
||||
|
||||
|
||||
del _map_error, ER
|
||||
|
||||
|
||||
def _get_error_info(data):
|
||||
errno = struct.unpack('<h', data[1:3])[0]
|
||||
if sys.version_info[0] == 3:
|
||||
is_41 = data[3] == ord("#")
|
||||
else:
|
||||
is_41 = data[3] == "#"
|
||||
if is_41:
|
||||
# version 4.1
|
||||
sqlstate = data[4:9].decode("utf8")
|
||||
errorvalue = data[9:].decode("utf8")
|
||||
return (errno, sqlstate, errorvalue)
|
||||
else:
|
||||
# version 4.0
|
||||
return (errno, None, data[3:].decode("utf8"))
|
||||
|
||||
def _check_mysql_exception(errinfo):
|
||||
errno, sqlstate, errorvalue = errinfo
|
||||
errorclass = error_map.get(errno, None)
|
||||
if errorclass:
|
||||
raise errorclass, (errno,errorvalue)
|
||||
|
||||
# couldn't find the right error number
|
||||
raise InternalError, (errno, errorvalue)
|
||||
|
||||
def raise_mysql_exception(data):
|
||||
errinfo = _get_error_info(data)
|
||||
_check_mysql_exception(errinfo)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
errno = struct.unpack('<h', data[1:3])[0]
|
||||
is_41 = data[3:4] == b"#"
|
||||
if is_41:
|
||||
# client protocol 4.1
|
||||
errval = data[9:].decode('utf-8', 'replace')
|
||||
else:
|
||||
errval = data[3:].decode('utf-8', 'replace')
|
||||
errorclass = error_map.get(errno, InternalError)
|
||||
raise errorclass(errno, errval)
|
||||
|
||||
Executable
+20
@@ -0,0 +1,20 @@
|
||||
from ._compat import PY2
|
||||
|
||||
if PY2:
|
||||
import ConfigParser as configparser
|
||||
else:
|
||||
import configparser
|
||||
|
||||
|
||||
class Parser(configparser.RawConfigParser):
|
||||
|
||||
def __remove_quotes(self, value):
|
||||
quotes = ["'", "\""]
|
||||
for quote in quotes:
|
||||
if len(value) >= 2 and value[0] == value[-1] == quote:
|
||||
return value[1:-1]
|
||||
return value
|
||||
|
||||
def get(self, section, option):
|
||||
value = configparser.RawConfigParser.get(self, section, option)
|
||||
return self.__remove_quotes(value)
|
||||
@@ -1,13 +1,18 @@
|
||||
from pymysql.tests.test_issues import *
|
||||
from pymysql.tests.test_example import *
|
||||
from pymysql.tests.test_basic import *
|
||||
# Sorted by alphabetical order
|
||||
from pymysql.tests.test_DictCursor import *
|
||||
from pymysql.tests.test_SSCursor import *
|
||||
from pymysql.tests.test_basic import *
|
||||
from pymysql.tests.test_connection import *
|
||||
from pymysql.tests.test_converters import *
|
||||
from pymysql.tests.test_cursor import *
|
||||
from pymysql.tests.test_err import *
|
||||
from pymysql.tests.test_issues import *
|
||||
from pymysql.tests.test_load_local import *
|
||||
from pymysql.tests.test_nextset import *
|
||||
from pymysql.tests.test_optionfile import *
|
||||
|
||||
import sys
|
||||
if sys.version_info[0] == 2:
|
||||
# MySQLdb tests were designed for Python 3
|
||||
from pymysql.tests.thirdparty import *
|
||||
from pymysql.tests.thirdparty import *
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
unittest.main()
|
||||
import unittest2
|
||||
unittest2.main()
|
||||
|
||||
@@ -1,20 +1,86 @@
|
||||
import pymysql
|
||||
import unittest
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
|
||||
class PyMySQLTestCase(unittest.TestCase):
|
||||
# Edit this to suit your test environment.
|
||||
databases = [
|
||||
{"host":"localhost","user":"root",
|
||||
"passwd":"","db":"test_pymysql", "use_unicode": True},
|
||||
{"host":"localhost","user":"root","passwd":"","db":"test_pymysql2"}]
|
||||
import unittest2
|
||||
|
||||
import pymysql
|
||||
from .._compat import CPYTHON
|
||||
|
||||
|
||||
class PyMySQLTestCase(unittest2.TestCase):
|
||||
# You can specify your test environment creating a file named
|
||||
# "databases.json" or editing the `databases` variable below.
|
||||
fname = os.path.join(os.path.dirname(__file__), "databases.json")
|
||||
if os.path.exists(fname):
|
||||
with open(fname) as f:
|
||||
databases = json.load(f)
|
||||
else:
|
||||
databases = [
|
||||
{"host":"localhost","user":"root",
|
||||
"passwd":"","db":"test_pymysql", "use_unicode": True, 'local_infile': True},
|
||||
{"host":"localhost","user":"root","passwd":"","db":"test_pymysql2"}]
|
||||
|
||||
def mysql_server_is(self, conn, version_tuple):
|
||||
"""Return True if the given connection is on the version given or
|
||||
greater.
|
||||
|
||||
e.g.::
|
||||
|
||||
if self.mysql_server_is(conn, (5, 6, 4)):
|
||||
# do something for MySQL 5.6.4 and above
|
||||
"""
|
||||
server_version = conn.get_server_info()
|
||||
server_version_tuple = tuple(
|
||||
(int(dig) if dig is not None else 0)
|
||||
for dig in
|
||||
re.match(r'(\d+)\.(\d+)\.(\d+)', server_version).group(1, 2, 3)
|
||||
)
|
||||
return server_version_tuple >= version_tuple
|
||||
|
||||
def setUp(self):
|
||||
self.connections = []
|
||||
|
||||
for params in self.databases:
|
||||
self.connections.append(pymysql.connect(**params))
|
||||
self.addCleanup(self._teardown_connections)
|
||||
|
||||
def tearDown(self):
|
||||
def _teardown_connections(self):
|
||||
for connection in self.connections:
|
||||
connection.close()
|
||||
|
||||
def safe_create_table(self, connection, tablename, ddl, cleanup=True):
|
||||
"""create a table.
|
||||
|
||||
Ensures any existing version of that table is first dropped.
|
||||
|
||||
Also adds a cleanup rule to drop the table after the test
|
||||
completes.
|
||||
"""
|
||||
cursor = connection.cursor()
|
||||
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore")
|
||||
cursor.execute("drop table if exists `%s`" % (tablename,))
|
||||
cursor.execute(ddl)
|
||||
cursor.close()
|
||||
if cleanup:
|
||||
self.addCleanup(self.drop_table, connection, tablename)
|
||||
|
||||
def drop_table(self, connection, tablename):
|
||||
cursor = connection.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore")
|
||||
cursor.execute("drop table if exists `%s`" % (tablename,))
|
||||
cursor.close()
|
||||
|
||||
def safe_gc_collect(self):
|
||||
"""Ensure cycles are collected via gc.
|
||||
|
||||
Runs additional times on non-CPython platforms.
|
||||
|
||||
"""
|
||||
gc.collect()
|
||||
if not CPYTHON:
|
||||
gc.collect()
|
||||
|
||||
+22749
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,50 @@
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
5,6,
|
||||
7,8,
|
||||
1,2,
|
||||
3,4,
|
||||
@@ -2,54 +2,117 @@ from pymysql.tests import base
|
||||
import pymysql.cursors
|
||||
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
|
||||
class TestDictCursor(base.PyMySQLTestCase):
|
||||
bob = {'name': 'bob', 'age': 21, 'DOB': datetime.datetime(1990, 2, 6, 23, 4, 56)}
|
||||
jim = {'name': 'jim', 'age': 56, 'DOB': datetime.datetime(1955, 5, 9, 13, 12, 45)}
|
||||
fred = {'name': 'fred', 'age': 100, 'DOB': datetime.datetime(1911, 9, 12, 1, 1, 1)}
|
||||
|
||||
cursor_type = pymysql.cursors.DictCursor
|
||||
|
||||
def setUp(self):
|
||||
super(TestDictCursor, self).setUp()
|
||||
self.conn = conn = self.connections[0]
|
||||
c = conn.cursor(self.cursor_type)
|
||||
|
||||
# create a table ane some data to query
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists dictcursor")
|
||||
# include in filterwarnings since for unbuffered dict cursor warning for lack of table
|
||||
# will only be propagated at start of next execute() call
|
||||
c.execute("""CREATE TABLE dictcursor (name char(20), age int , DOB datetime)""")
|
||||
data = [("bob", 21, "1990-02-06 23:04:56"),
|
||||
("jim", 56, "1955-05-09 13:12:45"),
|
||||
("fred", 100, "1911-09-12 01:01:01")]
|
||||
c.executemany("insert into dictcursor values (%s,%s,%s)", data)
|
||||
|
||||
def tearDown(self):
|
||||
c = self.conn.cursor()
|
||||
c.execute("drop table dictcursor")
|
||||
super(TestDictCursor, self).tearDown()
|
||||
|
||||
def _ensure_cursor_expired(self, cursor):
|
||||
pass
|
||||
|
||||
def test_DictCursor(self):
|
||||
#all assert test compare to the structure as would come out from MySQLdb
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor(pymysql.cursors.DictCursor)
|
||||
# create a table ane some data to query
|
||||
c.execute("""CREATE TABLE dictcursor (name char(20), age int , DOB datetime)""")
|
||||
data = (("bob",21,"1990-02-06 23:04:56"),
|
||||
("jim",56,"1955-05-09 13:12:45"),
|
||||
("fred",100,"1911-09-12 01:01:01"))
|
||||
bob = {'name':'bob','age':21,'DOB':datetime.datetime(1990, 02, 6, 23, 04, 56)}
|
||||
jim = {'name':'jim','age':56,'DOB':datetime.datetime(1955, 05, 9, 13, 12, 45)}
|
||||
fred = {'name':'fred','age':100,'DOB':datetime.datetime(1911, 9, 12, 1, 1, 1)}
|
||||
try:
|
||||
c.executemany("insert into dictcursor values (%s,%s,%s)", data)
|
||||
# try an update which should return no rows
|
||||
c.execute("update dictcursor set age=20 where name='bob'")
|
||||
bob['age'] = 20
|
||||
# pull back the single row dict for bob and check
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
r = c.fetchone()
|
||||
self.assertEqual(bob,r,"fetchone via DictCursor failed")
|
||||
# same again, but via fetchall => tuple)
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
r = c.fetchall()
|
||||
self.assertEqual((bob,),r,"fetch a 1 row result via fetchall failed via DictCursor")
|
||||
# same test again but iterate over the
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
for r in c:
|
||||
self.assertEqual(bob, r,"fetch a 1 row result via iteration failed via DictCursor")
|
||||
# get all 3 row via fetchall
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = c.fetchall()
|
||||
self.assertEqual((bob,jim,fred), r, "fetchall failed via DictCursor")
|
||||
#same test again but do a list comprehension
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = [x for x in c]
|
||||
self.assertEqual([bob,jim,fred], r, "list comprehension failed via DictCursor")
|
||||
# get all 2 row via fetchmany
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = c.fetchmany(2)
|
||||
self.assertEqual((bob,jim), r, "fetchmany failed via DictCursor")
|
||||
finally:
|
||||
c.execute("drop table dictcursor")
|
||||
bob, jim, fred = self.bob.copy(), self.jim.copy(), self.fred.copy()
|
||||
#all assert test compare to the structure as would come out from MySQLdb
|
||||
conn = self.conn
|
||||
c = conn.cursor(self.cursor_type)
|
||||
|
||||
__all__ = ["TestDictCursor"]
|
||||
# try an update which should return no rows
|
||||
c.execute("update dictcursor set age=20 where name='bob'")
|
||||
bob['age'] = 20
|
||||
# pull back the single row dict for bob and check
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
r = c.fetchone()
|
||||
self.assertEqual(bob, r, "fetchone via DictCursor failed")
|
||||
self._ensure_cursor_expired(c)
|
||||
|
||||
# same again, but via fetchall => tuple)
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
r = c.fetchall()
|
||||
self.assertEqual([bob], r, "fetch a 1 row result via fetchall failed via DictCursor")
|
||||
# same test again but iterate over the
|
||||
c.execute("SELECT * from dictcursor where name='bob'")
|
||||
for r in c:
|
||||
self.assertEqual(bob, r, "fetch a 1 row result via iteration failed via DictCursor")
|
||||
# get all 3 row via fetchall
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = c.fetchall()
|
||||
self.assertEqual([bob,jim,fred], r, "fetchall failed via DictCursor")
|
||||
#same test again but do a list comprehension
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = list(c)
|
||||
self.assertEqual([bob,jim,fred], r, "DictCursor should be iterable")
|
||||
# get all 2 row via fetchmany
|
||||
c.execute("SELECT * from dictcursor")
|
||||
r = c.fetchmany(2)
|
||||
self.assertEqual([bob, jim], r, "fetchmany failed via DictCursor")
|
||||
self._ensure_cursor_expired(c)
|
||||
|
||||
def test_custom_dict(self):
|
||||
class MyDict(dict): pass
|
||||
|
||||
class MyDictCursor(self.cursor_type):
|
||||
dict_type = MyDict
|
||||
|
||||
keys = ['name', 'age', 'DOB']
|
||||
bob = MyDict([(k, self.bob[k]) for k in keys])
|
||||
jim = MyDict([(k, self.jim[k]) for k in keys])
|
||||
fred = MyDict([(k, self.fred[k]) for k in keys])
|
||||
|
||||
cur = self.conn.cursor(MyDictCursor)
|
||||
cur.execute("SELECT * FROM dictcursor WHERE name='bob'")
|
||||
r = cur.fetchone()
|
||||
self.assertEqual(bob, r, "fetchone() returns MyDictCursor")
|
||||
self._ensure_cursor_expired(cur)
|
||||
|
||||
cur.execute("SELECT * FROM dictcursor")
|
||||
r = cur.fetchall()
|
||||
self.assertEqual([bob, jim, fred], r,
|
||||
"fetchall failed via MyDictCursor")
|
||||
|
||||
cur.execute("SELECT * FROM dictcursor")
|
||||
r = list(cur)
|
||||
self.assertEqual([bob, jim, fred], r,
|
||||
"list failed via MyDictCursor")
|
||||
|
||||
cur.execute("SELECT * FROM dictcursor")
|
||||
r = cur.fetchmany(2)
|
||||
self.assertEqual([bob, jim], r,
|
||||
"list failed via MyDictCursor")
|
||||
self._ensure_cursor_expired(cur)
|
||||
|
||||
|
||||
class TestSSDictCursor(TestDictCursor):
|
||||
cursor_type = pymysql.cursors.SSDictCursor
|
||||
|
||||
def _ensure_cursor_expired(self, cursor):
|
||||
list(cursor.fetchall_unbuffered())
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
|
||||
@@ -3,7 +3,7 @@ import sys
|
||||
try:
|
||||
from pymysql.tests import base
|
||||
import pymysql.cursors
|
||||
except:
|
||||
except Exception:
|
||||
# For local testing from top-level directory, without installing
|
||||
sys.path.append('../pymysql')
|
||||
from pymysql.tests import base
|
||||
@@ -12,7 +12,7 @@ except:
|
||||
class TestSSCursor(base.PyMySQLTestCase):
|
||||
def test_SSCursor(self):
|
||||
affected_rows = 18446744073709551615
|
||||
|
||||
|
||||
conn = self.connections[0]
|
||||
data = [
|
||||
('America', '', 'America/Jamaica'),
|
||||
@@ -25,22 +25,23 @@ class TestSSCursor(base.PyMySQLTestCase):
|
||||
('America', '', 'America/Costa_Rica'),
|
||||
('America', '', 'America/Denver'),
|
||||
('America', '', 'America/Detroit'),]
|
||||
|
||||
|
||||
try:
|
||||
cursor = conn.cursor(pymysql.cursors.SSCursor)
|
||||
|
||||
|
||||
# Create table
|
||||
cursor.execute(('CREATE TABLE tz_data ('
|
||||
'region VARCHAR(64),'
|
||||
'zone VARCHAR(64),'
|
||||
'name VARCHAR(64))'))
|
||||
|
||||
|
||||
conn.begin()
|
||||
# Test INSERT
|
||||
for i in data:
|
||||
cursor.execute('INSERT INTO tz_data VALUES (%s, %s, %s)', i)
|
||||
self.assertEqual(conn.affected_rows(), 1, 'affected_rows does not match')
|
||||
conn.commit()
|
||||
|
||||
|
||||
# Test fetchone()
|
||||
iter = 0
|
||||
cursor.execute('SELECT * FROM tz_data')
|
||||
@@ -49,46 +50,55 @@ class TestSSCursor(base.PyMySQLTestCase):
|
||||
if row is None:
|
||||
break
|
||||
iter += 1
|
||||
|
||||
|
||||
# Test cursor.rowcount
|
||||
self.assertEqual(cursor.rowcount, affected_rows,
|
||||
'cursor.rowcount != %s' % (str(affected_rows)))
|
||||
|
||||
|
||||
# Test cursor.rownumber
|
||||
self.assertEqual(cursor.rownumber, iter,
|
||||
'cursor.rowcount != %s' % (str(iter)))
|
||||
|
||||
|
||||
# Test row came out the same as it went in
|
||||
self.assertEqual((row in data), True,
|
||||
'Row not found in source data')
|
||||
|
||||
|
||||
# Test fetchall
|
||||
cursor.execute('SELECT * FROM tz_data')
|
||||
self.assertEqual(len(cursor.fetchall()), len(data),
|
||||
'fetchall failed. Number of rows does not match')
|
||||
|
||||
|
||||
# Test fetchmany
|
||||
cursor.execute('SELECT * FROM tz_data')
|
||||
self.assertEqual(len(cursor.fetchmany(2)), 2,
|
||||
'fetchmany failed. Number of rows does not match')
|
||||
|
||||
|
||||
# So MySQLdb won't throw "Commands out of sync"
|
||||
while True:
|
||||
res = cursor.fetchone()
|
||||
if res is None:
|
||||
break
|
||||
|
||||
|
||||
# Test update, affected_rows()
|
||||
cursor.execute('UPDATE tz_data SET zone = %s', ['Foo'])
|
||||
conn.commit()
|
||||
self.assertEqual(cursor.rowcount, len(data),
|
||||
'Update failed. affected_rows != %s' % (str(len(data))))
|
||||
|
||||
|
||||
# Test executemany
|
||||
cursor.executemany('INSERT INTO tz_data VALUES (%s, %s, %s)', data)
|
||||
self.assertEqual(cursor.rowcount, len(data),
|
||||
'executemany failed. cursor.rowcount != %s' % (str(len(data))))
|
||||
|
||||
|
||||
# Test multiple datasets
|
||||
cursor.execute('SELECT 1; SELECT 2; SELECT 3')
|
||||
self.assertListEqual(list(cursor), [(1, )])
|
||||
self.assertTrue(cursor.nextset())
|
||||
self.assertListEqual(list(cursor), [(2, )])
|
||||
self.assertTrue(cursor.nextset())
|
||||
self.assertListEqual(list(cursor), [(3, )])
|
||||
self.assertFalse(cursor.nextset())
|
||||
|
||||
finally:
|
||||
cursor.execute('DROP TABLE tz_data')
|
||||
cursor.close()
|
||||
|
||||
@@ -1,8 +1,19 @@
|
||||
from pymysql.tests import base
|
||||
from pymysql import util
|
||||
|
||||
import time
|
||||
# coding: utf-8
|
||||
import datetime
|
||||
import json
|
||||
import time
|
||||
import warnings
|
||||
|
||||
from unittest2 import SkipTest
|
||||
|
||||
from pymysql import util
|
||||
import pymysql.cursors
|
||||
from pymysql.tests import base
|
||||
from pymysql.err import ProgrammingError
|
||||
|
||||
|
||||
__all__ = ["TestConversion", "TestCursor", "TestBulkInserts"]
|
||||
|
||||
|
||||
class TestConversion(base.PyMySQLTestCase):
|
||||
def test_datatypes(self):
|
||||
@@ -12,16 +23,13 @@ class TestConversion(base.PyMySQLTestCase):
|
||||
c.execute("create table test_datatypes (b bit, i int, l bigint, f real, s varchar(32), u varchar(32), bb blob, d date, dt datetime, ts timestamp, td time, t time, st datetime)")
|
||||
try:
|
||||
# insert values
|
||||
v = (True, -3, 123456789012, 5.7, "hello'\" world", u"Espa\xc3\xb1ol", "binary\x00data".encode(conn.charset), datetime.date(1988,2,2), datetime.datetime.now(), datetime.timedelta(5,6), datetime.time(16,32), time.localtime())
|
||||
|
||||
v = (True, -3, 123456789012, 5.7, "hello'\" world", u"Espa\xc3\xb1ol", "binary\x00data".encode(conn.charset), datetime.date(1988,2,2), datetime.datetime(2014, 5, 15, 7, 45, 57), datetime.timedelta(5,6), datetime.time(16,32), time.localtime())
|
||||
c.execute("insert into test_datatypes (b,i,l,f,s,u,bb,d,dt,td,t,st) values (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)", v)
|
||||
c.execute("select b,i,l,f,s,u,bb,d,dt,td,t,st from test_datatypes")
|
||||
r = c.fetchone()
|
||||
self.assertEqual(util.int2byte(1), r[0])
|
||||
self.assertEqual(v[1:8], r[1:8])
|
||||
# mysql throws away microseconds so we need to check datetimes
|
||||
# specially. additionally times are turned into timedeltas.
|
||||
self.assertEqual(datetime.datetime(*v[8].timetuple()[:6]), r[8])
|
||||
self.assertEqual(v[9], r[9]) # just timedeltas
|
||||
self.assertEqual(v[1:10], r[1:10])
|
||||
self.assertEqual(datetime.timedelta(0, 60 * (v[10].hour * 60 + v[10].minute)), r[10])
|
||||
self.assertEqual(datetime.datetime(*v[-1][:6]), r[-1])
|
||||
|
||||
@@ -35,11 +43,15 @@ class TestConversion(base.PyMySQLTestCase):
|
||||
|
||||
c.execute("delete from test_datatypes")
|
||||
|
||||
# check sequence type
|
||||
c.execute("insert into test_datatypes (i, l) values (2,4), (6,8), (10,12)")
|
||||
c.execute("select l from test_datatypes where i in %s order by i", ((2,6),))
|
||||
r = c.fetchall()
|
||||
self.assertEqual(((4,),(8,)), r)
|
||||
# check sequences type
|
||||
for seq_type in (tuple, list, set, frozenset):
|
||||
c.execute("insert into test_datatypes (i, l) values (2,4), (6,8), (10,12)")
|
||||
seq = seq_type([2,6])
|
||||
c.execute("select l from test_datatypes where i in %s order by i", (seq,))
|
||||
r = c.fetchall()
|
||||
self.assertEqual(((4,),(8,)), r)
|
||||
c.execute("delete from test_datatypes")
|
||||
|
||||
finally:
|
||||
c.execute("drop table test_datatypes")
|
||||
|
||||
@@ -79,20 +91,18 @@ class TestConversion(base.PyMySQLTestCase):
|
||||
finally:
|
||||
c.execute("drop table test_dict")
|
||||
|
||||
|
||||
def test_big_blob(self):
|
||||
""" test tons of data """
|
||||
def test_blob(self):
|
||||
"""test binary data"""
|
||||
data = bytes(bytearray(range(256)) * 4)
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
c.execute("create table test_big_blob (b blob)")
|
||||
try:
|
||||
data = "pymysql" * 1024
|
||||
c.execute("insert into test_big_blob (b) values (%s)", (data,))
|
||||
c.execute("select b from test_big_blob")
|
||||
self.assertEqual(data.encode(conn.charset), c.fetchone()[0])
|
||||
finally:
|
||||
c.execute("drop table test_big_blob")
|
||||
|
||||
self.safe_create_table(
|
||||
conn, "test_blob", "create table test_blob (b blob)")
|
||||
|
||||
with conn.cursor() as c:
|
||||
c.execute("insert into test_blob (b) values (%s)", (data,))
|
||||
c.execute("select b from test_blob")
|
||||
self.assertEqual(data, c.fetchone()[0])
|
||||
|
||||
def test_untyped(self):
|
||||
""" test conversion of null, empty string """
|
||||
conn = self.connections[0]
|
||||
@@ -101,17 +111,40 @@ class TestConversion(base.PyMySQLTestCase):
|
||||
self.assertEqual((None,u''), c.fetchone())
|
||||
c.execute("select '',null")
|
||||
self.assertEqual((u'',None), c.fetchone())
|
||||
|
||||
def test_datetime(self):
|
||||
""" test conversion of null, empty string """
|
||||
|
||||
def test_timedelta(self):
|
||||
""" test timedelta conversion """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
c.execute("select time('12:30'), time('23:12:59'), time('23:12:59.05100')")
|
||||
c.execute("select time('12:30'), time('23:12:59'), time('23:12:59.05100'), time('-12:30'), time('-23:12:59'), time('-23:12:59.05100'), time('-00:30')")
|
||||
self.assertEqual((datetime.timedelta(0, 45000),
|
||||
datetime.timedelta(0, 83579),
|
||||
datetime.timedelta(0, 83579, 51000)),
|
||||
datetime.timedelta(0, 83579, 51000),
|
||||
-datetime.timedelta(0, 45000),
|
||||
-datetime.timedelta(0, 83579),
|
||||
-datetime.timedelta(0, 83579, 51000),
|
||||
-datetime.timedelta(0, 1800)),
|
||||
c.fetchone())
|
||||
|
||||
def test_datetime_microseconds(self):
|
||||
""" test datetime conversion w microseconds"""
|
||||
|
||||
conn = self.connections[0]
|
||||
if not self.mysql_server_is(conn, (5, 6, 4)):
|
||||
raise SkipTest("target backend does not support microseconds")
|
||||
c = conn.cursor()
|
||||
dt = datetime.datetime(2013, 11, 12, 9, 9, 9, 123450)
|
||||
c.execute("create table test_datetime (id int, ts datetime(6))")
|
||||
try:
|
||||
c.execute(
|
||||
"insert into test_datetime values (%s, %s)",
|
||||
(1, dt)
|
||||
)
|
||||
c.execute("select ts from test_datetime")
|
||||
self.assertEqual((dt,), c.fetchone())
|
||||
finally:
|
||||
c.execute("drop table test_datetime")
|
||||
|
||||
|
||||
class TestCursor(base.PyMySQLTestCase):
|
||||
# this test case does not work quite right yet, however,
|
||||
@@ -185,7 +218,7 @@ class TestCursor(base.PyMySQLTestCase):
|
||||
c = conn.cursor()
|
||||
try:
|
||||
c.execute('create table test_aggregates (i integer)')
|
||||
for i in xrange(0, 10):
|
||||
for i in range(0, 10):
|
||||
c.execute('insert into test_aggregates (i) values (%s)', (i,))
|
||||
c.execute('select sum(i) from test_aggregates')
|
||||
r, = c.fetchone()
|
||||
@@ -197,17 +230,150 @@ class TestCursor(base.PyMySQLTestCase):
|
||||
""" test a single tuple """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
try:
|
||||
c.execute("create table mystuff (id integer primary key)")
|
||||
c.execute("insert into mystuff (id) values (1)")
|
||||
c.execute("insert into mystuff (id) values (2)")
|
||||
c.execute("select id from mystuff where id in %s", ((1,),))
|
||||
self.assertEqual([(1,)], list(c.fetchall()))
|
||||
finally:
|
||||
c.execute("drop table mystuff")
|
||||
self.safe_create_table(
|
||||
conn, 'mystuff',
|
||||
"create table mystuff (id integer primary key)")
|
||||
c.execute("insert into mystuff (id) values (1)")
|
||||
c.execute("insert into mystuff (id) values (2)")
|
||||
c.execute("select id from mystuff where id in %s", ((1,),))
|
||||
self.assertEqual([(1,)], list(c.fetchall()))
|
||||
c.close()
|
||||
|
||||
__all__ = ["TestConversion","TestCursor"]
|
||||
def test_json(self):
|
||||
args = self.databases[0].copy()
|
||||
args["charset"] = "utf8mb4"
|
||||
conn = pymysql.connect(**args)
|
||||
if not self.mysql_server_is(conn, (5, 7, 0)):
|
||||
raise SkipTest("JSON type is not supported on MySQL <= 5.6")
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
unittest.main()
|
||||
self.safe_create_table(conn, "test_json", """\
|
||||
create table test_json (
|
||||
id int not null,
|
||||
json JSON not null,
|
||||
primary key (id)
|
||||
);""")
|
||||
cur = conn.cursor()
|
||||
|
||||
json_str = u'{"hello": "こんにちは"}'
|
||||
cur.execute("INSERT INTO test_json (id, `json`) values (42, %s)", (json_str,))
|
||||
cur.execute("SELECT `json` from `test_json` WHERE `id`=42")
|
||||
res = cur.fetchone()[0]
|
||||
self.assertEqual(json.loads(res), json.loads(json_str))
|
||||
|
||||
cur.execute("SELECT CAST(%s AS JSON) AS x", (json_str,))
|
||||
res = cur.fetchone()[0]
|
||||
self.assertEqual(json.loads(res), json.loads(json_str))
|
||||
|
||||
|
||||
class TestBulkInserts(base.PyMySQLTestCase):
|
||||
|
||||
cursor_type = pymysql.cursors.DictCursor
|
||||
|
||||
def setUp(self):
|
||||
super(TestBulkInserts, self).setUp()
|
||||
self.conn = conn = self.connections[0]
|
||||
c = conn.cursor(self.cursor_type)
|
||||
|
||||
# create a table ane some data to query
|
||||
self.safe_create_table(conn, 'bulkinsert', """\
|
||||
CREATE TABLE bulkinsert
|
||||
(
|
||||
id int(11),
|
||||
name char(20),
|
||||
age int,
|
||||
height int,
|
||||
PRIMARY KEY (id)
|
||||
)
|
||||
""")
|
||||
|
||||
def _verify_records(self, data):
|
||||
conn = self.connections[0]
|
||||
cursor = conn.cursor()
|
||||
cursor.execute("SELECT id, name, age, height from bulkinsert")
|
||||
result = cursor.fetchall()
|
||||
self.assertEqual(sorted(data), sorted(result))
|
||||
|
||||
def test_bulk_insert(self):
|
||||
conn = self.connections[0]
|
||||
cursor = conn.cursor()
|
||||
|
||||
data = [(0, "bob", 21, 123), (1, "jim", 56, 45), (2, "fred", 100, 180)]
|
||||
cursor.executemany("insert into bulkinsert (id, name, age, height) "
|
||||
"values (%s,%s,%s,%s)", data)
|
||||
self.assertEqual(
|
||||
cursor._last_executed, bytearray(
|
||||
b"insert into bulkinsert (id, name, age, height) values "
|
||||
b"(0,'bob',21,123),(1,'jim',56,45),(2,'fred',100,180)"))
|
||||
cursor.execute('commit')
|
||||
self._verify_records(data)
|
||||
|
||||
def test_bulk_insert_multiline_statement(self):
|
||||
conn = self.connections[0]
|
||||
cursor = conn.cursor()
|
||||
data = [(0, "bob", 21, 123), (1, "jim", 56, 45), (2, "fred", 100, 180)]
|
||||
cursor.executemany("""insert
|
||||
into bulkinsert (id, name,
|
||||
age, height)
|
||||
values (%s,
|
||||
%s , %s,
|
||||
%s )
|
||||
""", data)
|
||||
self.assertEqual(cursor._last_executed.strip(), bytearray(b"""insert
|
||||
into bulkinsert (id, name,
|
||||
age, height)
|
||||
values (0,
|
||||
'bob' , 21,
|
||||
123 ),(1,
|
||||
'jim' , 56,
|
||||
45 ),(2,
|
||||
'fred' , 100,
|
||||
180 )"""))
|
||||
cursor.execute('commit')
|
||||
self._verify_records(data)
|
||||
|
||||
def test_bulk_insert_single_record(self):
|
||||
conn = self.connections[0]
|
||||
cursor = conn.cursor()
|
||||
data = [(0, "bob", 21, 123)]
|
||||
cursor.executemany("insert into bulkinsert (id, name, age, height) "
|
||||
"values (%s,%s,%s,%s)", data)
|
||||
cursor.execute('commit')
|
||||
self._verify_records(data)
|
||||
|
||||
def test_issue_288(self):
|
||||
"""executemany should work with "insert ... on update" """
|
||||
conn = self.connections[0]
|
||||
cursor = conn.cursor()
|
||||
data = [(0, "bob", 21, 123), (1, "jim", 56, 45), (2, "fred", 100, 180)]
|
||||
cursor.executemany("""insert
|
||||
into bulkinsert (id, name,
|
||||
age, height)
|
||||
values (%s,
|
||||
%s , %s,
|
||||
%s ) on duplicate key update
|
||||
age = values(age)
|
||||
""", data)
|
||||
self.assertEqual(cursor._last_executed.strip(), bytearray(b"""insert
|
||||
into bulkinsert (id, name,
|
||||
age, height)
|
||||
values (0,
|
||||
'bob' , 21,
|
||||
123 ),(1,
|
||||
'jim' , 56,
|
||||
45 ),(2,
|
||||
'fred' , 100,
|
||||
180 ) on duplicate key update
|
||||
age = values(age)"""))
|
||||
cursor.execute('commit')
|
||||
self._verify_records(data)
|
||||
|
||||
def test_warnings(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
with warnings.catch_warnings(record=True) as ws:
|
||||
warnings.simplefilter("always")
|
||||
cur.execute("drop table if exists no_exists_table")
|
||||
self.assertEqual(len(ws), 1)
|
||||
self.assertEqual(ws[0].category, pymysql.Warning)
|
||||
if u"no_exists_table" not in str(ws[0].message):
|
||||
self.fail("'no_exists_table' not in %s" % (str(ws[0].message),))
|
||||
|
||||
+576
@@ -0,0 +1,576 @@
|
||||
import datetime
|
||||
import sys
|
||||
import time
|
||||
import unittest2
|
||||
import pymysql
|
||||
from pymysql.tests import base
|
||||
from pymysql._compat import text_type
|
||||
|
||||
|
||||
class TempUser:
|
||||
def __init__(self, c, user, db, auth=None, authdata=None, password=None):
|
||||
self._c = c
|
||||
self._user = user
|
||||
self._db = db
|
||||
create = "CREATE USER " + user
|
||||
if password is not None:
|
||||
create += " IDENTIFIED BY '%s'" % password
|
||||
elif auth is not None:
|
||||
create += " IDENTIFIED WITH %s" % auth
|
||||
if authdata is not None:
|
||||
create += " AS '%s'" % authdata
|
||||
try:
|
||||
c.execute(create)
|
||||
self._created = True
|
||||
except pymysql.err.InternalError:
|
||||
# already exists - TODO need to check the same plugin applies
|
||||
self._created = False
|
||||
try:
|
||||
c.execute("GRANT SELECT ON %s.* TO %s" % (db, user))
|
||||
self._grant = True
|
||||
except pymysql.err.InternalError:
|
||||
self._grant = False
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
if self._grant:
|
||||
self._c.execute("REVOKE SELECT ON %s.* FROM %s" % (self._db, self._user))
|
||||
if self._created:
|
||||
self._c.execute("DROP USER %s" % self._user)
|
||||
|
||||
|
||||
class TestAuthentication(base.PyMySQLTestCase):
|
||||
|
||||
socket_auth = False
|
||||
socket_found = False
|
||||
two_questions_found = False
|
||||
three_attempts_found = False
|
||||
pam_found = False
|
||||
mysql_old_password_found = False
|
||||
sha256_password_found = False
|
||||
|
||||
import os
|
||||
osuser = os.environ.get('USER')
|
||||
|
||||
# socket auth requires the current user and for the connection to be a socket
|
||||
# rest do grants @localhost due to incomplete logic - TODO change to @% then
|
||||
db = base.PyMySQLTestCase.databases[0].copy()
|
||||
|
||||
socket_auth = db.get('unix_socket') is not None \
|
||||
and db.get('host') in ('localhost', '127.0.0.1')
|
||||
|
||||
cur = pymysql.connect(**db).cursor()
|
||||
del db['user']
|
||||
cur.execute("SHOW PLUGINS")
|
||||
for r in cur:
|
||||
if (r[1], r[2]) != (u'ACTIVE', u'AUTHENTICATION'):
|
||||
continue
|
||||
if r[3] == u'auth_socket.so':
|
||||
socket_plugin_name = r[0]
|
||||
socket_found = True
|
||||
elif r[3] == u'dialog_examples.so':
|
||||
if r[0] == 'two_questions':
|
||||
two_questions_found = True
|
||||
elif r[0] == 'three_attempts':
|
||||
three_attempts_found = True
|
||||
elif r[0] == u'pam':
|
||||
pam_found = True
|
||||
pam_plugin_name = r[3].split('.')[0]
|
||||
if pam_plugin_name == 'auth_pam':
|
||||
pam_plugin_name = 'pam'
|
||||
# MySQL: authentication_pam
|
||||
# https://dev.mysql.com/doc/refman/5.5/en/pam-authentication-plugin.html
|
||||
|
||||
# MariaDB: pam
|
||||
# https://mariadb.com/kb/en/mariadb/pam-authentication-plugin/
|
||||
|
||||
# Names differ but functionality is close
|
||||
elif r[0] == u'mysql_old_password':
|
||||
mysql_old_password_found = True
|
||||
elif r[0] == u'sha256_password':
|
||||
sha256_password_found = True
|
||||
#else:
|
||||
# print("plugin: %r" % r[0])
|
||||
|
||||
def test_plugin(self):
|
||||
# Bit of an assumption that the current user is a native password
|
||||
self.assertEqual('mysql_native_password', self.connections[0]._auth_plugin_name)
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipIf(socket_found, "socket plugin already installed")
|
||||
def testSocketAuthInstallPlugin(self):
|
||||
# needs plugin. lets install it.
|
||||
cur = self.connections[0].cursor()
|
||||
try:
|
||||
cur.execute("install plugin auth_socket soname 'auth_socket.so'")
|
||||
TestAuthentication.socket_found = True
|
||||
self.socket_plugin_name = 'auth_socket'
|
||||
self.realtestSocketAuth()
|
||||
except pymysql.err.InternalError:
|
||||
try:
|
||||
cur.execute("install soname 'auth_socket'")
|
||||
TestAuthentication.socket_found = True
|
||||
self.socket_plugin_name = 'unix_socket'
|
||||
self.realtestSocketAuth()
|
||||
except pymysql.err.InternalError:
|
||||
TestAuthentication.socket_found = False
|
||||
raise unittest2.SkipTest('we couldn\'t install the socket plugin')
|
||||
finally:
|
||||
if TestAuthentication.socket_found:
|
||||
cur.execute("uninstall plugin %s" % self.socket_plugin_name)
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(socket_found, "no socket plugin")
|
||||
def testSocketAuth(self):
|
||||
self.realtestSocketAuth()
|
||||
|
||||
def realtestSocketAuth(self):
|
||||
with TempUser(self.connections[0].cursor(), TestAuthentication.osuser + '@localhost',
|
||||
self.databases[0]['db'], self.socket_plugin_name) as u:
|
||||
c = pymysql.connect(user=TestAuthentication.osuser, **self.db)
|
||||
|
||||
class Dialog(object):
|
||||
fail=False
|
||||
|
||||
def __init__(self, con):
|
||||
self.fail=TestAuthentication.Dialog.fail
|
||||
pass
|
||||
|
||||
def prompt(self, echo, prompt):
|
||||
if self.fail:
|
||||
self.fail=False
|
||||
return b'bad guess at a password'
|
||||
return self.m.get(prompt)
|
||||
|
||||
class DialogHandler(object):
|
||||
|
||||
def __init__(self, con):
|
||||
self.con=con
|
||||
|
||||
def authenticate(self, pkt):
|
||||
while True:
|
||||
flag = pkt.read_uint8()
|
||||
echo = (flag & 0x06) == 0x02
|
||||
last = (flag & 0x01) == 0x01
|
||||
prompt = pkt.read_all()
|
||||
|
||||
if prompt == b'Password, please:':
|
||||
self.con.write_packet(b'stillnotverysecret\0')
|
||||
else:
|
||||
self.con.write_packet(b'no idea what to do with this prompt\0')
|
||||
pkt = self.con._read_packet()
|
||||
pkt.check_error()
|
||||
if pkt.is_ok_packet() or last:
|
||||
break
|
||||
return pkt
|
||||
|
||||
class DefectiveHandler(object):
|
||||
def __init__(self, con):
|
||||
self.con=con
|
||||
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipIf(two_questions_found, "two_questions plugin already installed")
|
||||
def testDialogAuthTwoQuestionsInstallPlugin(self):
|
||||
# needs plugin. lets install it.
|
||||
cur = self.connections[0].cursor()
|
||||
try:
|
||||
cur.execute("install plugin two_questions soname 'dialog_examples.so'")
|
||||
TestAuthentication.two_questions_found = True
|
||||
self.realTestDialogAuthTwoQuestions()
|
||||
except pymysql.err.InternalError:
|
||||
raise unittest2.SkipTest('we couldn\'t install the two_questions plugin')
|
||||
finally:
|
||||
if TestAuthentication.two_questions_found:
|
||||
cur.execute("uninstall plugin two_questions")
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(two_questions_found, "no two questions auth plugin")
|
||||
def testDialogAuthTwoQuestions(self):
|
||||
self.realTestDialogAuthTwoQuestions()
|
||||
|
||||
def realTestDialogAuthTwoQuestions(self):
|
||||
TestAuthentication.Dialog.fail=False
|
||||
TestAuthentication.Dialog.m = {b'Password, please:': b'notverysecret',
|
||||
b'Are you sure ?': b'yes, of course'}
|
||||
with TempUser(self.connections[0].cursor(), 'pymysql_2q@localhost',
|
||||
self.databases[0]['db'], 'two_questions', 'notverysecret') as u:
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_2q', **self.db)
|
||||
pymysql.connect(user='pymysql_2q', auth_plugin_map={b'dialog': TestAuthentication.Dialog}, **self.db)
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipIf(three_attempts_found, "three_attempts plugin already installed")
|
||||
def testDialogAuthThreeAttemptsQuestionsInstallPlugin(self):
|
||||
# needs plugin. lets install it.
|
||||
cur = self.connections[0].cursor()
|
||||
try:
|
||||
cur.execute("install plugin three_attempts soname 'dialog_examples.so'")
|
||||
TestAuthentication.three_attempts_found = True
|
||||
self.realTestDialogAuthThreeAttempts()
|
||||
except pymysql.err.InternalError:
|
||||
raise unittest2.SkipTest('we couldn\'t install the three_attempts plugin')
|
||||
finally:
|
||||
if TestAuthentication.three_attempts_found:
|
||||
cur.execute("uninstall plugin three_attempts")
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(three_attempts_found, "no three attempts plugin")
|
||||
def testDialogAuthThreeAttempts(self):
|
||||
self.realTestDialogAuthThreeAttempts()
|
||||
|
||||
def realTestDialogAuthThreeAttempts(self):
|
||||
TestAuthentication.Dialog.m = {b'Password, please:': b'stillnotverysecret'}
|
||||
TestAuthentication.Dialog.fail=True # fail just once. We've got three attempts after all
|
||||
with TempUser(self.connections[0].cursor(), 'pymysql_3a@localhost',
|
||||
self.databases[0]['db'], 'three_attempts', 'stillnotverysecret') as u:
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': TestAuthentication.Dialog}, **self.db)
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': TestAuthentication.DialogHandler}, **self.db)
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': object}, **self.db)
|
||||
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': TestAuthentication.DefectiveHandler}, **self.db)
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'notdialogplugin': TestAuthentication.Dialog}, **self.db)
|
||||
TestAuthentication.Dialog.m = {b'Password, please:': b'I do not know'}
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': TestAuthentication.Dialog}, **self.db)
|
||||
TestAuthentication.Dialog.m = {b'Password, please:': None}
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_3a', auth_plugin_map={b'dialog': TestAuthentication.Dialog}, **self.db)
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipIf(pam_found, "pam plugin already installed")
|
||||
@unittest2.skipIf(os.environ.get('PASSWORD') is None, "PASSWORD env var required")
|
||||
@unittest2.skipIf(os.environ.get('PAMSERVICE') is None, "PAMSERVICE env var required")
|
||||
def testPamAuthInstallPlugin(self):
|
||||
# needs plugin. lets install it.
|
||||
cur = self.connections[0].cursor()
|
||||
try:
|
||||
cur.execute("install plugin pam soname 'auth_pam.so'")
|
||||
TestAuthentication.pam_found = True
|
||||
self.realTestPamAuth()
|
||||
except pymysql.err.InternalError:
|
||||
raise unittest2.SkipTest('we couldn\'t install the auth_pam plugin')
|
||||
finally:
|
||||
if TestAuthentication.pam_found:
|
||||
cur.execute("uninstall plugin pam")
|
||||
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(pam_found, "no pam plugin")
|
||||
@unittest2.skipIf(os.environ.get('PASSWORD') is None, "PASSWORD env var required")
|
||||
@unittest2.skipIf(os.environ.get('PAMSERVICE') is None, "PAMSERVICE env var required")
|
||||
def testPamAuth(self):
|
||||
self.realTestPamAuth()
|
||||
|
||||
def realTestPamAuth(self):
|
||||
db = self.db.copy()
|
||||
import os
|
||||
db['password'] = os.environ.get('PASSWORD')
|
||||
cur = self.connections[0].cursor()
|
||||
try:
|
||||
cur.execute('show grants for ' + TestAuthentication.osuser + '@localhost')
|
||||
grants = cur.fetchone()[0]
|
||||
cur.execute('drop user ' + TestAuthentication.osuser + '@localhost')
|
||||
except pymysql.OperationalError as e:
|
||||
# assuming the user doesn't exist which is ok too
|
||||
self.assertEqual(1045, e.args[0])
|
||||
grants = None
|
||||
with TempUser(cur, TestAuthentication.osuser + '@localhost',
|
||||
self.databases[0]['db'], 'pam', os.environ.get('PAMSERVICE')) as u:
|
||||
try:
|
||||
c = pymysql.connect(user=TestAuthentication.osuser, **db)
|
||||
db['password'] = 'very bad guess at password'
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user=TestAuthentication.osuser,
|
||||
auth_plugin_map={b'mysql_cleartext_password': TestAuthentication.DefectiveHandler},
|
||||
**self.db)
|
||||
except pymysql.OperationalError as e:
|
||||
self.assertEqual(1045, e.args[0])
|
||||
# we had 'bad guess at password' work with pam. Well at least we get a permission denied here
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user=TestAuthentication.osuser,
|
||||
auth_plugin_map={b'mysql_cleartext_password': TestAuthentication.DefectiveHandler},
|
||||
**self.db)
|
||||
if grants:
|
||||
# recreate the user
|
||||
cur.execute(grants)
|
||||
|
||||
# select old_password("crummy p\tassword");
|
||||
#| old_password("crummy p\tassword") |
|
||||
#| 2a01785203b08770 |
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(mysql_old_password_found, "no mysql_old_password plugin")
|
||||
def testMySQLOldPasswordAuth(self):
|
||||
if self.mysql_server_is(self.connections[0], (5, 7, 0)):
|
||||
raise unittest2.SkipTest('Old passwords aren\'t supported in 5.7')
|
||||
# pymysql.err.OperationalError: (1045, "Access denied for user 'old_pass_user'@'localhost' (using password: YES)")
|
||||
# from login in MySQL-5.6
|
||||
if self.mysql_server_is(self.connections[0], (5, 6, 0)):
|
||||
raise unittest2.SkipTest('Old passwords don\'t authenticate in 5.6')
|
||||
db = self.db.copy()
|
||||
db['password'] = "crummy p\tassword"
|
||||
with self.connections[0] as c:
|
||||
# deprecated in 5.6
|
||||
if sys.version_info[0:2] >= (3,2) and self.mysql_server_is(self.connections[0], (5, 6, 0)):
|
||||
with self.assertWarns(pymysql.err.Warning) as cm:
|
||||
c.execute("SELECT OLD_PASSWORD('%s')" % db['password'])
|
||||
else:
|
||||
c.execute("SELECT OLD_PASSWORD('%s')" % db['password'])
|
||||
v = c.fetchone()[0]
|
||||
self.assertEqual(v, '2a01785203b08770')
|
||||
# only works in MariaDB and MySQL-5.6 - can't separate out by version
|
||||
#if self.mysql_server_is(self.connections[0], (5, 5, 0)):
|
||||
# with TempUser(c, 'old_pass_user@localhost',
|
||||
# self.databases[0]['db'], 'mysql_old_password', '2a01785203b08770') as u:
|
||||
# cur = pymysql.connect(user='old_pass_user', **db).cursor()
|
||||
# cur.execute("SELECT VERSION()")
|
||||
c.execute("SELECT @@secure_auth")
|
||||
secure_auth_setting = c.fetchone()[0]
|
||||
c.execute('set old_passwords=1')
|
||||
# pymysql.err.Warning: 'pre-4.1 password hash' is deprecated and will be removed in a future release. Please use post-4.1 password hash instead
|
||||
if sys.version_info[0:2] >= (3,2) and self.mysql_server_is(self.connections[0], (5, 6, 0)):
|
||||
with self.assertWarns(pymysql.err.Warning) as cm:
|
||||
c.execute('set global secure_auth=0')
|
||||
else:
|
||||
c.execute('set global secure_auth=0')
|
||||
with TempUser(c, 'old_pass_user@localhost',
|
||||
self.databases[0]['db'], password=db['password']) as u:
|
||||
cur = pymysql.connect(user='old_pass_user', **db).cursor()
|
||||
cur.execute("SELECT VERSION()")
|
||||
c.execute('set global secure_auth=%r' % secure_auth_setting)
|
||||
|
||||
@unittest2.skipUnless(socket_auth, "connection to unix_socket required")
|
||||
@unittest2.skipUnless(sha256_password_found, "no sha256 password authentication plugin found")
|
||||
def testAuthSHA256(self):
|
||||
c = self.connections[0].cursor()
|
||||
with TempUser(c, 'pymysql_sha256@localhost',
|
||||
self.databases[0]['db'], 'sha256_password') as u:
|
||||
if self.mysql_server_is(self.connections[0], (5, 7, 0)):
|
||||
c.execute("SET PASSWORD FOR 'pymysql_sha256'@'localhost' ='Sh@256Pa33'")
|
||||
else:
|
||||
c.execute('SET old_passwords = 2')
|
||||
c.execute("SET PASSWORD FOR 'pymysql_sha256'@'localhost' = PASSWORD('Sh@256Pa33')")
|
||||
db = self.db.copy()
|
||||
db['password'] = "Sh@256Pa33"
|
||||
# not implemented yet so thows error
|
||||
with self.assertRaises(pymysql.err.OperationalError):
|
||||
pymysql.connect(user='pymysql_256', **db)
|
||||
|
||||
class TestConnection(base.PyMySQLTestCase):
|
||||
|
||||
def test_utf8mb4(self):
|
||||
"""This test requires MySQL >= 5.5"""
|
||||
arg = self.databases[0].copy()
|
||||
arg['charset'] = 'utf8mb4'
|
||||
conn = pymysql.connect(**arg)
|
||||
|
||||
def test_largedata(self):
|
||||
"""Large query and response (>=16MB)"""
|
||||
cur = self.connections[0].cursor()
|
||||
cur.execute("SELECT @@max_allowed_packet")
|
||||
if cur.fetchone()[0] < 16*1024*1024 + 10:
|
||||
print("Set max_allowed_packet to bigger than 17MB")
|
||||
return
|
||||
t = 'a' * (16*1024*1024)
|
||||
cur.execute("SELECT '" + t + "'")
|
||||
assert cur.fetchone()[0] == t
|
||||
|
||||
def test_autocommit(self):
|
||||
con = self.connections[0]
|
||||
self.assertFalse(con.get_autocommit())
|
||||
|
||||
cur = con.cursor()
|
||||
cur.execute("SET AUTOCOMMIT=1")
|
||||
self.assertTrue(con.get_autocommit())
|
||||
|
||||
con.autocommit(False)
|
||||
self.assertFalse(con.get_autocommit())
|
||||
cur.execute("SELECT @@AUTOCOMMIT")
|
||||
self.assertEqual(cur.fetchone()[0], 0)
|
||||
|
||||
def test_select_db(self):
|
||||
con = self.connections[0]
|
||||
current_db = self.databases[0]['db']
|
||||
other_db = self.databases[1]['db']
|
||||
|
||||
cur = con.cursor()
|
||||
cur.execute('SELECT database()')
|
||||
self.assertEqual(cur.fetchone()[0], current_db)
|
||||
|
||||
con.select_db(other_db)
|
||||
cur.execute('SELECT database()')
|
||||
self.assertEqual(cur.fetchone()[0], other_db)
|
||||
|
||||
def test_connection_gone_away(self):
|
||||
"""
|
||||
http://dev.mysql.com/doc/refman/5.0/en/gone-away.html
|
||||
http://dev.mysql.com/doc/refman/5.0/en/error-messages-client.html#error_cr_server_gone_error
|
||||
"""
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
cur.execute("SET wait_timeout=1")
|
||||
time.sleep(2)
|
||||
with self.assertRaises(pymysql.OperationalError) as cm:
|
||||
cur.execute("SELECT 1+1")
|
||||
# error occures while reading, not writing because of socket buffer.
|
||||
#self.assertEqual(cm.exception.args[0], 2006)
|
||||
self.assertIn(cm.exception.args[0], (2006, 2013))
|
||||
|
||||
def test_init_command(self):
|
||||
conn = pymysql.connect(
|
||||
init_command='SELECT "bar"; SELECT "baz"',
|
||||
**self.databases[0]
|
||||
)
|
||||
c = conn.cursor()
|
||||
c.execute('select "foobar";')
|
||||
self.assertEqual(('foobar',), c.fetchone())
|
||||
conn.close()
|
||||
with self.assertRaises(pymysql.err.Error):
|
||||
conn.ping(reconnect=False)
|
||||
|
||||
def test_read_default_group(self):
|
||||
conn = pymysql.connect(
|
||||
read_default_group='client',
|
||||
**self.databases[0]
|
||||
)
|
||||
self.assertTrue(conn.open)
|
||||
|
||||
def test_context(self):
|
||||
with self.assertRaises(ValueError):
|
||||
c = pymysql.connect(**self.databases[0])
|
||||
with c as cur:
|
||||
cur.execute('create table test ( a int )')
|
||||
c.begin()
|
||||
cur.execute('insert into test values ((1))')
|
||||
raise ValueError('pseudo abort')
|
||||
c.commit()
|
||||
c = pymysql.connect(**self.databases[0])
|
||||
with c as cur:
|
||||
cur.execute('select count(*) from test')
|
||||
self.assertEqual(0, cur.fetchone()[0])
|
||||
cur.execute('insert into test values ((1))')
|
||||
with c as cur:
|
||||
cur.execute('select count(*) from test')
|
||||
self.assertEqual(1,cur.fetchone()[0])
|
||||
cur.execute('drop table test')
|
||||
|
||||
def test_set_charset(self):
|
||||
c = pymysql.connect(**self.databases[0])
|
||||
c.set_charset('utf8')
|
||||
# TODO validate setting here
|
||||
|
||||
def test_defer_connect(self):
|
||||
import socket
|
||||
for db in self.databases:
|
||||
d = db.copy()
|
||||
try:
|
||||
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
sock.connect(d['unix_socket'])
|
||||
except KeyError:
|
||||
sock = socket.create_connection(
|
||||
(d.get('host', 'localhost'), d.get('port', 3306)))
|
||||
for k in ['unix_socket', 'host', 'port']:
|
||||
try:
|
||||
del d[k]
|
||||
except KeyError:
|
||||
pass
|
||||
|
||||
c = pymysql.connect(defer_connect=True, **d)
|
||||
self.assertFalse(c.open)
|
||||
c.connect(sock)
|
||||
c.close()
|
||||
sock.close()
|
||||
|
||||
@unittest2.skipUnless(sys.version_info[0:2] >= (3,2), "required py-3.2")
|
||||
def test_no_delay_warning(self):
|
||||
current_db = self.databases[0].copy()
|
||||
current_db['no_delay'] = True
|
||||
with self.assertWarns(DeprecationWarning) as cm:
|
||||
conn = pymysql.connect(**current_db)
|
||||
|
||||
|
||||
# A custom type and function to escape it
|
||||
class Foo(object):
|
||||
value = "bar"
|
||||
|
||||
|
||||
def escape_foo(x, d):
|
||||
return x.value
|
||||
|
||||
|
||||
class TestEscape(base.PyMySQLTestCase):
|
||||
def test_escape_string(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
self.assertEqual(con.escape("foo'bar"), "'foo\\'bar'")
|
||||
# added NO_AUTO_CREATE_USER as not including it in 5.7 generates warnings
|
||||
cur.execute("SET sql_mode='NO_BACKSLASH_ESCAPES,NO_AUTO_CREATE_USER'")
|
||||
self.assertEqual(con.escape("foo'bar"), "'foo''bar'")
|
||||
|
||||
def test_escape_builtin_encoders(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
val = datetime.datetime(2012, 3, 4, 5, 6)
|
||||
self.assertEqual(con.escape(val, con.encoders), "'2012-03-04 05:06:00'")
|
||||
|
||||
def test_escape_custom_object(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
mapping = {Foo: escape_foo}
|
||||
self.assertEqual(con.escape(Foo(), mapping), "bar")
|
||||
|
||||
def test_escape_fallback_encoder(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
class Custom(str):
|
||||
pass
|
||||
|
||||
mapping = {text_type: pymysql.escape_string}
|
||||
self.assertEqual(con.escape(Custom('foobar'), mapping), "'foobar'")
|
||||
|
||||
def test_escape_no_default(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
self.assertRaises(TypeError, con.escape, 42, {})
|
||||
|
||||
def test_escape_dict_value(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
mapping = con.encoders.copy()
|
||||
mapping[Foo] = escape_foo
|
||||
self.assertEqual(con.escape({'foo': Foo()}, mapping), {'foo': "bar"})
|
||||
|
||||
def test_escape_list_item(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
|
||||
mapping = con.encoders.copy()
|
||||
mapping[Foo] = escape_foo
|
||||
self.assertEqual(con.escape([Foo()], mapping), "(bar)")
|
||||
|
||||
def test_previous_cursor_not_closed(self):
|
||||
con = self.connections[0]
|
||||
cur1 = con.cursor()
|
||||
cur1.execute("SELECT 1; SELECT 2")
|
||||
cur2 = con.cursor()
|
||||
cur2.execute("SELECT 3")
|
||||
self.assertEqual(cur2.fetchone()[0], 3)
|
||||
|
||||
def test_commit_during_multi_result(self):
|
||||
con = self.connections[0]
|
||||
cur = con.cursor()
|
||||
cur.execute("SELECT 1; SELECT 2")
|
||||
con.commit()
|
||||
cur.execute("SELECT 3")
|
||||
self.assertEqual(cur.fetchone()[0], 3)
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
import datetime
|
||||
from unittest import TestCase
|
||||
|
||||
from pymysql._compat import PY2
|
||||
from pymysql import converters
|
||||
|
||||
|
||||
__all__ = ["TestConverter"]
|
||||
|
||||
|
||||
class TestConverter(TestCase):
|
||||
|
||||
def test_escape_string(self):
|
||||
self.assertEqual(
|
||||
converters.escape_string(u"foo\nbar"),
|
||||
u"foo\\nbar"
|
||||
)
|
||||
|
||||
if PY2:
|
||||
def test_escape_string_bytes(self):
|
||||
self.assertEqual(
|
||||
converters.escape_string(b"foo\nbar"),
|
||||
b"foo\\nbar"
|
||||
)
|
||||
|
||||
def test_convert_datetime(self):
|
||||
expected = datetime.datetime(2007, 2, 24, 23, 6, 20)
|
||||
dt = converters.convert_datetime('2007-02-24 23:06:20')
|
||||
self.assertEqual(dt, expected)
|
||||
|
||||
def test_convert_datetime_with_fsp(self):
|
||||
expected = datetime.datetime(2007, 2, 24, 23, 6, 20, 511581)
|
||||
dt = converters.convert_datetime('2007-02-24 23:06:20.511581')
|
||||
self.assertEqual(dt, expected)
|
||||
|
||||
def _test_convert_timedelta(self, with_negate=False, with_fsp=False):
|
||||
d = {'hours': 789, 'minutes': 12, 'seconds': 34}
|
||||
s = '%(hours)s:%(minutes)s:%(seconds)s' % d
|
||||
if with_fsp:
|
||||
d['microseconds'] = 511581
|
||||
s += '.%(microseconds)s' % d
|
||||
|
||||
expected = datetime.timedelta(**d)
|
||||
if with_negate:
|
||||
expected = -expected
|
||||
s = '-' + s
|
||||
|
||||
tdelta = converters.convert_timedelta(s)
|
||||
self.assertEqual(tdelta, expected)
|
||||
|
||||
def test_convert_timedelta(self):
|
||||
self._test_convert_timedelta(with_negate=False, with_fsp=False)
|
||||
self._test_convert_timedelta(with_negate=True, with_fsp=False)
|
||||
|
||||
def test_convert_timedelta_with_fsp(self):
|
||||
self._test_convert_timedelta(with_negate=False, with_fsp=True)
|
||||
self._test_convert_timedelta(with_negate=False, with_fsp=True)
|
||||
|
||||
def test_convert_time(self):
|
||||
expected = datetime.time(23, 6, 20)
|
||||
time_obj = converters.convert_time('23:06:20')
|
||||
self.assertEqual(time_obj, expected)
|
||||
|
||||
def test_convert_time_with_fsp(self):
|
||||
expected = datetime.time(23, 6, 20, 511581)
|
||||
time_obj = converters.convert_time('23:06:20.511581')
|
||||
self.assertEqual(time_obj, expected)
|
||||
Executable
+104
@@ -0,0 +1,104 @@
|
||||
import warnings
|
||||
|
||||
from pymysql.tests import base
|
||||
import pymysql.cursors
|
||||
|
||||
class CursorTest(base.PyMySQLTestCase):
|
||||
def setUp(self):
|
||||
super(CursorTest, self).setUp()
|
||||
|
||||
conn = self.connections[0]
|
||||
self.safe_create_table(
|
||||
conn,
|
||||
"test", "create table test (data varchar(10))",
|
||||
)
|
||||
cursor = conn.cursor()
|
||||
cursor.execute(
|
||||
"insert into test (data) values "
|
||||
"('row1'), ('row2'), ('row3'), ('row4'), ('row5')")
|
||||
cursor.close()
|
||||
self.test_connection = pymysql.connect(**self.databases[0])
|
||||
self.addCleanup(self.test_connection.close)
|
||||
|
||||
def test_cleanup_rows_unbuffered(self):
|
||||
conn = self.test_connection
|
||||
cursor = conn.cursor(pymysql.cursors.SSCursor)
|
||||
|
||||
cursor.execute("select * from test as t1, test as t2")
|
||||
for counter, row in enumerate(cursor):
|
||||
if counter > 10:
|
||||
break
|
||||
|
||||
del cursor
|
||||
self.safe_gc_collect()
|
||||
|
||||
c2 = conn.cursor()
|
||||
|
||||
c2.execute("select 1")
|
||||
self.assertEqual(c2.fetchone(), (1,))
|
||||
self.assertIsNone(c2.fetchone())
|
||||
|
||||
def test_cleanup_rows_buffered(self):
|
||||
conn = self.test_connection
|
||||
cursor = conn.cursor(pymysql.cursors.Cursor)
|
||||
|
||||
cursor.execute("select * from test as t1, test as t2")
|
||||
for counter, row in enumerate(cursor):
|
||||
if counter > 10:
|
||||
break
|
||||
|
||||
del cursor
|
||||
self.safe_gc_collect()
|
||||
|
||||
c2 = conn.cursor()
|
||||
|
||||
c2.execute("select 1")
|
||||
|
||||
self.assertEqual(
|
||||
c2.fetchone(), (1,)
|
||||
)
|
||||
self.assertIsNone(c2.fetchone())
|
||||
|
||||
def test_executemany(self):
|
||||
conn = self.test_connection
|
||||
cursor = conn.cursor(pymysql.cursors.Cursor)
|
||||
|
||||
m = pymysql.cursors.RE_INSERT_VALUES.match("INSERT INTO TEST (ID, NAME) VALUES (%s, %s)")
|
||||
self.assertIsNotNone(m, 'error parse %s')
|
||||
self.assertEqual(m.group(3), '', 'group 3 not blank, bug in RE_INSERT_VALUES?')
|
||||
|
||||
m = pymysql.cursors.RE_INSERT_VALUES.match("INSERT INTO TEST (ID, NAME) VALUES (%(id)s, %(name)s)")
|
||||
self.assertIsNotNone(m, 'error parse %(name)s')
|
||||
self.assertEqual(m.group(3), '', 'group 3 not blank, bug in RE_INSERT_VALUES?')
|
||||
|
||||
m = pymysql.cursors.RE_INSERT_VALUES.match("INSERT INTO TEST (ID, NAME) VALUES (%(id_name)s, %(name)s)")
|
||||
self.assertIsNotNone(m, 'error parse %(id_name)s')
|
||||
self.assertEqual(m.group(3), '', 'group 3 not blank, bug in RE_INSERT_VALUES?')
|
||||
|
||||
m = pymysql.cursors.RE_INSERT_VALUES.match("INSERT INTO TEST (ID, NAME) VALUES (%(id_name)s, %(name)s) ON duplicate update")
|
||||
self.assertIsNotNone(m, 'error parse %(id_name)s')
|
||||
self.assertEqual(m.group(3), ' ON duplicate update', 'group 3 not ON duplicate update, bug in RE_INSERT_VALUES?')
|
||||
|
||||
# cursor._executed must bee "insert into test (data) values (0),(1),(2),(3),(4),(5),(6),(7),(8),(9)"
|
||||
# list args
|
||||
data = range(10)
|
||||
cursor.executemany("insert into test (data) values (%s)", data)
|
||||
self.assertTrue(cursor._executed.endswith(b",(7),(8),(9)"), 'execute many with %s not in one query')
|
||||
|
||||
# dict args
|
||||
data_dict = [{'data': i} for i in range(10)]
|
||||
cursor.executemany("insert into test (data) values (%(data)s)", data_dict)
|
||||
self.assertTrue(cursor._executed.endswith(b",(7),(8),(9)"), 'execute many with %(data)s not in one query')
|
||||
|
||||
# %% in column set
|
||||
cursor.execute("""\
|
||||
CREATE TABLE percent_test (
|
||||
`A%` INTEGER,
|
||||
`B%` INTEGER)""")
|
||||
try:
|
||||
q = "INSERT INTO percent_test (`A%%`, `B%%`) VALUES (%s, %s)"
|
||||
self.assertIsNotNone(pymysql.cursors.RE_INSERT_VALUES.match(q))
|
||||
cursor.executemany(q, [(3, 4), (5, 6)])
|
||||
self.assertTrue(cursor._executed.endswith(b"(3, 4),(5, 6)"), "executemany with %% not in one query")
|
||||
finally:
|
||||
cursor.execute("DROP TABLE IF EXISTS percent_test")
|
||||
Executable
+21
@@ -0,0 +1,21 @@
|
||||
import unittest2
|
||||
|
||||
from pymysql import err
|
||||
|
||||
|
||||
__all__ = ["TestRaiseException"]
|
||||
|
||||
|
||||
class TestRaiseException(unittest2.TestCase):
|
||||
|
||||
def test_raise_mysql_exception(self):
|
||||
data = b"\xff\x15\x04Access denied"
|
||||
with self.assertRaises(err.OperationalError) as cm:
|
||||
err.raise_mysql_exception(data)
|
||||
self.assertEqual(cm.exception.args, (1045, 'Access denied'))
|
||||
|
||||
def test_raise_mysql_exception_client_protocol_41(self):
|
||||
data = b"\xff\x15\x04#28000Access denied"
|
||||
with self.assertRaises(err.OperationalError) as cm:
|
||||
err.raise_mysql_exception(data)
|
||||
self.assertEqual(cm.exception.args, (1045, 'Access denied'))
|
||||
@@ -1,26 +1,31 @@
|
||||
import pymysql
|
||||
from pymysql.tests import base
|
||||
import unittest
|
||||
|
||||
import datetime
|
||||
import time
|
||||
import warnings
|
||||
import sys
|
||||
|
||||
import pymysql
|
||||
from pymysql import cursors
|
||||
from pymysql._compat import text_type
|
||||
from pymysql.tests import base
|
||||
import unittest2
|
||||
|
||||
try:
|
||||
import imp
|
||||
reload = imp.reload
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
import datetime
|
||||
|
||||
# backwards compatibility:
|
||||
if not hasattr(unittest, "skip"):
|
||||
unittest.skip = lambda message: lambda f: f
|
||||
__all__ = ["TestOldIssues", "TestNewIssues", "TestGitHubIssues"]
|
||||
|
||||
class TestOldIssues(base.PyMySQLTestCase):
|
||||
def test_issue_3(self):
|
||||
""" undefined methods datetime_or_None, date_or_None """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue3")
|
||||
c.execute("create table issue3 (d date, t time, dt datetime, ts timestamp)")
|
||||
try:
|
||||
c.execute("insert into issue3 (d, t, dt, ts) values (%s,%s,%s,%s)", (None, None, None, None))
|
||||
@@ -39,6 +44,9 @@ class TestOldIssues(base.PyMySQLTestCase):
|
||||
""" can't retrieve TIMESTAMP fields """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue4")
|
||||
c.execute("create table issue4 (ts timestamp)")
|
||||
try:
|
||||
c.execute("insert into issue4 (ts) values (now())")
|
||||
@@ -55,7 +63,10 @@ class TestOldIssues(base.PyMySQLTestCase):
|
||||
|
||||
def test_issue_6(self):
|
||||
""" exception: TypeError: ord() expected a character, but string of length 0 found """
|
||||
conn = pymysql.connect(host="localhost",user="root",passwd="",db="mysql")
|
||||
# ToDo: this test requires access to db 'mysql'.
|
||||
kwargs = self.databases[0].copy()
|
||||
kwargs['db'] = "mysql"
|
||||
conn = pymysql.connect(**kwargs)
|
||||
c = conn.cursor()
|
||||
c.execute("select * from user")
|
||||
conn.close()
|
||||
@@ -64,8 +75,11 @@ class TestOldIssues(base.PyMySQLTestCase):
|
||||
""" Primary Key and Index error when selecting data """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists test")
|
||||
c.execute("""CREATE TABLE `test` (`station` int(10) NOT NULL DEFAULT '0', `dh`
|
||||
datetime NOT NULL DEFAULT '0000-00-00 00:00:00', `echeance` int(1) NOT NULL
|
||||
datetime NOT NULL DEFAULT '2015-01-01 00:00:00', `echeance` int(1) NOT NULL
|
||||
DEFAULT '0', `me` double DEFAULT NULL, `mo` double DEFAULT NULL, PRIMARY
|
||||
KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
try:
|
||||
@@ -82,18 +96,13 @@ KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
except DeprecationWarning:
|
||||
self.fail()
|
||||
|
||||
def test_issue_10(self):
|
||||
""" Allocate a variable to return when the exception handler is permissive """
|
||||
conn = self.connections[0]
|
||||
conn.errorhandler = lambda cursor, errorclass, errorvalue: None
|
||||
cur = conn.cursor()
|
||||
cur.execute( "create table t( n int )" )
|
||||
cur.execute( "create table t( n int )" )
|
||||
|
||||
def test_issue_13(self):
|
||||
""" can't handle large result fields """
|
||||
conn = self.connections[0]
|
||||
cur = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
cur.execute("drop table if exists issue13")
|
||||
try:
|
||||
cur.execute("create table issue13 (t text)")
|
||||
# ticket says 18k
|
||||
@@ -106,18 +115,13 @@ KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
finally:
|
||||
cur.execute("drop table issue13")
|
||||
|
||||
def test_issue_14(self):
|
||||
""" typo in converters.py """
|
||||
self.assertEqual('1', pymysql.converters.escape_item(1, "utf8"))
|
||||
self.assertEqual('1', pymysql.converters.escape_item(1L, "utf8"))
|
||||
|
||||
self.assertEqual('1', pymysql.converters.escape_object(1))
|
||||
self.assertEqual('1', pymysql.converters.escape_object(1L))
|
||||
|
||||
def test_issue_15(self):
|
||||
""" query should be expanded before perform character encoding """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue15")
|
||||
c.execute("create table issue15 (t varchar(32))")
|
||||
try:
|
||||
c.execute("insert into issue15 (t) values (%s)", (u'\xe4\xf6\xfc',))
|
||||
@@ -130,6 +134,9 @@ KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
""" Patch for string and tuple escaping """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue16")
|
||||
c.execute("create table issue16 (name varchar(32) primary key, email varchar(32))")
|
||||
try:
|
||||
c.execute("insert into issue16 (name, email) values ('pete', 'floydophone')")
|
||||
@@ -138,20 +145,24 @@ KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
finally:
|
||||
c.execute("drop table issue16")
|
||||
|
||||
@unittest.skip("test_issue_17() requires a custom, legacy MySQL configuration and will not be run.")
|
||||
@unittest2.skip("test_issue_17() requires a custom, legacy MySQL configuration and will not be run.")
|
||||
def test_issue_17(self):
|
||||
""" could not connect mysql use passwod """
|
||||
"""could not connect mysql use passwod"""
|
||||
conn = self.connections[0]
|
||||
host = self.databases[0]["host"]
|
||||
db = self.databases[0]["db"]
|
||||
c = conn.cursor()
|
||||
|
||||
# grant access to a table to a user with a password
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue17")
|
||||
c.execute("create table issue17 (x varchar(32) primary key)")
|
||||
c.execute("insert into issue17 (x) values ('hello, world!')")
|
||||
c.execute("grant all privileges on %s.issue17 to 'issue17user'@'%%' identified by '1234'" % db)
|
||||
conn.commit()
|
||||
|
||||
|
||||
conn2 = pymysql.connect(host=host, user="issue17user", passwd="1234", db=db)
|
||||
c2 = conn2.cursor()
|
||||
c2.execute("select x from issue17")
|
||||
@@ -159,71 +170,71 @@ KEY (`station`,`dh`,`echeance`)) ENGINE=MyISAM DEFAULT CHARSET=latin1;""")
|
||||
finally:
|
||||
c.execute("drop table issue17")
|
||||
|
||||
def _uni(s, e):
|
||||
# hack for py3
|
||||
if sys.version_info[0] > 2:
|
||||
return unicode(bytes(s, sys.getdefaultencoding()), e)
|
||||
else:
|
||||
return unicode(s, e)
|
||||
|
||||
class TestNewIssues(base.PyMySQLTestCase):
|
||||
def test_issue_34(self):
|
||||
try:
|
||||
pymysql.connect(host="localhost", port=1237, user="root")
|
||||
self.fail()
|
||||
except pymysql.OperationalError, e:
|
||||
except pymysql.OperationalError as e:
|
||||
self.assertEqual(2003, e.args[0])
|
||||
except:
|
||||
except Exception:
|
||||
self.fail()
|
||||
|
||||
def test_issue_33(self):
|
||||
conn = pymysql.connect(host="localhost", user="root", db=self.databases[0]["db"], charset="utf8")
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
self.safe_create_table(conn, u'hei\xdfe',
|
||||
u'create table hei\xdfe (name varchar(32))')
|
||||
c = conn.cursor()
|
||||
try:
|
||||
c.execute(_uni("create table hei\xc3\x9fe (name varchar(32))", "utf8"))
|
||||
c.execute(_uni("insert into hei\xc3\x9fe (name) values ('Pi\xc3\xb1ata')", "utf8"))
|
||||
c.execute(_uni("select name from hei\xc3\x9fe", "utf8"))
|
||||
self.assertEqual(_uni("Pi\xc3\xb1ata","utf8"), c.fetchone()[0])
|
||||
finally:
|
||||
c.execute(_uni("drop table hei\xc3\x9fe", "utf8"))
|
||||
c.execute(u"insert into hei\xdfe (name) values ('Pi\xdfata')")
|
||||
c.execute(u"select name from hei\xdfe")
|
||||
self.assertEqual(u"Pi\xdfata", c.fetchone()[0])
|
||||
|
||||
@unittest.skip("This test requires manual intervention")
|
||||
@unittest2.skip("This test requires manual intervention")
|
||||
def test_issue_35(self):
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
print "sudo killall -9 mysqld within the next 10 seconds"
|
||||
print("sudo killall -9 mysqld within the next 10 seconds")
|
||||
try:
|
||||
c.execute("select sleep(10)")
|
||||
self.fail()
|
||||
except pymysql.OperationalError, e:
|
||||
except pymysql.OperationalError as e:
|
||||
self.assertEqual(2013, e.args[0])
|
||||
|
||||
def test_issue_36(self):
|
||||
conn = self.connections[0]
|
||||
# connection 0 is super user, connection 1 isn't
|
||||
conn = self.connections[1]
|
||||
c = conn.cursor()
|
||||
# kill connections[0]
|
||||
c.execute("show processlist")
|
||||
kill_id = None
|
||||
for id,user,host,db,command,time,state,info in c.fetchall():
|
||||
for row in c.fetchall():
|
||||
id = row[0]
|
||||
info = row[7]
|
||||
if info == "show processlist":
|
||||
kill_id = id
|
||||
break
|
||||
self.assertEqual(kill_id, conn.thread_id())
|
||||
# now nuke the connection
|
||||
conn.kill(kill_id)
|
||||
self.connections[0].kill(kill_id)
|
||||
# make sure this connection has broken
|
||||
try:
|
||||
c.execute("show tables")
|
||||
self.fail()
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
c.close()
|
||||
conn.close()
|
||||
|
||||
# check the process list from the other connection
|
||||
try:
|
||||
c = self.connections[1].cursor()
|
||||
# Wait since Travis-CI sometimes fail this test.
|
||||
time.sleep(0.1)
|
||||
|
||||
c = self.connections[0].cursor()
|
||||
c.execute("show processlist")
|
||||
ids = [row[0] for row in c.fetchall()]
|
||||
self.assertFalse(kill_id in ids)
|
||||
finally:
|
||||
del self.connections[0]
|
||||
del self.connections[1]
|
||||
|
||||
def test_issue_37(self):
|
||||
conn = self.connections[0]
|
||||
@@ -237,8 +248,11 @@ class TestNewIssues(base.PyMySQLTestCase):
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
datum = "a" * 1024 * 1023 # reduced size for most default mysql installs
|
||||
|
||||
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue38")
|
||||
c.execute("create table issue38 (id integer, data mediumblob)")
|
||||
c.execute("insert into issue38 values (1, %s)", (datum,))
|
||||
finally:
|
||||
@@ -247,8 +261,11 @@ class TestNewIssues(base.PyMySQLTestCase):
|
||||
def disabled_test_issue_54(self):
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue54")
|
||||
big_sql = "select * from issue54 where "
|
||||
big_sql += " and ".join("%d=%d" % (i,i) for i in xrange(0, 100000))
|
||||
big_sql += " and ".join("%d=%d" % (i,i) for i in range(0, 100000))
|
||||
|
||||
try:
|
||||
c.execute("create table issue54 (id integer primary key)")
|
||||
@@ -260,10 +277,14 @@ class TestNewIssues(base.PyMySQLTestCase):
|
||||
|
||||
class TestGitHubIssues(base.PyMySQLTestCase):
|
||||
def test_issue_66(self):
|
||||
""" 'Connection' object has no attribute 'insert_id' """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
self.assertEqual(0, conn.insert_id())
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists issue66")
|
||||
c.execute("create table issue66 (id integer primary key auto_increment, x integer)")
|
||||
c.execute("insert into issue66 (x) values (1)")
|
||||
c.execute("insert into issue66 (x) values (1)")
|
||||
@@ -271,8 +292,224 @@ class TestGitHubIssues(base.PyMySQLTestCase):
|
||||
finally:
|
||||
c.execute("drop table issue66")
|
||||
|
||||
__all__ = ["TestOldIssues", "TestNewIssues", "TestGitHubIssues"]
|
||||
def test_issue_79(self):
|
||||
""" Duplicate field overwrites the previous one in the result of DictCursor """
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor(pymysql.cursors.DictCursor)
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
unittest.main()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
c.execute("drop table if exists a")
|
||||
c.execute("drop table if exists b")
|
||||
c.execute("""CREATE TABLE a (id int, value int)""")
|
||||
c.execute("""CREATE TABLE b (id int, value int)""")
|
||||
|
||||
a=(1,11)
|
||||
b=(1,22)
|
||||
try:
|
||||
c.execute("insert into a values (%s, %s)", a)
|
||||
c.execute("insert into b values (%s, %s)", b)
|
||||
|
||||
c.execute("SELECT * FROM a inner join b on a.id = b.id")
|
||||
r = c.fetchall()[0]
|
||||
self.assertEqual(r['id'], 1)
|
||||
self.assertEqual(r['value'], 11)
|
||||
self.assertEqual(r['b.value'], 22)
|
||||
finally:
|
||||
c.execute("drop table a")
|
||||
c.execute("drop table b")
|
||||
|
||||
def test_issue_95(self):
|
||||
""" Leftover trailing OK packet for "CALL my_sp" queries """
|
||||
conn = self.connections[0]
|
||||
cur = conn.cursor()
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
cur.execute("DROP PROCEDURE IF EXISTS `foo`")
|
||||
cur.execute("""CREATE PROCEDURE `foo` ()
|
||||
BEGIN
|
||||
SELECT 1;
|
||||
END""")
|
||||
try:
|
||||
cur.execute("""CALL foo()""")
|
||||
cur.execute("""SELECT 1""")
|
||||
self.assertEqual(cur.fetchone()[0], 1)
|
||||
finally:
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
cur.execute("DROP PROCEDURE IF EXISTS `foo`")
|
||||
|
||||
def test_issue_114(self):
|
||||
""" autocommit is not set after reconnecting with ping() """
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
conn.autocommit(False)
|
||||
c = conn.cursor()
|
||||
c.execute("""select @@autocommit;""")
|
||||
self.assertFalse(c.fetchone()[0])
|
||||
conn.close()
|
||||
conn.ping()
|
||||
c.execute("""select @@autocommit;""")
|
||||
self.assertFalse(c.fetchone()[0])
|
||||
conn.close()
|
||||
|
||||
# Ensure autocommit() is still working
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
c = conn.cursor()
|
||||
c.execute("""select @@autocommit;""")
|
||||
self.assertFalse(c.fetchone()[0])
|
||||
conn.close()
|
||||
conn.ping()
|
||||
conn.autocommit(True)
|
||||
c.execute("""select @@autocommit;""")
|
||||
self.assertTrue(c.fetchone()[0])
|
||||
conn.close()
|
||||
|
||||
def test_issue_175(self):
|
||||
""" The number of fields returned by server is read in wrong way """
|
||||
conn = self.connections[0]
|
||||
cur = conn.cursor()
|
||||
for length in (200, 300):
|
||||
columns = ', '.join('c{0} integer'.format(i) for i in range(length))
|
||||
sql = 'create table test_field_count ({0})'.format(columns)
|
||||
try:
|
||||
cur.execute(sql)
|
||||
cur.execute('select * from test_field_count')
|
||||
assert len(cur.description) == length
|
||||
finally:
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore")
|
||||
cur.execute('drop table if exists test_field_count')
|
||||
|
||||
def test_issue_321(self):
|
||||
""" Test iterable as query argument. """
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
self.safe_create_table(
|
||||
conn, "issue321",
|
||||
"create table issue321 (value_1 varchar(1), value_2 varchar(1))")
|
||||
|
||||
sql_insert = "insert into issue321 (value_1, value_2) values (%s, %s)"
|
||||
sql_dict_insert = ("insert into issue321 (value_1, value_2) "
|
||||
"values (%(value_1)s, %(value_2)s)")
|
||||
sql_select = ("select * from issue321 where "
|
||||
"value_1 in %s and value_2=%s")
|
||||
data = [
|
||||
[(u"a", ), u"\u0430"],
|
||||
[[u"b"], u"\u0430"],
|
||||
{"value_1": [[u"c"]], "value_2": u"\u0430"}
|
||||
]
|
||||
cur = conn.cursor()
|
||||
self.assertEqual(cur.execute(sql_insert, data[0]), 1)
|
||||
self.assertEqual(cur.execute(sql_insert, data[1]), 1)
|
||||
self.assertEqual(cur.execute(sql_dict_insert, data[2]), 1)
|
||||
self.assertEqual(
|
||||
cur.execute(sql_select, [(u"a", u"b", u"c"), u"\u0430"]), 3)
|
||||
self.assertEqual(cur.fetchone(), (u"a", u"\u0430"))
|
||||
self.assertEqual(cur.fetchone(), (u"b", u"\u0430"))
|
||||
self.assertEqual(cur.fetchone(), (u"c", u"\u0430"))
|
||||
|
||||
def test_issue_364(self):
|
||||
""" Test mixed unicode/binary arguments in executemany. """
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
self.safe_create_table(
|
||||
conn, "issue364",
|
||||
"create table issue364 (value_1 binary(3), value_2 varchar(3)) "
|
||||
"engine=InnoDB default charset=utf8")
|
||||
|
||||
sql = "insert into issue364 (value_1, value_2) values (%s, %s)"
|
||||
usql = u"insert into issue364 (value_1, value_2) values (%s, %s)"
|
||||
values = [pymysql.Binary(b"\x00\xff\x00"), u"\xe4\xf6\xfc"]
|
||||
|
||||
# test single insert and select
|
||||
cur = conn.cursor()
|
||||
cur.execute(sql, args=values)
|
||||
cur.execute("select * from issue364")
|
||||
self.assertEqual(cur.fetchone(), tuple(values))
|
||||
|
||||
# test single insert unicode query
|
||||
cur.execute(usql, args=values)
|
||||
|
||||
# test multi insert and select
|
||||
cur.executemany(sql, args=(values, values, values))
|
||||
cur.execute("select * from issue364")
|
||||
for row in cur.fetchall():
|
||||
self.assertEqual(row, tuple(values))
|
||||
|
||||
# test multi insert with unicode query
|
||||
cur.executemany(usql, args=(values, values, values))
|
||||
|
||||
def test_issue_363(self):
|
||||
""" Test binary / geometry types. """
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
self.safe_create_table(
|
||||
conn, "issue363",
|
||||
"CREATE TABLE issue363 ( "
|
||||
"id INTEGER PRIMARY KEY, geom LINESTRING NOT NULL, "
|
||||
"SPATIAL KEY geom (geom)) "
|
||||
"ENGINE=MyISAM default charset=utf8")
|
||||
|
||||
cur = conn.cursor()
|
||||
query = ("INSERT INTO issue363 (id, geom) VALUES"
|
||||
"(1998, GeomFromText('LINESTRING(1.1 1.1,2.2 2.2)'))")
|
||||
# From MySQL 5.7, ST_GeomFromText is added and GeomFromText is deprecated.
|
||||
if self.mysql_server_is(conn, (5, 7, 0)):
|
||||
with self.assertWarns(pymysql.err.Warning) as cm:
|
||||
cur.execute(query)
|
||||
else:
|
||||
cur.execute(query)
|
||||
|
||||
# select WKT
|
||||
query = "SELECT AsText(geom) FROM issue363"
|
||||
if self.mysql_server_is(conn, (5, 7, 0)):
|
||||
with self.assertWarns(pymysql.err.Warning) as cm:
|
||||
cur.execute(query)
|
||||
else:
|
||||
cur.execute(query)
|
||||
row = cur.fetchone()
|
||||
self.assertEqual(row, ("LINESTRING(1.1 1.1,2.2 2.2)", ))
|
||||
|
||||
# select WKB
|
||||
query = "SELECT AsBinary(geom) FROM issue363"
|
||||
if self.mysql_server_is(conn, (5, 7, 0)):
|
||||
with self.assertWarns(pymysql.err.Warning) as cm:
|
||||
cur.execute(query)
|
||||
else:
|
||||
cur.execute(query)
|
||||
row = cur.fetchone()
|
||||
self.assertEqual(row,
|
||||
(b"\x01\x02\x00\x00\x00\x02\x00\x00\x00"
|
||||
b"\x9a\x99\x99\x99\x99\x99\xf1?"
|
||||
b"\x9a\x99\x99\x99\x99\x99\xf1?"
|
||||
b"\x9a\x99\x99\x99\x99\x99\x01@"
|
||||
b"\x9a\x99\x99\x99\x99\x99\x01@", ))
|
||||
|
||||
# select internal binary
|
||||
cur.execute("SELECT geom FROM issue363")
|
||||
row = cur.fetchone()
|
||||
# don't assert the exact internal binary value, as it could
|
||||
# vary across implementations
|
||||
self.assertTrue(isinstance(row[0], bytes))
|
||||
|
||||
def test_issue_491(self):
|
||||
""" Test warning propagation """
|
||||
conn = pymysql.connect(charset="utf8", **self.databases[0])
|
||||
|
||||
with warnings.catch_warnings():
|
||||
# Ignore all warnings other than pymysql generated ones
|
||||
warnings.simplefilter("ignore")
|
||||
warnings.simplefilter("error", category=pymysql.Warning)
|
||||
|
||||
# verify for both buffered and unbuffered cursor types
|
||||
for cursor_class in (cursors.Cursor, cursors.SSCursor):
|
||||
c = conn.cursor(cursor_class)
|
||||
try:
|
||||
c.execute("SELECT CAST('124b' AS SIGNED)")
|
||||
c.fetchall()
|
||||
except pymysql.Warning as e:
|
||||
# Warnings should have errorcode and string message, just like exceptions
|
||||
self.assertEqual(len(e.args), 2)
|
||||
self.assertEqual(e.args[0], 1292)
|
||||
self.assertTrue(isinstance(e.args[1], text_type))
|
||||
else:
|
||||
self.fail("Should raise Warning")
|
||||
finally:
|
||||
c.close()
|
||||
|
||||
+93
@@ -0,0 +1,93 @@
|
||||
from pymysql import cursors, OperationalError, Warning
|
||||
from pymysql.tests import base
|
||||
|
||||
import os
|
||||
import warnings
|
||||
|
||||
__all__ = ["TestLoadLocal"]
|
||||
|
||||
|
||||
class TestLoadLocal(base.PyMySQLTestCase):
|
||||
def test_no_file(self):
|
||||
"""Test load local infile when the file does not exist"""
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
c.execute("CREATE TABLE test_load_local (a INTEGER, b INTEGER)")
|
||||
try:
|
||||
self.assertRaises(
|
||||
OperationalError,
|
||||
c.execute,
|
||||
("LOAD DATA LOCAL INFILE 'no_data.txt' INTO TABLE "
|
||||
"test_load_local fields terminated by ','")
|
||||
)
|
||||
finally:
|
||||
c.execute("DROP TABLE test_load_local")
|
||||
c.close()
|
||||
|
||||
def test_load_file(self):
|
||||
"""Test load local infile with a valid file"""
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
c.execute("CREATE TABLE test_load_local (a INTEGER, b INTEGER)")
|
||||
filename = os.path.join(os.path.dirname(os.path.realpath(__file__)),
|
||||
'data',
|
||||
'load_local_data.txt')
|
||||
try:
|
||||
c.execute(
|
||||
("LOAD DATA LOCAL INFILE '{0}' INTO TABLE " +
|
||||
"test_load_local FIELDS TERMINATED BY ','").format(filename)
|
||||
)
|
||||
c.execute("SELECT COUNT(*) FROM test_load_local")
|
||||
self.assertEqual(22749, c.fetchone()[0])
|
||||
finally:
|
||||
c.execute("DROP TABLE test_load_local")
|
||||
|
||||
def test_unbuffered_load_file(self):
|
||||
"""Test unbuffered load local infile with a valid file"""
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor(cursors.SSCursor)
|
||||
c.execute("CREATE TABLE test_load_local (a INTEGER, b INTEGER)")
|
||||
filename = os.path.join(os.path.dirname(os.path.realpath(__file__)),
|
||||
'data',
|
||||
'load_local_data.txt')
|
||||
try:
|
||||
c.execute(
|
||||
("LOAD DATA LOCAL INFILE '{0}' INTO TABLE " +
|
||||
"test_load_local FIELDS TERMINATED BY ','").format(filename)
|
||||
)
|
||||
c.execute("SELECT COUNT(*) FROM test_load_local")
|
||||
self.assertEqual(22749, c.fetchone()[0])
|
||||
finally:
|
||||
c.close()
|
||||
conn.close()
|
||||
conn.connect()
|
||||
c = conn.cursor()
|
||||
c.execute("DROP TABLE test_load_local")
|
||||
|
||||
def test_load_warnings(self):
|
||||
"""Test load local infile produces the appropriate warnings"""
|
||||
conn = self.connections[0]
|
||||
c = conn.cursor()
|
||||
c.execute("CREATE TABLE test_load_local (a INTEGER, b INTEGER)")
|
||||
filename = os.path.join(os.path.dirname(os.path.realpath(__file__)),
|
||||
'data',
|
||||
'load_local_warn_data.txt')
|
||||
try:
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter('always')
|
||||
c.execute(
|
||||
("LOAD DATA LOCAL INFILE '{0}' INTO TABLE " +
|
||||
"test_load_local FIELDS TERMINATED BY ','").format(filename)
|
||||
)
|
||||
self.assertEqual(w[0].category, Warning)
|
||||
expected_message = "Incorrect integer value"
|
||||
if expected_message not in str(w[-1].message):
|
||||
self.fail("%r not in %r" % (expected_message, w[-1].message))
|
||||
finally:
|
||||
c.execute("DROP TABLE test_load_local")
|
||||
c.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
unittest.main()
|
||||
Executable
+68
@@ -0,0 +1,68 @@
|
||||
import unittest2
|
||||
|
||||
from pymysql.tests import base
|
||||
from pymysql import util
|
||||
|
||||
|
||||
class TestNextset(base.PyMySQLTestCase):
|
||||
|
||||
def setUp(self):
|
||||
super(TestNextset, self).setUp()
|
||||
self.con = self.connections[0]
|
||||
|
||||
def test_nextset(self):
|
||||
cur = self.con.cursor()
|
||||
cur.execute("SELECT 1; SELECT 2;")
|
||||
self.assertEqual([(1,)], list(cur))
|
||||
|
||||
r = cur.nextset()
|
||||
self.assertTrue(r)
|
||||
|
||||
self.assertEqual([(2,)], list(cur))
|
||||
self.assertIsNone(cur.nextset())
|
||||
|
||||
def test_skip_nextset(self):
|
||||
cur = self.con.cursor()
|
||||
cur.execute("SELECT 1; SELECT 2;")
|
||||
self.assertEqual([(1,)], list(cur))
|
||||
|
||||
cur.execute("SELECT 42")
|
||||
self.assertEqual([(42,)], list(cur))
|
||||
|
||||
def test_ok_and_next(self):
|
||||
cur = self.con.cursor()
|
||||
cur.execute("SELECT 1; commit; SELECT 2;")
|
||||
self.assertEqual([(1,)], list(cur))
|
||||
self.assertTrue(cur.nextset())
|
||||
self.assertTrue(cur.nextset())
|
||||
self.assertEqual([(2,)], list(cur))
|
||||
self.assertFalse(bool(cur.nextset()))
|
||||
|
||||
@unittest2.expectedFailure
|
||||
def test_multi_cursor(self):
|
||||
cur1 = self.con.cursor()
|
||||
cur2 = self.con.cursor()
|
||||
|
||||
cur1.execute("SELECT 1; SELECT 2;")
|
||||
cur2.execute("SELECT 42")
|
||||
|
||||
self.assertEqual([(1,)], list(cur1))
|
||||
self.assertEqual([(42,)], list(cur2))
|
||||
|
||||
r = cur1.nextset()
|
||||
self.assertTrue(r)
|
||||
|
||||
self.assertEqual([(2,)], list(cur1))
|
||||
self.assertIsNone(cur1.nextset())
|
||||
|
||||
def test_multi_statement_warnings(self):
|
||||
cursor = self.con.cursor()
|
||||
|
||||
try:
|
||||
cursor.execute('DROP TABLE IF EXISTS a; '
|
||||
'DROP TABLE IF EXISTS b;')
|
||||
except TypeError:
|
||||
self.fail()
|
||||
|
||||
#TODO: How about SSCursor and nextset?
|
||||
# It's very hard to implement correctly...
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
from pymysql.optionfile import Parser
|
||||
from unittest import TestCase
|
||||
from pymysql._compat import PY2
|
||||
|
||||
try:
|
||||
from cStringIO import StringIO
|
||||
except ImportError:
|
||||
from io import StringIO
|
||||
|
||||
|
||||
__all__ = ['TestParser']
|
||||
|
||||
|
||||
_cfg_file = (r"""
|
||||
[default]
|
||||
string = foo
|
||||
quoted = "bar"
|
||||
single_quoted = 'foobar'
|
||||
""")
|
||||
|
||||
|
||||
class TestParser(TestCase):
|
||||
|
||||
def test_string(self):
|
||||
parser = Parser()
|
||||
if PY2:
|
||||
parser.readfp(StringIO(_cfg_file))
|
||||
else:
|
||||
parser.read_file(StringIO(_cfg_file))
|
||||
self.assertEqual(parser.get("default", "string"), "foo")
|
||||
self.assertEqual(parser.get("default", "quoted"), "bar")
|
||||
self.assertEqual(parser.get("default", "single_quoted"), "foobar")
|
||||
+8
@@ -0,0 +1,8 @@
|
||||
from .test_MySQLdb import *
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
unittest.main()
|
||||
+7
@@ -0,0 +1,7 @@
|
||||
from .test_MySQLdb_capabilities import test_MySQLdb as test_capabilities
|
||||
from .test_MySQLdb_nonstandard import *
|
||||
from .test_MySQLdb_dbapi20 import test_MySQLdb as test_dbapi2
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
unittest.main()
|
||||
+298
@@ -0,0 +1,298 @@
|
||||
#!/usr/bin/env python -O
|
||||
""" Script to test database capabilities and the DB-API interface
|
||||
for functionality and memory leaks.
|
||||
|
||||
Adapted from a script by M-A Lemburg.
|
||||
|
||||
"""
|
||||
import sys
|
||||
from time import time
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
|
||||
PY2 = sys.version_info[0] == 2
|
||||
|
||||
class DatabaseTest(unittest.TestCase):
|
||||
|
||||
db_module = None
|
||||
connect_args = ()
|
||||
connect_kwargs = dict(use_unicode=True, charset="utf8")
|
||||
create_table_extra = "ENGINE=INNODB CHARACTER SET UTF8"
|
||||
rows = 10
|
||||
debug = False
|
||||
|
||||
def setUp(self):
|
||||
db = self.db_module.connect(*self.connect_args, **self.connect_kwargs)
|
||||
self.connection = db
|
||||
self.cursor = db.cursor()
|
||||
self.BLOBText = ''.join([chr(i) for i in range(256)] * 100);
|
||||
if PY2:
|
||||
self.BLOBUText = unicode().join(unichr(i) for i in range(16834))
|
||||
else:
|
||||
self.BLOBUText = "".join(chr(i) for i in range(16834))
|
||||
data = bytearray(range(256)) * 16
|
||||
self.BLOBBinary = self.db_module.Binary(data)
|
||||
|
||||
leak_test = True
|
||||
|
||||
def tearDown(self):
|
||||
if self.leak_test:
|
||||
import gc
|
||||
del self.cursor
|
||||
orphans = gc.collect()
|
||||
self.assertFalse(orphans, "%d orphaned objects found after deleting cursor" % orphans)
|
||||
|
||||
del self.connection
|
||||
orphans = gc.collect()
|
||||
self.assertFalse(orphans, "%d orphaned objects found after deleting connection" % orphans)
|
||||
|
||||
def table_exists(self, name):
|
||||
try:
|
||||
self.cursor.execute('select * from %s where 1=0' % name)
|
||||
except Exception:
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
|
||||
def quote_identifier(self, ident):
|
||||
return '"%s"' % ident
|
||||
|
||||
def new_table_name(self):
|
||||
i = id(self.cursor)
|
||||
while True:
|
||||
name = self.quote_identifier('tb%08x' % i)
|
||||
if not self.table_exists(name):
|
||||
return name
|
||||
i = i + 1
|
||||
|
||||
def create_table(self, columndefs):
|
||||
|
||||
""" Create a table using a list of column definitions given in
|
||||
columndefs.
|
||||
|
||||
generator must be a function taking arguments (row_number,
|
||||
col_number) returning a suitable data object for insertion
|
||||
into the table.
|
||||
|
||||
"""
|
||||
self.table = self.new_table_name()
|
||||
self.cursor.execute('CREATE TABLE %s (%s) %s' %
|
||||
(self.table,
|
||||
',\n'.join(columndefs),
|
||||
self.create_table_extra))
|
||||
|
||||
def check_data_integrity(self, columndefs, generator):
|
||||
# insert
|
||||
self.create_table(columndefs)
|
||||
insert_statement = ('INSERT INTO %s VALUES (%s)' %
|
||||
(self.table,
|
||||
','.join(['%s'] * len(columndefs))))
|
||||
data = [ [ generator(i,j) for j in range(len(columndefs)) ]
|
||||
for i in range(self.rows) ]
|
||||
if self.debug:
|
||||
print(data)
|
||||
self.cursor.executemany(insert_statement, data)
|
||||
self.connection.commit()
|
||||
# verify
|
||||
self.cursor.execute('select * from %s' % self.table)
|
||||
l = self.cursor.fetchall()
|
||||
if self.debug:
|
||||
print(l)
|
||||
self.assertEqual(len(l), self.rows)
|
||||
try:
|
||||
for i in range(self.rows):
|
||||
for j in range(len(columndefs)):
|
||||
self.assertEqual(l[i][j], generator(i,j))
|
||||
finally:
|
||||
if not self.debug:
|
||||
self.cursor.execute('drop table %s' % (self.table))
|
||||
|
||||
def test_transactions(self):
|
||||
columndefs = ( 'col1 INT', 'col2 VARCHAR(255)')
|
||||
def generator(row, col):
|
||||
if col == 0: return row
|
||||
else: return ('%i' % (row%10))*255
|
||||
self.create_table(columndefs)
|
||||
insert_statement = ('INSERT INTO %s VALUES (%s)' %
|
||||
(self.table,
|
||||
','.join(['%s'] * len(columndefs))))
|
||||
data = [ [ generator(i,j) for j in range(len(columndefs)) ]
|
||||
for i in range(self.rows) ]
|
||||
self.cursor.executemany(insert_statement, data)
|
||||
# verify
|
||||
self.connection.commit()
|
||||
self.cursor.execute('select * from %s' % self.table)
|
||||
l = self.cursor.fetchall()
|
||||
self.assertEqual(len(l), self.rows)
|
||||
for i in range(self.rows):
|
||||
for j in range(len(columndefs)):
|
||||
self.assertEqual(l[i][j], generator(i,j))
|
||||
delete_statement = 'delete from %s where col1=%%s' % self.table
|
||||
self.cursor.execute(delete_statement, (0,))
|
||||
self.cursor.execute('select col1 from %s where col1=%s' % \
|
||||
(self.table, 0))
|
||||
l = self.cursor.fetchall()
|
||||
self.assertFalse(l, "DELETE didn't work")
|
||||
self.connection.rollback()
|
||||
self.cursor.execute('select col1 from %s where col1=%s' % \
|
||||
(self.table, 0))
|
||||
l = self.cursor.fetchall()
|
||||
self.assertTrue(len(l) == 1, "ROLLBACK didn't work")
|
||||
self.cursor.execute('drop table %s' % (self.table))
|
||||
|
||||
def test_truncation(self):
|
||||
columndefs = ( 'col1 INT', 'col2 VARCHAR(255)')
|
||||
def generator(row, col):
|
||||
if col == 0: return row
|
||||
else: return ('%i' % (row%10))*((255-self.rows//2)+row)
|
||||
self.create_table(columndefs)
|
||||
insert_statement = ('INSERT INTO %s VALUES (%s)' %
|
||||
(self.table,
|
||||
','.join(['%s'] * len(columndefs))))
|
||||
|
||||
try:
|
||||
self.cursor.execute(insert_statement, (0, '0'*256))
|
||||
except Warning:
|
||||
if self.debug: print(self.cursor.messages)
|
||||
except self.connection.DataError:
|
||||
pass
|
||||
else:
|
||||
self.fail("Over-long column did not generate warnings/exception with single insert")
|
||||
|
||||
self.connection.rollback()
|
||||
|
||||
try:
|
||||
for i in range(self.rows):
|
||||
data = []
|
||||
for j in range(len(columndefs)):
|
||||
data.append(generator(i,j))
|
||||
self.cursor.execute(insert_statement,tuple(data))
|
||||
except Warning:
|
||||
if self.debug: print(self.cursor.messages)
|
||||
except self.connection.DataError:
|
||||
pass
|
||||
else:
|
||||
self.fail("Over-long columns did not generate warnings/exception with execute()")
|
||||
|
||||
self.connection.rollback()
|
||||
|
||||
try:
|
||||
data = [ [ generator(i,j) for j in range(len(columndefs)) ]
|
||||
for i in range(self.rows) ]
|
||||
self.cursor.executemany(insert_statement, data)
|
||||
except Warning:
|
||||
if self.debug: print(self.cursor.messages)
|
||||
except self.connection.DataError:
|
||||
pass
|
||||
else:
|
||||
self.fail("Over-long columns did not generate warnings/exception with executemany()")
|
||||
|
||||
self.connection.rollback()
|
||||
self.cursor.execute('drop table %s' % (self.table))
|
||||
|
||||
def test_CHAR(self):
|
||||
# Character data
|
||||
def generator(row,col):
|
||||
return ('%i' % ((row+col) % 10)) * 255
|
||||
self.check_data_integrity(
|
||||
('col1 char(255)','col2 char(255)'),
|
||||
generator)
|
||||
|
||||
def test_INT(self):
|
||||
# Number data
|
||||
def generator(row,col):
|
||||
return row*row
|
||||
self.check_data_integrity(
|
||||
('col1 INT',),
|
||||
generator)
|
||||
|
||||
def test_DECIMAL(self):
|
||||
# DECIMAL
|
||||
def generator(row,col):
|
||||
from decimal import Decimal
|
||||
return Decimal("%d.%02d" % (row, col))
|
||||
self.check_data_integrity(
|
||||
('col1 DECIMAL(5,2)',),
|
||||
generator)
|
||||
|
||||
def test_DATE(self):
|
||||
ticks = time()
|
||||
def generator(row,col):
|
||||
return self.db_module.DateFromTicks(ticks+row*86400-col*1313)
|
||||
self.check_data_integrity(
|
||||
('col1 DATE',),
|
||||
generator)
|
||||
|
||||
def test_TIME(self):
|
||||
ticks = time()
|
||||
def generator(row,col):
|
||||
return self.db_module.TimeFromTicks(ticks+row*86400-col*1313)
|
||||
self.check_data_integrity(
|
||||
('col1 TIME',),
|
||||
generator)
|
||||
|
||||
def test_DATETIME(self):
|
||||
ticks = time()
|
||||
def generator(row,col):
|
||||
return self.db_module.TimestampFromTicks(ticks+row*86400-col*1313)
|
||||
self.check_data_integrity(
|
||||
('col1 DATETIME',),
|
||||
generator)
|
||||
|
||||
def test_TIMESTAMP(self):
|
||||
ticks = time()
|
||||
def generator(row,col):
|
||||
return self.db_module.TimestampFromTicks(ticks+row*86400-col*1313)
|
||||
self.check_data_integrity(
|
||||
('col1 TIMESTAMP',),
|
||||
generator)
|
||||
|
||||
def test_fractional_TIMESTAMP(self):
|
||||
ticks = time()
|
||||
def generator(row,col):
|
||||
return self.db_module.TimestampFromTicks(ticks+row*86400-col*1313+row*0.7*col/3.0)
|
||||
self.check_data_integrity(
|
||||
('col1 TIMESTAMP',),
|
||||
generator)
|
||||
|
||||
def test_LONG(self):
|
||||
def generator(row,col):
|
||||
if col == 0:
|
||||
return row
|
||||
else:
|
||||
return self.BLOBUText # 'BLOB Text ' * 1024
|
||||
self.check_data_integrity(
|
||||
('col1 INT', 'col2 LONG'),
|
||||
generator)
|
||||
|
||||
def test_TEXT(self):
|
||||
def generator(row,col):
|
||||
if col == 0:
|
||||
return row
|
||||
else:
|
||||
return self.BLOBUText[:5192] # 'BLOB Text ' * 1024
|
||||
self.check_data_integrity(
|
||||
('col1 INT', 'col2 TEXT'),
|
||||
generator)
|
||||
|
||||
def test_LONG_BYTE(self):
|
||||
def generator(row,col):
|
||||
if col == 0:
|
||||
return row
|
||||
else:
|
||||
return self.BLOBBinary # 'BLOB\000Binary ' * 1024
|
||||
self.check_data_integrity(
|
||||
('col1 INT','col2 LONG BYTE'),
|
||||
generator)
|
||||
|
||||
def test_BLOB(self):
|
||||
def generator(row,col):
|
||||
if col == 0:
|
||||
return row
|
||||
else:
|
||||
return self.BLOBBinary # 'BLOB\000Binary ' * 1024
|
||||
self.check_data_integrity(
|
||||
('col1 INT','col2 BLOB'),
|
||||
generator)
|
||||
+856
@@ -0,0 +1,856 @@
|
||||
#!/usr/bin/env python
|
||||
''' Python DB API 2.0 driver compliance unit test suite.
|
||||
|
||||
This software is Public Domain and may be used without restrictions.
|
||||
|
||||
"Now we have booze and barflies entering the discussion, plus rumours of
|
||||
DBAs on drugs... and I won't tell you what flashes through my mind each
|
||||
time I read the subject line with 'Anal Compliance' in it. All around
|
||||
this is turning out to be a thoroughly unwholesome unit test."
|
||||
|
||||
-- Ian Bicking
|
||||
'''
|
||||
|
||||
__rcs_id__ = '$Id$'
|
||||
__version__ = '$Revision$'[11:-2]
|
||||
__author__ = 'Stuart Bishop <zen@shangri-la.dropbear.id.au>'
|
||||
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
|
||||
import time
|
||||
|
||||
# $Log$
|
||||
# Revision 1.1.2.1 2006/02/25 03:44:32 adustman
|
||||
# Generic DB-API unit test module
|
||||
#
|
||||
# Revision 1.10 2003/10/09 03:14:14 zenzen
|
||||
# Add test for DB API 2.0 optional extension, where database exceptions
|
||||
# are exposed as attributes on the Connection object.
|
||||
#
|
||||
# Revision 1.9 2003/08/13 01:16:36 zenzen
|
||||
# Minor tweak from Stefan Fleiter
|
||||
#
|
||||
# Revision 1.8 2003/04/10 00:13:25 zenzen
|
||||
# Changes, as per suggestions by M.-A. Lemburg
|
||||
# - Add a table prefix, to ensure namespace collisions can always be avoided
|
||||
#
|
||||
# Revision 1.7 2003/02/26 23:33:37 zenzen
|
||||
# Break out DDL into helper functions, as per request by David Rushby
|
||||
#
|
||||
# Revision 1.6 2003/02/21 03:04:33 zenzen
|
||||
# Stuff from Henrik Ekelund:
|
||||
# added test_None
|
||||
# added test_nextset & hooks
|
||||
#
|
||||
# Revision 1.5 2003/02/17 22:08:43 zenzen
|
||||
# Implement suggestions and code from Henrik Eklund - test that cursor.arraysize
|
||||
# defaults to 1 & generic cursor.callproc test added
|
||||
#
|
||||
# Revision 1.4 2003/02/15 00:16:33 zenzen
|
||||
# Changes, as per suggestions and bug reports by M.-A. Lemburg,
|
||||
# Matthew T. Kromer, Federico Di Gregorio and Daniel Dittmar
|
||||
# - Class renamed
|
||||
# - Now a subclass of TestCase, to avoid requiring the driver stub
|
||||
# to use multiple inheritance
|
||||
# - Reversed the polarity of buggy test in test_description
|
||||
# - Test exception heirarchy correctly
|
||||
# - self.populate is now self._populate(), so if a driver stub
|
||||
# overrides self.ddl1 this change propogates
|
||||
# - VARCHAR columns now have a width, which will hopefully make the
|
||||
# DDL even more portible (this will be reversed if it causes more problems)
|
||||
# - cursor.rowcount being checked after various execute and fetchXXX methods
|
||||
# - Check for fetchall and fetchmany returning empty lists after results
|
||||
# are exhausted (already checking for empty lists if select retrieved
|
||||
# nothing
|
||||
# - Fix bugs in test_setoutputsize_basic and test_setinputsizes
|
||||
#
|
||||
|
||||
class DatabaseAPI20Test(unittest.TestCase):
|
||||
''' Test a database self.driver for DB API 2.0 compatibility.
|
||||
This implementation tests Gadfly, but the TestCase
|
||||
is structured so that other self.drivers can subclass this
|
||||
test case to ensure compiliance with the DB-API. It is
|
||||
expected that this TestCase may be expanded in the future
|
||||
if ambiguities or edge conditions are discovered.
|
||||
|
||||
The 'Optional Extensions' are not yet being tested.
|
||||
|
||||
self.drivers should subclass this test, overriding setUp, tearDown,
|
||||
self.driver, connect_args and connect_kw_args. Class specification
|
||||
should be as follows:
|
||||
|
||||
import dbapi20
|
||||
class mytest(dbapi20.DatabaseAPI20Test):
|
||||
[...]
|
||||
|
||||
Don't 'import DatabaseAPI20Test from dbapi20', or you will
|
||||
confuse the unit tester - just 'import dbapi20'.
|
||||
'''
|
||||
|
||||
# The self.driver module. This should be the module where the 'connect'
|
||||
# method is to be found
|
||||
driver = None
|
||||
connect_args = () # List of arguments to pass to connect
|
||||
connect_kw_args = {} # Keyword arguments for connect
|
||||
table_prefix = 'dbapi20test_' # If you need to specify a prefix for tables
|
||||
|
||||
ddl1 = 'create table %sbooze (name varchar(20))' % table_prefix
|
||||
ddl2 = 'create table %sbarflys (name varchar(20))' % table_prefix
|
||||
xddl1 = 'drop table %sbooze' % table_prefix
|
||||
xddl2 = 'drop table %sbarflys' % table_prefix
|
||||
|
||||
lowerfunc = 'lower' # Name of stored procedure to convert string->lowercase
|
||||
|
||||
# Some drivers may need to override these helpers, for example adding
|
||||
# a 'commit' after the execute.
|
||||
def executeDDL1(self,cursor):
|
||||
cursor.execute(self.ddl1)
|
||||
|
||||
def executeDDL2(self,cursor):
|
||||
cursor.execute(self.ddl2)
|
||||
|
||||
def setUp(self):
|
||||
''' self.drivers should override this method to perform required setup
|
||||
if any is necessary, such as creating the database.
|
||||
'''
|
||||
pass
|
||||
|
||||
def tearDown(self):
|
||||
''' self.drivers should override this method to perform required cleanup
|
||||
if any is necessary, such as deleting the test database.
|
||||
The default drops the tables that may be created.
|
||||
'''
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
for ddl in (self.xddl1,self.xddl2):
|
||||
try:
|
||||
cur.execute(ddl)
|
||||
con.commit()
|
||||
except self.driver.Error:
|
||||
# Assume table didn't exist. Other tests will check if
|
||||
# execute is busted.
|
||||
pass
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def _connect(self):
|
||||
try:
|
||||
return self.driver.connect(
|
||||
*self.connect_args,**self.connect_kw_args
|
||||
)
|
||||
except AttributeError:
|
||||
self.fail("No connect method found in self.driver module")
|
||||
|
||||
def test_connect(self):
|
||||
con = self._connect()
|
||||
con.close()
|
||||
|
||||
def test_apilevel(self):
|
||||
try:
|
||||
# Must exist
|
||||
apilevel = self.driver.apilevel
|
||||
# Must equal 2.0
|
||||
self.assertEqual(apilevel,'2.0')
|
||||
except AttributeError:
|
||||
self.fail("Driver doesn't define apilevel")
|
||||
|
||||
def test_threadsafety(self):
|
||||
try:
|
||||
# Must exist
|
||||
threadsafety = self.driver.threadsafety
|
||||
# Must be a valid value
|
||||
self.assertTrue(threadsafety in (0,1,2,3))
|
||||
except AttributeError:
|
||||
self.fail("Driver doesn't define threadsafety")
|
||||
|
||||
def test_paramstyle(self):
|
||||
try:
|
||||
# Must exist
|
||||
paramstyle = self.driver.paramstyle
|
||||
# Must be a valid value
|
||||
self.assertTrue(paramstyle in (
|
||||
'qmark','numeric','named','format','pyformat'
|
||||
))
|
||||
except AttributeError:
|
||||
self.fail("Driver doesn't define paramstyle")
|
||||
|
||||
def test_Exceptions(self):
|
||||
# Make sure required exceptions exist, and are in the
|
||||
# defined heirarchy.
|
||||
self.assertTrue(issubclass(self.driver.Warning,Exception))
|
||||
self.assertTrue(issubclass(self.driver.Error,Exception))
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.InterfaceError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.DatabaseError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.OperationalError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.IntegrityError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.InternalError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.ProgrammingError,self.driver.Error)
|
||||
)
|
||||
self.assertTrue(
|
||||
issubclass(self.driver.NotSupportedError,self.driver.Error)
|
||||
)
|
||||
|
||||
def test_ExceptionsAsConnectionAttributes(self):
|
||||
# OPTIONAL EXTENSION
|
||||
# Test for the optional DB API 2.0 extension, where the exceptions
|
||||
# are exposed as attributes on the Connection object
|
||||
# I figure this optional extension will be implemented by any
|
||||
# driver author who is using this test suite, so it is enabled
|
||||
# by default.
|
||||
con = self._connect()
|
||||
drv = self.driver
|
||||
self.assertTrue(con.Warning is drv.Warning)
|
||||
self.assertTrue(con.Error is drv.Error)
|
||||
self.assertTrue(con.InterfaceError is drv.InterfaceError)
|
||||
self.assertTrue(con.DatabaseError is drv.DatabaseError)
|
||||
self.assertTrue(con.OperationalError is drv.OperationalError)
|
||||
self.assertTrue(con.IntegrityError is drv.IntegrityError)
|
||||
self.assertTrue(con.InternalError is drv.InternalError)
|
||||
self.assertTrue(con.ProgrammingError is drv.ProgrammingError)
|
||||
self.assertTrue(con.NotSupportedError is drv.NotSupportedError)
|
||||
|
||||
|
||||
def test_commit(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
# Commit must work, even if it doesn't do anything
|
||||
con.commit()
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_rollback(self):
|
||||
con = self._connect()
|
||||
# If rollback is defined, it should either work or throw
|
||||
# the documented exception
|
||||
if hasattr(con,'rollback'):
|
||||
try:
|
||||
con.rollback()
|
||||
except self.driver.NotSupportedError:
|
||||
pass
|
||||
|
||||
def test_cursor(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_cursor_isolation(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
# Make sure cursors created from the same connection have
|
||||
# the documented transaction isolation level
|
||||
cur1 = con.cursor()
|
||||
cur2 = con.cursor()
|
||||
self.executeDDL1(cur1)
|
||||
cur1.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
cur2.execute("select name from %sbooze" % self.table_prefix)
|
||||
booze = cur2.fetchall()
|
||||
self.assertEqual(len(booze),1)
|
||||
self.assertEqual(len(booze[0]),1)
|
||||
self.assertEqual(booze[0][0],'Victoria Bitter')
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_description(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
self.assertEqual(cur.description,None,
|
||||
'cursor.description should be none after executing a '
|
||||
'statement that can return no rows (such as DDL)'
|
||||
)
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
self.assertEqual(len(cur.description),1,
|
||||
'cursor.description describes too many columns'
|
||||
)
|
||||
self.assertEqual(len(cur.description[0]),7,
|
||||
'cursor.description[x] tuples must have 7 elements'
|
||||
)
|
||||
self.assertEqual(cur.description[0][0].lower(),'name',
|
||||
'cursor.description[x][0] must return column name'
|
||||
)
|
||||
self.assertEqual(cur.description[0][1],self.driver.STRING,
|
||||
'cursor.description[x][1] must return column type. Got %r'
|
||||
% cur.description[0][1]
|
||||
)
|
||||
|
||||
# Make sure self.description gets reset
|
||||
self.executeDDL2(cur)
|
||||
self.assertEqual(cur.description,None,
|
||||
'cursor.description not being set to None when executing '
|
||||
'no-result statements (eg. DDL)'
|
||||
)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_rowcount(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
self.assertEqual(cur.rowcount,-1,
|
||||
'cursor.rowcount should be -1 after executing no-result '
|
||||
'statements'
|
||||
)
|
||||
cur.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
self.assertTrue(cur.rowcount in (-1,1),
|
||||
'cursor.rowcount should == number or rows inserted, or '
|
||||
'set to -1 after executing an insert statement'
|
||||
)
|
||||
cur.execute("select name from %sbooze" % self.table_prefix)
|
||||
self.assertTrue(cur.rowcount in (-1,1),
|
||||
'cursor.rowcount should == number of rows returned, or '
|
||||
'set to -1 after executing a select statement'
|
||||
)
|
||||
self.executeDDL2(cur)
|
||||
self.assertEqual(cur.rowcount,-1,
|
||||
'cursor.rowcount not being reset to -1 after executing '
|
||||
'no-result statements'
|
||||
)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
lower_func = 'lower'
|
||||
def test_callproc(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
if self.lower_func and hasattr(cur,'callproc'):
|
||||
r = cur.callproc(self.lower_func,('FOO',))
|
||||
self.assertEqual(len(r),1)
|
||||
self.assertEqual(r[0],'FOO')
|
||||
r = cur.fetchall()
|
||||
self.assertEqual(len(r),1,'callproc produced no result set')
|
||||
self.assertEqual(len(r[0]),1,
|
||||
'callproc produced invalid result set'
|
||||
)
|
||||
self.assertEqual(r[0][0],'foo',
|
||||
'callproc produced invalid results'
|
||||
)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_close(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
# cursor.execute should raise an Error if called after connection
|
||||
# closed
|
||||
self.assertRaises(self.driver.Error,self.executeDDL1,cur)
|
||||
|
||||
# connection.commit should raise an Error if called after connection'
|
||||
# closed.'
|
||||
self.assertRaises(self.driver.Error,con.commit)
|
||||
|
||||
# connection.close should raise an Error if called more than once
|
||||
self.assertRaises(self.driver.Error,con.close)
|
||||
|
||||
def test_execute(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self._paraminsert(cur)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def _paraminsert(self,cur):
|
||||
self.executeDDL1(cur)
|
||||
cur.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
self.assertTrue(cur.rowcount in (-1,1))
|
||||
|
||||
if self.driver.paramstyle == 'qmark':
|
||||
cur.execute(
|
||||
'insert into %sbooze values (?)' % self.table_prefix,
|
||||
("Cooper's",)
|
||||
)
|
||||
elif self.driver.paramstyle == 'numeric':
|
||||
cur.execute(
|
||||
'insert into %sbooze values (:1)' % self.table_prefix,
|
||||
("Cooper's",)
|
||||
)
|
||||
elif self.driver.paramstyle == 'named':
|
||||
cur.execute(
|
||||
'insert into %sbooze values (:beer)' % self.table_prefix,
|
||||
{'beer':"Cooper's"}
|
||||
)
|
||||
elif self.driver.paramstyle == 'format':
|
||||
cur.execute(
|
||||
'insert into %sbooze values (%%s)' % self.table_prefix,
|
||||
("Cooper's",)
|
||||
)
|
||||
elif self.driver.paramstyle == 'pyformat':
|
||||
cur.execute(
|
||||
'insert into %sbooze values (%%(beer)s)' % self.table_prefix,
|
||||
{'beer':"Cooper's"}
|
||||
)
|
||||
else:
|
||||
self.fail('Invalid paramstyle')
|
||||
self.assertTrue(cur.rowcount in (-1,1))
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
res = cur.fetchall()
|
||||
self.assertEqual(len(res),2,'cursor.fetchall returned too few rows')
|
||||
beers = [res[0][0],res[1][0]]
|
||||
beers.sort()
|
||||
self.assertEqual(beers[0],"Cooper's",
|
||||
'cursor.fetchall retrieved incorrect data, or data inserted '
|
||||
'incorrectly'
|
||||
)
|
||||
self.assertEqual(beers[1],"Victoria Bitter",
|
||||
'cursor.fetchall retrieved incorrect data, or data inserted '
|
||||
'incorrectly'
|
||||
)
|
||||
|
||||
def test_executemany(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
largs = [ ("Cooper's",) , ("Boag's",) ]
|
||||
margs = [ {'beer': "Cooper's"}, {'beer': "Boag's"} ]
|
||||
if self.driver.paramstyle == 'qmark':
|
||||
cur.executemany(
|
||||
'insert into %sbooze values (?)' % self.table_prefix,
|
||||
largs
|
||||
)
|
||||
elif self.driver.paramstyle == 'numeric':
|
||||
cur.executemany(
|
||||
'insert into %sbooze values (:1)' % self.table_prefix,
|
||||
largs
|
||||
)
|
||||
elif self.driver.paramstyle == 'named':
|
||||
cur.executemany(
|
||||
'insert into %sbooze values (:beer)' % self.table_prefix,
|
||||
margs
|
||||
)
|
||||
elif self.driver.paramstyle == 'format':
|
||||
cur.executemany(
|
||||
'insert into %sbooze values (%%s)' % self.table_prefix,
|
||||
largs
|
||||
)
|
||||
elif self.driver.paramstyle == 'pyformat':
|
||||
cur.executemany(
|
||||
'insert into %sbooze values (%%(beer)s)' % (
|
||||
self.table_prefix
|
||||
),
|
||||
margs
|
||||
)
|
||||
else:
|
||||
self.fail('Unknown paramstyle')
|
||||
self.assertTrue(cur.rowcount in (-1,2),
|
||||
'insert using cursor.executemany set cursor.rowcount to '
|
||||
'incorrect value %r' % cur.rowcount
|
||||
)
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
res = cur.fetchall()
|
||||
self.assertEqual(len(res),2,
|
||||
'cursor.fetchall retrieved incorrect number of rows'
|
||||
)
|
||||
beers = [res[0][0],res[1][0]]
|
||||
beers.sort()
|
||||
self.assertEqual(beers[0],"Boag's",'incorrect data retrieved')
|
||||
self.assertEqual(beers[1],"Cooper's",'incorrect data retrieved')
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_fetchone(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
|
||||
# cursor.fetchone should raise an Error if called before
|
||||
# executing a select-type query
|
||||
self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
# cursor.fetchone should raise an Error if called after
|
||||
# executing a query that cannnot return rows
|
||||
self.executeDDL1(cur)
|
||||
self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
self.assertEqual(cur.fetchone(),None,
|
||||
'cursor.fetchone should return None if a query retrieves '
|
||||
'no rows'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,0))
|
||||
|
||||
# cursor.fetchone should raise an Error if called after
|
||||
# executing a query that cannnot return rows
|
||||
cur.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
r = cur.fetchone()
|
||||
self.assertEqual(len(r),1,
|
||||
'cursor.fetchone should have retrieved a single row'
|
||||
)
|
||||
self.assertEqual(r[0],'Victoria Bitter',
|
||||
'cursor.fetchone retrieved incorrect data'
|
||||
)
|
||||
self.assertEqual(cur.fetchone(),None,
|
||||
'cursor.fetchone should return None if no more rows available'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,1))
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
samples = [
|
||||
'Carlton Cold',
|
||||
'Carlton Draft',
|
||||
'Mountain Goat',
|
||||
'Redback',
|
||||
'Victoria Bitter',
|
||||
'XXXX'
|
||||
]
|
||||
|
||||
def _populate(self):
|
||||
''' Return a list of sql commands to setup the DB for the fetch
|
||||
tests.
|
||||
'''
|
||||
populate = [
|
||||
"insert into %sbooze values ('%s')" % (self.table_prefix,s)
|
||||
for s in self.samples
|
||||
]
|
||||
return populate
|
||||
|
||||
def test_fetchmany(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
|
||||
# cursor.fetchmany should raise an Error if called without
|
||||
#issuing a query
|
||||
self.assertRaises(self.driver.Error,cur.fetchmany,4)
|
||||
|
||||
self.executeDDL1(cur)
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
r = cur.fetchmany()
|
||||
self.assertEqual(len(r),1,
|
||||
'cursor.fetchmany retrieved incorrect number of rows, '
|
||||
'default of arraysize is one.'
|
||||
)
|
||||
cur.arraysize=10
|
||||
r = cur.fetchmany(3) # Should get 3 rows
|
||||
self.assertEqual(len(r),3,
|
||||
'cursor.fetchmany retrieved incorrect number of rows'
|
||||
)
|
||||
r = cur.fetchmany(4) # Should get 2 more
|
||||
self.assertEqual(len(r),2,
|
||||
'cursor.fetchmany retrieved incorrect number of rows'
|
||||
)
|
||||
r = cur.fetchmany(4) # Should be an empty sequence
|
||||
self.assertEqual(len(r),0,
|
||||
'cursor.fetchmany should return an empty sequence after '
|
||||
'results are exhausted'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,6))
|
||||
|
||||
# Same as above, using cursor.arraysize
|
||||
cur.arraysize=4
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
r = cur.fetchmany() # Should get 4 rows
|
||||
self.assertEqual(len(r),4,
|
||||
'cursor.arraysize not being honoured by fetchmany'
|
||||
)
|
||||
r = cur.fetchmany() # Should get 2 more
|
||||
self.assertEqual(len(r),2)
|
||||
r = cur.fetchmany() # Should be an empty sequence
|
||||
self.assertEqual(len(r),0)
|
||||
self.assertTrue(cur.rowcount in (-1,6))
|
||||
|
||||
cur.arraysize=6
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
rows = cur.fetchmany() # Should get all rows
|
||||
self.assertTrue(cur.rowcount in (-1,6))
|
||||
self.assertEqual(len(rows),6)
|
||||
self.assertEqual(len(rows),6)
|
||||
rows = [r[0] for r in rows]
|
||||
rows.sort()
|
||||
|
||||
# Make sure we get the right data back out
|
||||
for i in range(0,6):
|
||||
self.assertEqual(rows[i],self.samples[i],
|
||||
'incorrect data retrieved by cursor.fetchmany'
|
||||
)
|
||||
|
||||
rows = cur.fetchmany() # Should return an empty list
|
||||
self.assertEqual(len(rows),0,
|
||||
'cursor.fetchmany should return an empty sequence if '
|
||||
'called after the whole result set has been fetched'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,6))
|
||||
|
||||
self.executeDDL2(cur)
|
||||
cur.execute('select name from %sbarflys' % self.table_prefix)
|
||||
r = cur.fetchmany() # Should get empty sequence
|
||||
self.assertEqual(len(r),0,
|
||||
'cursor.fetchmany should return an empty sequence if '
|
||||
'query retrieved no rows'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,0))
|
||||
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_fetchall(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
# cursor.fetchall should raise an Error if called
|
||||
# without executing a query that may return rows (such
|
||||
# as a select)
|
||||
self.assertRaises(self.driver.Error, cur.fetchall)
|
||||
|
||||
self.executeDDL1(cur)
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
# cursor.fetchall should raise an Error if called
|
||||
# after executing a a statement that cannot return rows
|
||||
self.assertRaises(self.driver.Error,cur.fetchall)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
rows = cur.fetchall()
|
||||
self.assertTrue(cur.rowcount in (-1,len(self.samples)))
|
||||
self.assertEqual(len(rows),len(self.samples),
|
||||
'cursor.fetchall did not retrieve all rows'
|
||||
)
|
||||
rows = [r[0] for r in rows]
|
||||
rows.sort()
|
||||
for i in range(0,len(self.samples)):
|
||||
self.assertEqual(rows[i],self.samples[i],
|
||||
'cursor.fetchall retrieved incorrect rows'
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
self.assertEqual(
|
||||
len(rows),0,
|
||||
'cursor.fetchall should return an empty list if called '
|
||||
'after the whole result set has been fetched'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,len(self.samples)))
|
||||
|
||||
self.executeDDL2(cur)
|
||||
cur.execute('select name from %sbarflys' % self.table_prefix)
|
||||
rows = cur.fetchall()
|
||||
self.assertTrue(cur.rowcount in (-1,0))
|
||||
self.assertEqual(len(rows),0,
|
||||
'cursor.fetchall should return an empty list if '
|
||||
'a select query returns no rows'
|
||||
)
|
||||
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_mixedfetch(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
rows1 = cur.fetchone()
|
||||
rows23 = cur.fetchmany(2)
|
||||
rows4 = cur.fetchone()
|
||||
rows56 = cur.fetchall()
|
||||
self.assertTrue(cur.rowcount in (-1,6))
|
||||
self.assertEqual(len(rows23),2,
|
||||
'fetchmany returned incorrect number of rows'
|
||||
)
|
||||
self.assertEqual(len(rows56),2,
|
||||
'fetchall returned incorrect number of rows'
|
||||
)
|
||||
|
||||
rows = [rows1[0]]
|
||||
rows.extend([rows23[0][0],rows23[1][0]])
|
||||
rows.append(rows4[0])
|
||||
rows.extend([rows56[0][0],rows56[1][0]])
|
||||
rows.sort()
|
||||
for i in range(0,len(self.samples)):
|
||||
self.assertEqual(rows[i],self.samples[i],
|
||||
'incorrect data retrieved or inserted'
|
||||
)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def help_nextset_setUp(self,cur):
|
||||
''' Should create a procedure called deleteme
|
||||
that returns two result sets, first the
|
||||
number of rows in booze then "name from booze"
|
||||
'''
|
||||
raise NotImplementedError('Helper not implemented')
|
||||
#sql="""
|
||||
# create procedure deleteme as
|
||||
# begin
|
||||
# select count(*) from booze
|
||||
# select name from booze
|
||||
# end
|
||||
#"""
|
||||
#cur.execute(sql)
|
||||
|
||||
def help_nextset_tearDown(self,cur):
|
||||
'If cleaning up is needed after nextSetTest'
|
||||
raise NotImplementedError('Helper not implemented')
|
||||
#cur.execute("drop procedure deleteme")
|
||||
|
||||
def test_nextset(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
if not hasattr(cur,'nextset'):
|
||||
return
|
||||
|
||||
try:
|
||||
self.executeDDL1(cur)
|
||||
sql=self._populate()
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
self.help_nextset_setUp(cur)
|
||||
|
||||
cur.callproc('deleteme')
|
||||
numberofrows=cur.fetchone()
|
||||
assert numberofrows[0]== len(self.samples)
|
||||
assert cur.nextset()
|
||||
names=cur.fetchall()
|
||||
assert len(names) == len(self.samples)
|
||||
s=cur.nextset()
|
||||
assert s == None,'No more return sets, should return None'
|
||||
finally:
|
||||
self.help_nextset_tearDown(cur)
|
||||
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_nextset(self):
|
||||
raise NotImplementedError('Drivers need to override this test')
|
||||
|
||||
def test_arraysize(self):
|
||||
# Not much here - rest of the tests for this are in test_fetchmany
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.assertTrue(hasattr(cur,'arraysize'),
|
||||
'cursor.arraysize must be defined'
|
||||
)
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_setinputsizes(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
cur.setinputsizes( (25,) )
|
||||
self._paraminsert(cur) # Make sure cursor still works
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_setoutputsize_basic(self):
|
||||
# Basic test is to make sure setoutputsize doesn't blow up
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
cur.setoutputsize(1000)
|
||||
cur.setoutputsize(2000,0)
|
||||
self._paraminsert(cur) # Make sure the cursor still works
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_setoutputsize(self):
|
||||
# Real test for setoutputsize is driver dependant
|
||||
raise NotImplementedError('Driver need to override this test')
|
||||
|
||||
def test_None(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
cur.execute('insert into %sbooze values (NULL)' % self.table_prefix)
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
r = cur.fetchall()
|
||||
self.assertEqual(len(r),1)
|
||||
self.assertEqual(len(r[0]),1)
|
||||
self.assertEqual(r[0][0],None,'NULL value not returned as None')
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_Date(self):
|
||||
d1 = self.driver.Date(2002,12,25)
|
||||
d2 = self.driver.DateFromTicks(time.mktime((2002,12,25,0,0,0,0,0,0)))
|
||||
# Can we assume this? API doesn't specify, but it seems implied
|
||||
# self.assertEqual(str(d1),str(d2))
|
||||
|
||||
def test_Time(self):
|
||||
t1 = self.driver.Time(13,45,30)
|
||||
t2 = self.driver.TimeFromTicks(time.mktime((2001,1,1,13,45,30,0,0,0)))
|
||||
# Can we assume this? API doesn't specify, but it seems implied
|
||||
# self.assertEqual(str(t1),str(t2))
|
||||
|
||||
def test_Timestamp(self):
|
||||
t1 = self.driver.Timestamp(2002,12,25,13,45,30)
|
||||
t2 = self.driver.TimestampFromTicks(
|
||||
time.mktime((2002,12,25,13,45,30,0,0,0))
|
||||
)
|
||||
# Can we assume this? API doesn't specify, but it seems implied
|
||||
# self.assertEqual(str(t1),str(t2))
|
||||
|
||||
def test_Binary(self):
|
||||
b = self.driver.Binary(b'Something')
|
||||
b = self.driver.Binary(b'')
|
||||
|
||||
def test_STRING(self):
|
||||
self.assertTrue(hasattr(self.driver,'STRING'),
|
||||
'module.STRING must be defined'
|
||||
)
|
||||
|
||||
def test_BINARY(self):
|
||||
self.assertTrue(hasattr(self.driver,'BINARY'),
|
||||
'module.BINARY must be defined.'
|
||||
)
|
||||
|
||||
def test_NUMBER(self):
|
||||
self.assertTrue(hasattr(self.driver,'NUMBER'),
|
||||
'module.NUMBER must be defined.'
|
||||
)
|
||||
|
||||
def test_DATETIME(self):
|
||||
self.assertTrue(hasattr(self.driver,'DATETIME'),
|
||||
'module.DATETIME must be defined.'
|
||||
)
|
||||
|
||||
def test_ROWID(self):
|
||||
self.assertTrue(hasattr(self.driver,'ROWID'),
|
||||
'module.ROWID must be defined.'
|
||||
)
|
||||
Vendored
Executable
+109
@@ -0,0 +1,109 @@
|
||||
#!/usr/bin/env python
|
||||
from . import capabilities
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
import pymysql
|
||||
from pymysql.tests import base
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('error')
|
||||
|
||||
class test_MySQLdb(capabilities.DatabaseTest):
|
||||
|
||||
db_module = pymysql
|
||||
connect_args = ()
|
||||
connect_kwargs = base.PyMySQLTestCase.databases[0].copy()
|
||||
connect_kwargs.update(dict(read_default_file='~/.my.cnf',
|
||||
use_unicode=True,
|
||||
charset='utf8', sql_mode="ANSI,STRICT_TRANS_TABLES,TRADITIONAL"))
|
||||
|
||||
create_table_extra = "ENGINE=INNODB CHARACTER SET UTF8"
|
||||
leak_test = False
|
||||
|
||||
def quote_identifier(self, ident):
|
||||
return "`%s`" % ident
|
||||
|
||||
def test_TIME(self):
|
||||
from datetime import timedelta
|
||||
def generator(row,col):
|
||||
return timedelta(0, row*8000)
|
||||
self.check_data_integrity(
|
||||
('col1 TIME',),
|
||||
generator)
|
||||
|
||||
def test_TINYINT(self):
|
||||
# Number data
|
||||
def generator(row,col):
|
||||
v = (row*row) % 256
|
||||
if v > 127:
|
||||
v = v-256
|
||||
return v
|
||||
self.check_data_integrity(
|
||||
('col1 TINYINT',),
|
||||
generator)
|
||||
|
||||
def test_stored_procedures(self):
|
||||
db = self.connection
|
||||
c = self.cursor
|
||||
try:
|
||||
self.create_table(('pos INT', 'tree CHAR(20)'))
|
||||
c.executemany("INSERT INTO %s (pos,tree) VALUES (%%s,%%s)" % self.table,
|
||||
list(enumerate('ash birch cedar larch pine'.split())))
|
||||
db.commit()
|
||||
|
||||
c.execute("""
|
||||
CREATE PROCEDURE test_sp(IN t VARCHAR(255))
|
||||
BEGIN
|
||||
SELECT pos FROM %s WHERE tree = t;
|
||||
END
|
||||
""" % self.table)
|
||||
db.commit()
|
||||
|
||||
c.callproc('test_sp', ('larch',))
|
||||
rows = c.fetchall()
|
||||
self.assertEqual(len(rows), 1)
|
||||
self.assertEqual(rows[0][0], 3)
|
||||
c.nextset()
|
||||
finally:
|
||||
c.execute("DROP PROCEDURE IF EXISTS test_sp")
|
||||
c.execute('drop table %s' % (self.table))
|
||||
|
||||
def test_small_CHAR(self):
|
||||
# Character data
|
||||
def generator(row,col):
|
||||
i = ((row+1)*(col+1)+62)%256
|
||||
if i == 62: return ''
|
||||
if i == 63: return None
|
||||
return chr(i)
|
||||
self.check_data_integrity(
|
||||
('col1 char(1)','col2 char(1)'),
|
||||
generator)
|
||||
|
||||
def test_bug_2671682(self):
|
||||
from pymysql.constants import ER
|
||||
try:
|
||||
self.cursor.execute("describe some_non_existent_table");
|
||||
except self.connection.ProgrammingError as msg:
|
||||
self.assertEqual(msg.args[0], ER.NO_SUCH_TABLE)
|
||||
|
||||
def test_ping(self):
|
||||
self.connection.ping()
|
||||
|
||||
def test_literal_int(self):
|
||||
self.assertTrue("2" == self.connection.literal(2))
|
||||
|
||||
def test_literal_float(self):
|
||||
self.assertTrue("3.1415" == self.connection.literal(3.1415))
|
||||
|
||||
def test_literal_string(self):
|
||||
self.assertTrue("'foo'" == self.connection.literal("foo"))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
if test_MySQLdb.leak_test:
|
||||
import gc
|
||||
gc.enable()
|
||||
gc.set_debug(gc.DEBUG_LEAK)
|
||||
unittest.main()
|
||||
+210
@@ -0,0 +1,210 @@
|
||||
#!/usr/bin/env python
|
||||
from . import dbapi20
|
||||
import pymysql
|
||||
from pymysql.tests import base
|
||||
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
|
||||
|
||||
class test_MySQLdb(dbapi20.DatabaseAPI20Test):
|
||||
driver = pymysql
|
||||
connect_args = ()
|
||||
connect_kw_args = base.PyMySQLTestCase.databases[0].copy()
|
||||
connect_kw_args.update(dict(read_default_file='~/.my.cnf',
|
||||
charset='utf8',
|
||||
sql_mode="ANSI,STRICT_TRANS_TABLES,TRADITIONAL"))
|
||||
|
||||
def test_setoutputsize(self): pass
|
||||
def test_setoutputsize_basic(self): pass
|
||||
def test_nextset(self): pass
|
||||
|
||||
"""The tests on fetchone and fetchall and rowcount bogusly
|
||||
test for an exception if the statement cannot return a
|
||||
result set. MySQL always returns a result set; it's just that
|
||||
some things return empty result sets."""
|
||||
|
||||
def test_fetchall(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
# cursor.fetchall should raise an Error if called
|
||||
# without executing a query that may return rows (such
|
||||
# as a select)
|
||||
self.assertRaises(self.driver.Error, cur.fetchall)
|
||||
|
||||
self.executeDDL1(cur)
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
# cursor.fetchall should raise an Error if called
|
||||
# after executing a a statement that cannot return rows
|
||||
## self.assertRaises(self.driver.Error,cur.fetchall)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
rows = cur.fetchall()
|
||||
self.assertTrue(cur.rowcount in (-1,len(self.samples)))
|
||||
self.assertEqual(len(rows),len(self.samples),
|
||||
'cursor.fetchall did not retrieve all rows'
|
||||
)
|
||||
rows = [r[0] for r in rows]
|
||||
rows.sort()
|
||||
for i in range(0,len(self.samples)):
|
||||
self.assertEqual(rows[i],self.samples[i],
|
||||
'cursor.fetchall retrieved incorrect rows'
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
self.assertEqual(
|
||||
len(rows),0,
|
||||
'cursor.fetchall should return an empty list if called '
|
||||
'after the whole result set has been fetched'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,len(self.samples)))
|
||||
|
||||
self.executeDDL2(cur)
|
||||
cur.execute('select name from %sbarflys' % self.table_prefix)
|
||||
rows = cur.fetchall()
|
||||
self.assertTrue(cur.rowcount in (-1,0))
|
||||
self.assertEqual(len(rows),0,
|
||||
'cursor.fetchall should return an empty list if '
|
||||
'a select query returns no rows'
|
||||
)
|
||||
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_fetchone(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
|
||||
# cursor.fetchone should raise an Error if called before
|
||||
# executing a select-type query
|
||||
self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
# cursor.fetchone should raise an Error if called after
|
||||
# executing a query that cannnot return rows
|
||||
self.executeDDL1(cur)
|
||||
## self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
self.assertEqual(cur.fetchone(),None,
|
||||
'cursor.fetchone should return None if a query retrieves '
|
||||
'no rows'
|
||||
)
|
||||
self.assertTrue(cur.rowcount in (-1,0))
|
||||
|
||||
# cursor.fetchone should raise an Error if called after
|
||||
# executing a query that cannnot return rows
|
||||
cur.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
## self.assertRaises(self.driver.Error,cur.fetchone)
|
||||
|
||||
cur.execute('select name from %sbooze' % self.table_prefix)
|
||||
r = cur.fetchone()
|
||||
self.assertEqual(len(r),1,
|
||||
'cursor.fetchone should have retrieved a single row'
|
||||
)
|
||||
self.assertEqual(r[0],'Victoria Bitter',
|
||||
'cursor.fetchone retrieved incorrect data'
|
||||
)
|
||||
## self.assertEqual(cur.fetchone(),None,
|
||||
## 'cursor.fetchone should return None if no more rows available'
|
||||
## )
|
||||
self.assertTrue(cur.rowcount in (-1,1))
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
# Same complaint as for fetchall and fetchone
|
||||
def test_rowcount(self):
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
self.executeDDL1(cur)
|
||||
## self.assertEqual(cur.rowcount,-1,
|
||||
## 'cursor.rowcount should be -1 after executing no-result '
|
||||
## 'statements'
|
||||
## )
|
||||
cur.execute("insert into %sbooze values ('Victoria Bitter')" % (
|
||||
self.table_prefix
|
||||
))
|
||||
## self.assertTrue(cur.rowcount in (-1,1),
|
||||
## 'cursor.rowcount should == number or rows inserted, or '
|
||||
## 'set to -1 after executing an insert statement'
|
||||
## )
|
||||
cur.execute("select name from %sbooze" % self.table_prefix)
|
||||
self.assertTrue(cur.rowcount in (-1,1),
|
||||
'cursor.rowcount should == number of rows returned, or '
|
||||
'set to -1 after executing a select statement'
|
||||
)
|
||||
self.executeDDL2(cur)
|
||||
## self.assertEqual(cur.rowcount,-1,
|
||||
## 'cursor.rowcount not being reset to -1 after executing '
|
||||
## 'no-result statements'
|
||||
## )
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
def test_callproc(self):
|
||||
pass # performed in test_MySQL_capabilities
|
||||
|
||||
def help_nextset_setUp(self,cur):
|
||||
''' Should create a procedure called deleteme
|
||||
that returns two result sets, first the
|
||||
number of rows in booze then "name from booze"
|
||||
'''
|
||||
sql="""
|
||||
create procedure deleteme()
|
||||
begin
|
||||
select count(*) from %(tp)sbooze;
|
||||
select name from %(tp)sbooze;
|
||||
end
|
||||
""" % dict(tp=self.table_prefix)
|
||||
cur.execute(sql)
|
||||
|
||||
def help_nextset_tearDown(self,cur):
|
||||
'If cleaning up is needed after nextSetTest'
|
||||
cur.execute("drop procedure deleteme")
|
||||
|
||||
def test_nextset(self):
|
||||
from warnings import warn
|
||||
con = self._connect()
|
||||
try:
|
||||
cur = con.cursor()
|
||||
if not hasattr(cur,'nextset'):
|
||||
return
|
||||
|
||||
try:
|
||||
self.executeDDL1(cur)
|
||||
sql=self._populate()
|
||||
for sql in self._populate():
|
||||
cur.execute(sql)
|
||||
|
||||
self.help_nextset_setUp(cur)
|
||||
|
||||
cur.callproc('deleteme')
|
||||
numberofrows=cur.fetchone()
|
||||
assert numberofrows[0]== len(self.samples)
|
||||
assert cur.nextset()
|
||||
names=cur.fetchall()
|
||||
assert len(names) == len(self.samples)
|
||||
s=cur.nextset()
|
||||
if s:
|
||||
empty = cur.fetchall()
|
||||
self.assertEqual(len(empty), 0,
|
||||
"non-empty result set after other result sets")
|
||||
#warn("Incompatibility: MySQL returns an empty result set for the CALL itself",
|
||||
# Warning)
|
||||
#assert s == None,'No more return sets, should return None'
|
||||
finally:
|
||||
self.help_nextset_tearDown(cur)
|
||||
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
+101
@@ -0,0 +1,101 @@
|
||||
import sys
|
||||
try:
|
||||
import unittest2 as unittest
|
||||
except ImportError:
|
||||
import unittest
|
||||
|
||||
import pymysql
|
||||
_mysql = pymysql
|
||||
from pymysql.constants import FIELD_TYPE
|
||||
from pymysql.tests import base
|
||||
from pymysql._compat import PY2, long_type
|
||||
|
||||
if not PY2:
|
||||
basestring = str
|
||||
|
||||
|
||||
class TestDBAPISet(unittest.TestCase):
|
||||
def test_set_equality(self):
|
||||
self.assertTrue(pymysql.STRING == pymysql.STRING)
|
||||
|
||||
def test_set_inequality(self):
|
||||
self.assertTrue(pymysql.STRING != pymysql.NUMBER)
|
||||
|
||||
def test_set_equality_membership(self):
|
||||
self.assertTrue(FIELD_TYPE.VAR_STRING == pymysql.STRING)
|
||||
|
||||
def test_set_inequality_membership(self):
|
||||
self.assertTrue(FIELD_TYPE.DATE != pymysql.STRING)
|
||||
|
||||
|
||||
class CoreModule(unittest.TestCase):
|
||||
"""Core _mysql module features."""
|
||||
|
||||
def test_NULL(self):
|
||||
"""Should have a NULL constant."""
|
||||
self.assertEqual(_mysql.NULL, 'NULL')
|
||||
|
||||
def test_version(self):
|
||||
"""Version information sanity."""
|
||||
self.assertTrue(isinstance(_mysql.__version__, basestring))
|
||||
|
||||
self.assertTrue(isinstance(_mysql.version_info, tuple))
|
||||
self.assertEqual(len(_mysql.version_info), 5)
|
||||
|
||||
def test_client_info(self):
|
||||
self.assertTrue(isinstance(_mysql.get_client_info(), basestring))
|
||||
|
||||
def test_thread_safe(self):
|
||||
self.assertTrue(isinstance(_mysql.thread_safe(), int))
|
||||
|
||||
|
||||
class CoreAPI(unittest.TestCase):
|
||||
"""Test _mysql interaction internals."""
|
||||
|
||||
def setUp(self):
|
||||
kwargs = base.PyMySQLTestCase.databases[0].copy()
|
||||
kwargs["read_default_file"] = "~/.my.cnf"
|
||||
self.conn = _mysql.connect(**kwargs)
|
||||
|
||||
def tearDown(self):
|
||||
self.conn.close()
|
||||
|
||||
def test_thread_id(self):
|
||||
tid = self.conn.thread_id()
|
||||
self.assertTrue(isinstance(tid, (int, long_type)),
|
||||
"thread_id didn't return an integral value.")
|
||||
|
||||
self.assertRaises(TypeError, self.conn.thread_id, ('evil',),
|
||||
"thread_id shouldn't accept arguments.")
|
||||
|
||||
def test_affected_rows(self):
|
||||
self.assertEqual(self.conn.affected_rows(), 0,
|
||||
"Should return 0 before we do anything.")
|
||||
|
||||
|
||||
#def test_debug(self):
|
||||
## FIXME Only actually tests if you lack SUPER
|
||||
#self.assertRaises(pymysql.OperationalError,
|
||||
#self.conn.dump_debug_info)
|
||||
|
||||
def test_charset_name(self):
|
||||
self.assertTrue(isinstance(self.conn.character_set_name(), basestring),
|
||||
"Should return a string.")
|
||||
|
||||
def test_host_info(self):
|
||||
assert isinstance(self.conn.get_host_info(), basestring), "should return a string"
|
||||
|
||||
def test_proto_info(self):
|
||||
self.assertTrue(isinstance(self.conn.get_proto_info(), int),
|
||||
"Should return an int.")
|
||||
|
||||
def test_server_info(self):
|
||||
if sys.version_info[0] == 2:
|
||||
self.assertTrue(isinstance(self.conn.get_server_info(), basestring),
|
||||
"Should return an str.")
|
||||
else:
|
||||
self.assertTrue(isinstance(self.conn.get_server_info(), basestring),
|
||||
"Should return an str.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,16 +1,20 @@
|
||||
from time import localtime
|
||||
from datetime import date, datetime, time, timedelta
|
||||
|
||||
|
||||
Date = date
|
||||
Time = time
|
||||
TimeDelta = timedelta
|
||||
Timestamp = datetime
|
||||
|
||||
|
||||
def DateFromTicks(ticks):
|
||||
return date(*localtime(ticks)[:3])
|
||||
|
||||
|
||||
def TimeFromTicks(ticks):
|
||||
return time(*localtime(ticks)[3:6])
|
||||
|
||||
|
||||
def TimestampFromTicks(ticks):
|
||||
return datetime(*localtime(ticks)[:6])
|
||||
|
||||
@@ -1,14 +1,17 @@
|
||||
import struct
|
||||
|
||||
|
||||
def byte2int(b):
|
||||
if isinstance(b, int):
|
||||
return b
|
||||
else:
|
||||
return struct.unpack("!B", b)[0]
|
||||
|
||||
|
||||
def int2byte(i):
|
||||
return struct.pack("!B", i)
|
||||
|
||||
|
||||
def join_bytes(bs):
|
||||
if len(bs) == 0:
|
||||
return ""
|
||||
|
||||
Reference in New Issue
Block a user