Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/pymysql/connections.py: 26%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1# Python implementation of the MySQL client-server protocol
2# http://dev.mysql.com/doc/internals/en/client-server-protocol.html
3# Error codes:
4# https://dev.mysql.com/doc/refman/5.5/en/error-handling.html
5from __future__ import annotations
7import contextlib
8import errno
9import os
10import socket
11import struct
12import sys
13import traceback
14import warnings
16from . import VERSION_STRING, _auth, converters, err
17from .charset import charset_by_id, charset_by_name
18from .constants import CLIENT, COMMAND, CR, ER, FIELD_TYPE, SERVER_STATUS
19from .cursors import Cursor
20from .optionfile import Parser
21from .protocol import (
22 EOFPacketWrapper,
23 FieldDescriptorPacket,
24 LoadLocalPacketWrapper,
25 MysqlPacket,
26 OKPacketWrapper,
27 dump_packet,
28)
30try:
31 import ssl
33 SSL_ENABLED = True
34except ImportError:
35 ssl = None
36 SSL_ENABLED = False
38try:
39 import getpass
41 DEFAULT_USER = getpass.getuser()
42 del getpass
43except (ImportError, KeyError, OSError):
44 # When there's no entry in OS database for a current user:
45 # KeyError is raised in Python 3.12 and below.
46 # OSError is raised in Python 3.13+
47 DEFAULT_USER = None
49DEBUG = False
50_DEFAULT_AUTH_PLUGIN = None # if this is not None, use it instead of server's default.
52TEXT_TYPES = {
53 FIELD_TYPE.BIT,
54 FIELD_TYPE.BLOB,
55 FIELD_TYPE.LONG_BLOB,
56 FIELD_TYPE.MEDIUM_BLOB,
57 FIELD_TYPE.STRING,
58 FIELD_TYPE.TINY_BLOB,
59 FIELD_TYPE.VAR_STRING,
60 FIELD_TYPE.VARCHAR,
61 FIELD_TYPE.GEOMETRY,
62}
65DEFAULT_CHARSET = "utf8mb4"
67MAX_PACKET_LEN = 2**24 - 1
70def _pack_int24(n):
71 return struct.pack("<I", n)[:3]
74# https://dev.mysql.com/doc/internals/en/integer.html#packet-Protocol::LengthEncodedInteger
75def _lenenc_int(i):
76 if i < 0:
77 raise ValueError(
78 "Encoding %d is less than 0 - no representation in LengthEncodedInteger" % i
79 )
80 elif i < 0xFB:
81 return bytes([i])
82 elif i < (1 << 16):
83 return b"\xfc" + struct.pack("<H", i)
84 elif i < (1 << 24):
85 return b"\xfd" + struct.pack("<I", i)[:3]
86 elif i < (1 << 64):
87 return b"\xfe" + struct.pack("<Q", i)
88 else:
89 raise ValueError(
90 f"Encoding {i:x} is larger than {1 << 64:x} - no representation in LengthEncodedInteger"
91 )
94class Connection:
95 """
96 Representation of a socket with a mysql server.
98 The proper way to get an instance of this class is to call
99 connect().
101 Establish a connection to the MySQL database. Accepts several
102 arguments:
104 :param host: Host where the database server is located.
105 :param user: Username to log in as.
106 :param password: Password to use.
107 :param database: Database to use, None to not use a particular one.
108 :param port: MySQL port to use, default is usually OK. (default: 3306)
109 :param bind_address: When the client has multiple network interfaces, specify
110 the interface from which to connect to the host. Argument can be
111 a hostname or an IP address.
112 :param unix_socket: Use a unix socket rather than TCP/IP.
113 :param read_timeout: The timeout for reading from the connection in seconds.
114 (default: None - no timeout)
115 :param write_timeout: The timeout for writing to the connection in seconds.
116 (default: None - no timeout)
117 :param str charset: Charset to use. "utf8" (or "utf8mb4") is recommended.
118 legacy multibyte encodings may pose security risks.
119 Do not use such encodings for public-facing systems.
120 :param str collation: Collation name to use.
121 :param sql_mode: Default SQL_MODE to use.
122 :param read_default_file:
123 Specifies my.cnf file to read these parameters from under the [client] section.
124 :param conv:
125 Conversion dictionary to use instead of the default one.
126 This is used to provide custom marshalling and unmarshalling of types.
127 See converters.
128 :param use_unicode:
129 Whether or not to default to unicode strings.
130 This option defaults to true.
131 :param client_flag: Custom flags to send to MySQL. Find potential values in constants.CLIENT.
132 :param cursorclass: Custom cursor class to use.
133 :param init_command: Initial SQL statement to run when connection is established.
134 :param connect_timeout: The timeout for connecting to the database in seconds.
135 (default: 10, min: 1, max: 31536000)
136 :param ssl: An ssl.SSLContext, or a dict of arguments similar to mysql_ssl_set()'s parameters.
137 Passing a dict is deprecated; use the individual ``ssl_*`` parameters or an
138 ``ssl.SSLContext`` instead.
139 :param ssl_ca: Path to the file that contains a PEM-formatted CA certificate.
140 :param ssl_cert: Path to the file that contains a PEM-formatted client certificate.
141 :param ssl_disabled: A boolean value that disables usage of TLS. Unlike other SSL options,
142 setting this to True explicitly prohibits the use of TLS, even if the server supports it.
143 :param ssl_key: Path to the file that contains a PEM-formatted private key for
144 the client certificate.
145 :param ssl_key_password: The password for the client certificate private key.
146 :param ssl_verify_cert: Set to true to check the server certificate's validity.
147 :param ssl_verify_identity: Set to true to check the server's identity.
148 :param read_default_group: Group to read from in the configuration file.
149 :param autocommit: Autocommit mode. None means use server default. (default: False)
150 :param local_infile: Boolean to enable the use of LOAD DATA LOCAL command. (default: False)
151 :param max_allowed_packet: Max size of packet sent to server in bytes. (default: 16MB)
152 Only used to limit size of "LOAD LOCAL INFILE" data packet smaller than default (16KB).
153 :param defer_connect: Don't explicitly connect on construction - wait for connect call.
154 (default: False)
155 :param auth_plugin_map: A dict of plugin names to a class that processes that plugin.
156 The class will take the Connection object as the argument to the constructor.
157 The class needs an authenticate method taking an authentication packet as
158 an argument. For the dialog plugin, a prompt(echo, prompt) method can be used
159 (if no authenticate method) for returning a string from the user. (experimental)
160 :param server_public_key: SHA256 authentication plugin public key value. (default: None)
161 :param binary_prefix: **DEPRECATED**
162 :param compress: Not supported.
163 :param named_pipe: Not supported.
164 :param db: **DEPRECATED** Alias for database.
165 :param passwd: **DEPRECATED** Alias for password.
167 See `Connection <https://www.python.org/dev/peps/pep-0249/#connection-objects>`_ in the
168 specification.
169 """
171 _sock = None
172 _rfile = None
173 _auth_plugin_name = ""
174 _closed = False
175 _secure = False
177 def __init__(
178 self,
179 *,
180 user=None, # The first four arguments is based on DB-API 2.0 recommendation.
181 password="",
182 host=None,
183 database=None,
184 unix_socket=None,
185 port=0,
186 charset="",
187 collation=None,
188 sql_mode=None,
189 read_default_file=None,
190 conv=None,
191 use_unicode=True,
192 client_flag=0,
193 cursorclass=Cursor,
194 init_command=None,
195 connect_timeout=10,
196 read_default_group=None,
197 autocommit=False,
198 local_infile=False,
199 max_allowed_packet=16 * 1024 * 1024,
200 defer_connect=False,
201 auth_plugin_map=None,
202 read_timeout=None,
203 write_timeout=None,
204 bind_address=None,
205 binary_prefix=False,
206 program_name=None,
207 server_public_key=None,
208 ssl=None,
209 ssl_ca=None,
210 ssl_cert=None,
211 ssl_disabled=None,
212 ssl_key=None,
213 ssl_key_password=None,
214 ssl_verify_cert=None,
215 ssl_verify_identity=None,
216 compress=None, # not supported
217 named_pipe=None, # not supported
218 passwd=None, # deprecated
219 db=None, # deprecated
220 ):
221 if db is not None and database is None:
222 warnings.warn("'db' is deprecated, use 'database'", DeprecationWarning, 3)
223 database = db
224 if passwd is not None and not password:
225 warnings.warn(
226 "'passwd' is deprecated, use 'password'", DeprecationWarning, 3
227 )
228 password = passwd
230 if compress or named_pipe:
231 raise NotImplementedError(
232 "compress and named_pipe arguments are not supported"
233 )
235 self._local_infile = bool(local_infile)
236 if self._local_infile:
237 client_flag |= CLIENT.LOCAL_FILES
239 if read_default_group and not read_default_file:
240 if sys.platform.startswith("win"):
241 read_default_file = "c:\\my.ini"
242 else:
243 read_default_file = "/etc/my.cnf"
245 if read_default_file:
246 if not read_default_group:
247 read_default_group = "client"
249 cfg = Parser()
250 cfg.read(os.path.expanduser(read_default_file))
252 def _config(key, arg):
253 if arg:
254 return arg
255 try:
256 return cfg.get(read_default_group, key)
257 except Exception:
258 return arg
260 user = _config("user", user)
261 password = _config("password", password)
262 host = _config("host", host)
263 database = _config("database", database)
264 unix_socket = _config("socket", unix_socket)
265 port = int(_config("port", port))
266 bind_address = _config("bind-address", bind_address)
267 charset = _config("default-character-set", charset)
268 if not ssl:
269 ssl = {}
270 if isinstance(ssl, dict):
271 for key in ["ca", "capath", "cert", "key", "password", "cipher"]:
272 value = _config("ssl-" + key, ssl.get(key))
273 if value:
274 ssl[key] = value
276 self.ssl = False
277 self._ssl_required = False
278 if not ssl_disabled:
279 if ssl_ca or ssl_cert or ssl_key or ssl_verify_cert or ssl_verify_identity:
280 ssl = {
281 "ca": ssl_ca,
282 "check_hostname": bool(ssl_verify_identity),
283 "verify_mode": ssl_verify_cert
284 if ssl_verify_cert is not None
285 else False,
286 }
287 if ssl_cert is not None:
288 ssl["cert"] = ssl_cert
289 if ssl_key is not None:
290 ssl["key"] = ssl_key
291 if ssl_key_password is not None:
292 ssl["password"] = ssl_key_password
293 if ssl:
294 if not SSL_ENABLED:
295 raise NotImplementedError("ssl module not found")
296 self.ssl = True
297 self._ssl_required = True
298 client_flag |= CLIENT.SSL
299 self.ctx = self._create_ssl_ctx(ssl)
300 elif SSL_ENABLED:
301 # No explicit SSL options specified: use PREFERRED mode.
302 # Attempt SSL but fall back gracefully if the server doesn't support it.
303 self.ssl = True
304 self._ssl_required = False
305 self.ctx = self._create_ssl_ctx({})
307 self.host = host or "localhost"
308 self.port = port or 3306
309 if type(self.port) is not int:
310 raise ValueError("port should be of type int")
311 self.user = user or DEFAULT_USER
312 self.password = password or b""
313 if isinstance(self.password, str):
314 self.password = self.password.encode("latin1")
315 self.db = database
316 self.unix_socket = unix_socket
317 self.bind_address = bind_address
318 if not (0 < connect_timeout <= 31536000):
319 raise ValueError("connect_timeout should be >0 and <=31536000")
320 self.connect_timeout = connect_timeout or None
321 if read_timeout is not None and read_timeout <= 0:
322 raise ValueError("read_timeout should be > 0")
323 self._read_timeout = read_timeout
324 if write_timeout is not None and write_timeout <= 0:
325 raise ValueError("write_timeout should be > 0")
326 self._write_timeout = write_timeout
328 self.charset = charset or DEFAULT_CHARSET
329 self.collation = collation
330 self.use_unicode = use_unicode
332 self.encoding = charset_by_name(self.charset).encoding
334 client_flag |= CLIENT.CAPABILITIES
335 if self.db:
336 client_flag |= CLIENT.CONNECT_WITH_DB
338 self.client_flag = client_flag
340 self.cursorclass = cursorclass
342 self._result = None
343 self._affected_rows = 0
344 self.host_info = "Not connected"
346 # specified autocommit mode. None means use server default.
347 self.autocommit_mode = autocommit
349 if conv is None:
350 conv = converters.conversions
352 # Need for MySQLdb compatibility.
353 self.encoders = {k: v for (k, v) in conv.items() if type(k) is not int}
354 self.decoders = {k: v for (k, v) in conv.items() if type(k) is int}
355 self.sql_mode = sql_mode
356 self.init_command = init_command
357 self.max_allowed_packet = max_allowed_packet
358 self._auth_plugin_map = auth_plugin_map or {}
359 self.server_public_key = server_public_key
361 self._connect_attrs = {
362 "_client_name": "pymysql",
363 "_client_version": VERSION_STRING,
364 "_pid": str(os.getpid()),
365 }
367 if program_name:
368 self._connect_attrs["program_name"] = program_name
370 if defer_connect:
371 self._sock = None
372 else:
373 self.connect()
375 def __enter__(self):
376 return self
378 def __exit__(self, *exc_info):
379 del exc_info
380 self.close()
382 def _create_ssl_ctx(self, sslp):
383 if isinstance(sslp, ssl.SSLContext):
384 return sslp
385 ca = sslp.get("ca")
386 capath = sslp.get("capath")
387 hasnoca = ca is None and capath is None
388 ctx = ssl.create_default_context(cafile=ca, capath=capath)
390 # Python 3.13 enables VERIFY_X509_STRICT by default.
391 # But self signed certificates that are generated by MySQL automatically
392 # doesn't pass the verification.
393 ctx.verify_flags &= ~ssl.VERIFY_X509_STRICT
395 ctx.check_hostname = not hasnoca and sslp.get("check_hostname", True)
396 verify_mode_value = sslp.get("verify_mode")
397 if verify_mode_value is None:
398 ctx.verify_mode = ssl.CERT_NONE if hasnoca else ssl.CERT_REQUIRED
399 elif isinstance(verify_mode_value, bool):
400 ctx.verify_mode = ssl.CERT_REQUIRED if verify_mode_value else ssl.CERT_NONE
401 else:
402 if isinstance(verify_mode_value, str):
403 verify_mode_value = verify_mode_value.lower()
404 if verify_mode_value in ("none", "0", "false", "no"):
405 ctx.verify_mode = ssl.CERT_NONE
406 elif verify_mode_value == "optional":
407 ctx.verify_mode = ssl.CERT_OPTIONAL
408 elif verify_mode_value in ("required", "1", "true", "yes"):
409 ctx.verify_mode = ssl.CERT_REQUIRED
410 else:
411 ctx.verify_mode = ssl.CERT_NONE if hasnoca else ssl.CERT_REQUIRED
412 if "cert" in sslp:
413 ctx.load_cert_chain(
414 sslp["cert"], keyfile=sslp.get("key"), password=sslp.get("password")
415 )
416 if "cipher" in sslp:
417 ctx.set_ciphers(sslp["cipher"])
418 ctx.options |= ssl.OP_NO_SSLv2
419 ctx.options |= ssl.OP_NO_SSLv3
420 return ctx
422 def close(self):
423 """
424 Send the quit message and close the socket.
426 See `Connection.close() <https://www.python.org/dev/peps/pep-0249/#Connection.close>`_
427 in the specification.
429 :raise Error: If the connection is already closed.
430 """
431 if self._closed:
432 raise err.Error("Already closed")
433 self._closed = True
434 if self._sock is None:
435 return
436 send_data = struct.pack("<iB", 1, COMMAND.COM_QUIT)
437 try:
438 with contextlib.suppress(Exception):
439 self._write_bytes(send_data)
440 finally:
441 self._force_close()
443 @property
444 def open(self):
445 """Return True if the connection is open."""
446 return self._sock is not None
448 def _force_close(self):
449 """Close connection without QUIT message."""
450 if self._rfile:
451 self._rfile.close()
452 if self._sock:
453 try:
454 self._sock.close()
455 except: # noqa
456 pass
457 self._sock = None
458 self._rfile = None
460 __del__ = _force_close
462 def autocommit(self, value):
463 self.autocommit_mode = bool(value)
464 current = self.get_autocommit()
465 if value != current:
466 self._send_autocommit_mode()
468 def get_autocommit(self):
469 return bool(self.server_status & SERVER_STATUS.SERVER_STATUS_AUTOCOMMIT)
471 def _read_ok_packet(self):
472 pkt = self._read_packet()
473 if not pkt.is_ok_packet():
474 raise err.OperationalError(
475 CR.CR_COMMANDS_OUT_OF_SYNC,
476 "Command Out of Sync",
477 )
478 ok = OKPacketWrapper(pkt)
479 self.server_status = ok.server_status
480 return ok
482 def _send_autocommit_mode(self):
483 """Set whether or not to commit after every execute()."""
484 self._execute_command(
485 COMMAND.COM_QUERY, "SET AUTOCOMMIT = %s" % self.escape(self.autocommit_mode)
486 )
487 self._read_ok_packet()
489 def begin(self):
490 """Begin transaction."""
491 self._execute_command(COMMAND.COM_QUERY, "BEGIN")
492 self._read_ok_packet()
494 def commit(self):
495 """
496 Commit changes to stable storage.
498 See `Connection.commit() <https://www.python.org/dev/peps/pep-0249/#commit>`_
499 in the specification.
500 """
501 self._execute_command(COMMAND.COM_QUERY, "COMMIT")
502 self._read_ok_packet()
504 def rollback(self):
505 """
506 Roll back the current transaction.
508 See `Connection.rollback() <https://www.python.org/dev/peps/pep-0249/#rollback>`_
509 in the specification.
510 """
511 self._execute_command(COMMAND.COM_QUERY, "ROLLBACK")
512 self._read_ok_packet()
514 def show_warnings(self):
515 """Send the "SHOW WARNINGS" SQL command."""
516 self._execute_command(COMMAND.COM_QUERY, "SHOW WARNINGS")
517 result = MySQLResult(self)
518 result.read()
519 return result.rows
521 def select_db(self, db):
522 """
523 Set current db.
525 :param db: The name of the db.
526 """
527 self._execute_command(COMMAND.COM_INIT_DB, db)
528 self._read_ok_packet()
530 def escape(self, obj, mapping=None) -> str:
531 """Escape whatever value is passed.
533 Non-standard, for internal use; do not use this in your applications.
534 """
535 if isinstance(obj, str):
536 return f"'{self._escape_string(obj)}'"
538 if isinstance(obj, (bytes, bytearray)):
539 return f"X'{obj.hex()}'"
541 if mapping is None:
542 mapping = self.encoders
543 return converters.escape_item(obj, self.encoding, mapping=mapping)
545 def literal(self, obj) -> str:
546 """Alias for escape().
548 Non-standard, for internal use; do not use this in your applications.
549 """
550 warnings.warn(
551 "literal() is deprecated and will be removed in the next version.",
552 DeprecationWarning,
553 stacklevel=2,
554 )
555 return self.escape(obj)
557 def _escape_string(self, s: str):
558 if self.server_status & SERVER_STATUS.SERVER_STATUS_NO_BACKSLASH_ESCAPES:
559 return s.replace("'", "''") # Escape only single quote. Use '' quote
560 return converters.escape_string(s)
562 def cursor(self, cursor=None):
563 """
564 Create a new cursor to execute queries with.
566 :param cursor: The type of cursor to create. None means use Cursor.
567 :type cursor: :py:class:`Cursor`, :py:class:`SSCursor`, :py:class:`DictCursor`,
568 or :py:class:`SSDictCursor`.
569 """
570 if cursor:
571 return cursor(self)
572 return self.cursorclass(self)
574 # The following methods are INTERNAL USE ONLY (called from Cursor)
575 def query(self, sql, unbuffered=False):
576 # if DEBUG:
577 # print("DEBUG: sending query:", sql)
578 if isinstance(sql, str):
579 sql = sql.encode(self.encoding)
580 self._execute_command(COMMAND.COM_QUERY, sql)
581 self._affected_rows = self._read_query_result(unbuffered=unbuffered)
582 return self._affected_rows
584 def next_result(self, unbuffered=False):
585 self._affected_rows = self._read_query_result(unbuffered=unbuffered)
586 return self._affected_rows
588 def affected_rows(self):
589 return self._affected_rows
591 def kill(self, thread_id):
592 if not isinstance(thread_id, int):
593 raise TypeError("thread_id must be an integer")
594 self.query(f"KILL {thread_id:d}")
596 def ping(self, reconnect=False):
597 """
598 Check if the server is alive.
600 `reconnect` is deprecated. Create a new connection if you want to reconnect.
602 :param reconnect: If the connection is closed, reconnect.
603 :type reconnect: boolean
605 :raise Error: If the connection is closed and reconnect=False.
606 """
607 # emit deprecation warning for reconnect.
608 if reconnect:
609 warnings.warn(
610 "The 'reconnect' argument is deprecated. Create a new connection if you want to reconnect.",
611 DeprecationWarning,
612 2,
613 )
614 if self._sock is None:
615 if reconnect:
616 self.connect()
617 reconnect = False
618 else:
619 raise err.Error("Already closed")
620 try:
621 self._execute_command(COMMAND.COM_PING, "")
622 self._read_ok_packet()
623 except Exception:
624 if reconnect:
625 self.connect()
626 self.ping(False)
627 else:
628 raise
630 def set_charset(self, charset):
631 """Deprecated. Use set_character_set() instead."""
632 warnings.warn(
633 "'set_charset' is deprecated, use 'set_character_set' instead",
634 DeprecationWarning,
635 2,
636 )
637 # This function has been implemented in old PyMySQL.
638 # But this name is different from MySQLdb.
639 # So we keep this function for compatibility and add
640 # new set_character_set() function.
641 self.set_character_set(charset)
643 def set_character_set(self, charset, collation=None):
644 """
645 Set charset (and collation)
647 Send "SET NAMES charset [COLLATE collation]" query.
648 Update Connection.encoding based on charset.
649 """
650 # Make sure charset is supported.
651 encoding = charset_by_name(charset).encoding
653 if collation:
654 query = f"SET NAMES {charset} COLLATE {collation}"
655 else:
656 query = f"SET NAMES {charset}"
657 self._execute_command(COMMAND.COM_QUERY, query)
658 self._read_packet()
659 self.charset = charset
660 self.encoding = encoding
661 self.collation = collation
663 def connect(self, sock=None):
664 self._closed = False
665 try:
666 if sock is None:
667 if self.unix_socket:
668 sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
669 sock.settimeout(self.connect_timeout)
670 sock.connect(self.unix_socket)
671 self.host_info = "Localhost via UNIX socket"
672 self._secure = True
673 if DEBUG:
674 print("connected using unix_socket")
675 else:
676 kwargs = {}
677 if self.bind_address is not None:
678 kwargs["source_address"] = (self.bind_address, 0)
679 while True:
680 try:
681 sock = socket.create_connection(
682 (self.host, self.port), self.connect_timeout, **kwargs
683 )
684 break
685 except OSError as e:
686 if e.errno == errno.EINTR:
687 continue
688 raise
689 self.host_info = "socket %s:%d" % (self.host, self.port)
690 if DEBUG:
691 print("connected using socket")
692 sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
693 sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
695 self._sock = sock
696 sock.settimeout(self._read_timeout)
697 self._current_timeout = self._read_timeout
698 self._rfile = sock.makefile("rb")
699 self._next_seq_id = 0
701 self._get_server_information()
702 self._request_authentication()
704 # Send "SET NAMES" query on init for:
705 # - Ensure charset (and collation) is set to the server.
706 # - collation_id in handshake packet may be ignored.
707 # - If collation is not specified, we don't know what is server's
708 # default collation for the charset. For example, default collation
709 # of utf8mb4 is:
710 # - MySQL 5.7, MariaDB 10.x: utf8mb4_general_ci
711 # - MySQL 8.0: utf8mb4_0900_ai_ci
712 #
713 # Reference:
714 # - https://github.com/PyMySQL/PyMySQL/issues/1092
715 # - https://github.com/wagtail/wagtail/issues/9477
716 # - https://zenn.dev/methane/articles/2023-mysql-collation (Japanese)
717 self.set_character_set(self.charset, self.collation)
719 if self.sql_mode is not None:
720 c = self.cursor()
721 c.execute("SET sql_mode=%s", (self.sql_mode,))
722 c.close()
724 if self.init_command is not None:
725 c = self.cursor()
726 c.execute(self.init_command)
727 c.close()
729 if self.autocommit_mode is not None:
730 self.autocommit(self.autocommit_mode)
731 except BaseException as e:
732 self._force_close()
734 if isinstance(e, OSError):
735 exc = err.OperationalError(
736 CR.CR_CONN_HOST_ERROR,
737 f"Can't connect to MySQL server on {self.host!r} ({e})",
738 )
739 # Keep original exception and traceback to investigate error.
740 exc.original_exception = e
741 exc.traceback = traceback.format_exc()
742 if DEBUG:
743 print(exc.traceback)
744 raise exc
746 # If e is neither DatabaseError or IOError, It's a bug.
747 # But raising AssertionError hides original error.
748 # So just reraise it.
749 raise
751 def write_packet(self, payload):
752 """Writes an entire "mysql packet" in its entirety to the network
753 adding its length and sequence number.
754 """
755 # Internal note: when you build packet manually and calls _write_bytes()
756 # directly, you should set self._next_seq_id properly.
757 data = _pack_int24(len(payload)) + bytes([self._next_seq_id]) + payload
758 if DEBUG:
759 dump_packet(data)
760 self._write_bytes(data)
761 self._next_seq_id = (self._next_seq_id + 1) % 256
763 def _read_packet(self, packet_type=MysqlPacket):
764 """Read an entire "mysql packet" in its entirety from the network
765 and return a MysqlPacket type that represents the results.
767 :raise OperationalError: If the connection to the MySQL server is lost.
768 :raise InternalError: If the packet sequence number is wrong.
769 """
770 # Although `socket.settimeout()` may appear fast, it temporarily releases
771 # the GIL, which can hurt performance in multithreaded applications.
772 # Avoid calling it repeatedly at high frequency.
773 if self._current_timeout != self._read_timeout:
774 self._sock.settimeout(self._read_timeout)
775 self._current_timeout = self._read_timeout
777 buff = []
778 while True:
779 packet_header = self._read_bytes(4)
780 # if DEBUG: dump_packet(packet_header)
782 btrl, btrh, packet_number = struct.unpack("<HBB", packet_header)
783 bytes_to_read = btrl + (btrh << 16)
784 if packet_number != self._next_seq_id:
785 self._force_close()
786 if packet_number == 0:
787 # MariaDB sends error packet with seqno==0 when shutdown
788 raise err.OperationalError(
789 CR.CR_SERVER_LOST,
790 "Lost connection to MySQL server during query",
791 )
792 raise err.InternalError(
793 "Packet sequence number wrong - got %d expected %d"
794 % (packet_number, self._next_seq_id)
795 )
796 self._next_seq_id = (self._next_seq_id + 1) % 256
798 recv_data = self._read_bytes(bytes_to_read)
799 if DEBUG:
800 dump_packet(recv_data)
801 buff.append(recv_data)
802 # https://dev.mysql.com/doc/internals/en/sending-more-than-16mbyte.html
803 if bytes_to_read < MAX_PACKET_LEN:
804 break
806 packet = packet_type(b"".join(buff), self.encoding)
807 if packet.is_error_packet():
808 if self._result is not None and self._result.unbuffered_active is True:
809 self._result.unbuffered_active = False
810 packet.raise_for_error()
811 return packet
813 def _read_bytes(self, num_bytes):
814 # NOTE: caller should call self._sock.settimeout(self._read_timeout)
815 # before first read.
816 while True:
817 try:
818 data = self._rfile.read(num_bytes)
819 break
820 except OSError as e:
821 if e.errno == errno.EINTR:
822 continue
823 self._force_close()
824 raise err.OperationalError(
825 CR.CR_SERVER_LOST,
826 f"Lost connection to MySQL server during query ({e})",
827 )
828 except BaseException:
829 # Don't convert unknown exception to MySQLError.
830 self._force_close()
831 raise
832 if len(data) < num_bytes:
833 self._force_close()
834 raise err.OperationalError(
835 CR.CR_SERVER_LOST, "Lost connection to MySQL server during query"
836 )
837 return data
839 def _write_bytes(self, data):
840 if self._current_timeout != self._write_timeout:
841 self._sock.settimeout(self._write_timeout)
842 self._current_timeout = self._write_timeout
843 try:
844 self._sock.sendall(data)
845 except OSError as e:
846 self._force_close()
847 raise err.OperationalError(
848 CR.CR_SERVER_GONE_ERROR, f"MySQL server has gone away ({e!r})"
849 )
851 def _read_query_result(self, unbuffered=False):
852 self._result = None
853 result = MySQLResult(self)
854 if unbuffered:
855 result.init_unbuffered_query()
856 else:
857 result.read()
858 self._result = result
859 if result.server_status is not None:
860 self.server_status = result.server_status
861 return result.affected_rows
863 def insert_id(self):
864 if self._result:
865 return self._result.insert_id
866 else:
867 return 0
869 def _execute_command(self, command, sql):
870 """
871 :raise InterfaceError: If the connection is closed.
872 :raise ValueError: If no username was specified.
873 """
874 if not self._sock:
875 raise err.InterfaceError(0, "")
877 # If the last query was unbuffered, make sure it finishes before
878 # sending new commands
879 if self._result is not None:
880 if self._result.unbuffered_active:
881 warnings.warn("Previous unbuffered result was left incomplete")
882 self._result._finish_unbuffered_query()
883 while self._result.has_next:
884 self.next_result()
885 self._result = None
887 if isinstance(sql, str):
888 sql = sql.encode(self.encoding)
890 packet_size = min(MAX_PACKET_LEN, len(sql) + 1) # +1 is for command
892 # tiny optimization: build first packet manually instead of
893 # calling self..write_packet()
894 prelude = struct.pack("<iB", packet_size, command)
895 packet = prelude + sql[: packet_size - 1]
896 self._write_bytes(packet)
897 if DEBUG:
898 dump_packet(packet)
899 self._next_seq_id = 1
901 if packet_size < MAX_PACKET_LEN:
902 return
904 sql = sql[packet_size - 1 :]
905 while True:
906 packet_size = min(MAX_PACKET_LEN, len(sql))
907 self.write_packet(sql[:packet_size])
908 sql = sql[packet_size:]
909 if not sql and packet_size < MAX_PACKET_LEN:
910 break
912 def _request_authentication(self):
913 # https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::HandshakeResponse
914 if int(self.server_version.split(".", 1)[0]) >= 5:
915 self.client_flag |= CLIENT.MULTI_RESULTS
917 if self.user is None:
918 raise ValueError("Did not specify a username")
920 charset_id = charset_by_name(self.charset).id
921 if isinstance(self.user, str):
922 self.user = self.user.encode(self.encoding)
924 # Determine flags for the initial handshake packet.
925 # CLIENT.SSL is added conditionally: for REQUIRED mode it is already set in
926 # self.client_flag, but for PREFERRED mode it is only added when the server
927 # also advertises SSL support.
928 # _do_ssl is set here and checked below for sha256_password auth.
929 client_flags = self.client_flag
930 if self.ssl:
931 if self.server_capabilities & CLIENT.SSL:
932 # SSL upgrade: include CLIENT.SSL flag and wrap the socket.
933 _do_ssl = True
934 client_flags |= CLIENT.SSL
935 elif self._ssl_required:
936 raise err.OperationalError(
937 CR.CR_SSL_CONNECTION_ERROR,
938 "SSL is required but the server doesn't support it",
939 )
940 else:
941 # PREFERRED mode: server doesn't support SSL, fall back to non-SSL.
942 _do_ssl = False
943 else:
944 _do_ssl = False
946 data_init = struct.pack(
947 "<iIB23s", client_flags, MAX_PACKET_LEN, charset_id, b""
948 )
950 if _do_ssl:
951 self.write_packet(data_init)
952 self._sock = self.ctx.wrap_socket(self._sock, server_hostname=self.host)
953 self._rfile = self._sock.makefile("rb")
954 self._secure = True
956 data = data_init + self.user + b"\0"
958 authresp = b""
959 plugin_name = None
961 if self._auth_plugin_name == "":
962 plugin_name = b""
963 authresp = _auth.scramble_native_password(self.password, self.salt)
964 elif self._auth_plugin_name == "mysql_native_password":
965 plugin_name = b"mysql_native_password"
966 authresp = _auth.scramble_native_password(self.password, self.salt)
967 elif self._auth_plugin_name == "caching_sha2_password":
968 plugin_name = b"caching_sha2_password"
969 if self.password:
970 if DEBUG:
971 print("caching_sha2: trying fast path")
972 authresp = _auth.scramble_caching_sha2(self.password, self.salt)
973 else:
974 if DEBUG:
975 print("caching_sha2: empty password")
976 elif self._auth_plugin_name == "sha256_password":
977 plugin_name = b"sha256_password"
978 if _do_ssl:
979 authresp = self.password + b"\0"
980 elif self.password:
981 authresp = b"\1" # request public key
982 else:
983 authresp = b"\0" # empty password
985 if self.server_capabilities & CLIENT.PLUGIN_AUTH_LENENC_CLIENT_DATA:
986 data += _lenenc_int(len(authresp)) + authresp
987 elif self.server_capabilities & CLIENT.SECURE_CONNECTION:
988 data += struct.pack("B", len(authresp)) + authresp
989 else: # pragma: no cover - not testing against servers without secure auth (>=5.0)
990 data += authresp + b"\0"
992 if self.db and self.server_capabilities & CLIENT.CONNECT_WITH_DB:
993 if isinstance(self.db, str):
994 self.db = self.db.encode(self.encoding)
995 data += self.db + b"\0"
997 if self.server_capabilities & CLIENT.PLUGIN_AUTH:
998 data += (plugin_name or b"") + b"\0"
1000 if self.server_capabilities & CLIENT.CONNECT_ATTRS:
1001 connect_attrs = b""
1002 for k, v in self._connect_attrs.items():
1003 k = k.encode("utf-8")
1004 connect_attrs += _lenenc_int(len(k)) + k
1005 v = v.encode("utf-8")
1006 connect_attrs += _lenenc_int(len(v)) + v
1007 data += _lenenc_int(len(connect_attrs)) + connect_attrs
1009 self.write_packet(data)
1010 auth_packet = self._read_packet()
1012 # Authentication is a state machine. An authentication plugin can return
1013 # another transition packet (for example, MySQL Router can request full
1014 # caching_sha2_password authentication and then switch to the backend's
1015 # mysql_native_password plugin), so keep dispatching until the server
1016 # sends a terminal packet.
1017 auth_plugin_name = self._auth_plugin_name
1018 if isinstance(auth_plugin_name, str):
1019 auth_plugin_name = auth_plugin_name.encode("ascii")
1020 auth_plugin_handler = self._get_auth_plugin_handler(auth_plugin_name)
1021 auth_switch_received = False
1023 while True:
1024 # Custom authentication handlers historically did not need to return
1025 # the final OK packet after consuming the complete exchange.
1026 if auth_packet is None and auth_plugin_handler:
1027 break
1029 # if authentication method isn't accepted the first byte
1030 # will have the octet 254
1031 if auth_packet.is_auth_switch_request():
1032 if auth_switch_received:
1033 raise err.OperationalError("received multiple auth switch requests")
1034 auth_switch_received = True
1035 if DEBUG:
1036 print("received auth switch")
1037 # https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::AuthSwitchRequest
1038 auth_packet.read_uint8() # 0xfe packet identifier
1039 plugin_name = auth_packet.read_string()
1040 if (
1041 self.server_capabilities & CLIENT.PLUGIN_AUTH
1042 and plugin_name is not None
1043 ):
1044 auth_plugin_name = plugin_name
1045 auth_plugin_handler = self._get_auth_plugin_handler(plugin_name)
1046 auth_packet = self._process_auth(
1047 plugin_name, auth_packet, auth_plugin_handler
1048 )
1049 continue
1050 raise err.OperationalError("received unknown auth switch request")
1052 if auth_packet.is_extra_auth_data():
1053 if DEBUG:
1054 print("received extra data")
1055 # https://dev.mysql.com/doc/internals/en/successful-authentication.html
1056 if auth_plugin_handler:
1057 auth_packet = auth_plugin_handler.authenticate(auth_packet)
1058 continue
1059 elif auth_plugin_name in (
1060 b"caching_sha2_password",
1061 "caching_sha2_password",
1062 ):
1063 auth_packet = _auth.caching_sha2_password_auth(self, auth_packet)
1064 continue
1065 if auth_plugin_name in (b"sha256_password", "sha256_password"):
1066 auth_packet = _auth.sha256_password_auth(self, auth_packet)
1067 continue
1068 raise err.OperationalError(
1069 "Received extra packet for auth method %r", auth_plugin_name
1070 )
1072 if auth_packet.is_ok_packet():
1073 break
1074 raise err.OperationalError("unexpected packet during authentication")
1076 if DEBUG:
1077 print("Succeed to auth")
1079 def _process_auth(self, plugin_name, auth_packet, handler=None):
1080 if handler:
1081 try:
1082 return handler.authenticate(auth_packet)
1083 except AttributeError:
1084 if plugin_name != b"dialog":
1085 raise err.OperationalError(
1086 CR.CR_AUTH_PLUGIN_CANNOT_LOAD,
1087 f"Authentication plugin '{plugin_name}'"
1088 f" not loaded: - {type(handler)!r} missing authenticate method",
1089 )
1090 if plugin_name == b"caching_sha2_password":
1091 return _auth.caching_sha2_password_auth(self, auth_packet)
1092 elif plugin_name == b"sha256_password":
1093 return _auth.sha256_password_auth(self, auth_packet)
1094 elif plugin_name == b"mysql_native_password":
1095 data = _auth.scramble_native_password(self.password, auth_packet.read_all())
1096 elif plugin_name == b"client_ed25519":
1097 data = _auth.ed25519_password(self.password, auth_packet.read_all())
1098 elif plugin_name == b"mysql_old_password":
1099 data = (
1100 _auth.scramble_old_password(self.password, auth_packet.read_all())
1101 + b"\0"
1102 )
1103 elif plugin_name == b"mysql_clear_password":
1104 # https://dev.mysql.com/doc/internals/en/clear-text-authentication.html
1105 data = self.password + b"\0"
1106 elif plugin_name == b"dialog":
1107 pkt = auth_packet
1108 while True:
1109 flag = pkt.read_uint8()
1110 echo = (flag & 0x06) == 0x02
1111 last = (flag & 0x01) == 0x01
1112 prompt = pkt.read_all()
1114 if prompt == b"Password: ":
1115 self.write_packet(self.password + b"\0")
1116 elif handler:
1117 resp = "no response - TypeError within plugin.prompt method"
1118 try:
1119 resp = handler.prompt(echo, prompt)
1120 self.write_packet(resp + b"\0")
1121 except AttributeError:
1122 raise err.OperationalError(
1123 CR.CR_AUTH_PLUGIN_CANNOT_LOAD,
1124 f"Authentication plugin '{plugin_name}'"
1125 f" not loaded: - {handler!r} missing prompt method",
1126 )
1127 except TypeError:
1128 raise err.OperationalError(
1129 CR.CR_AUTH_PLUGIN_ERR,
1130 f"Authentication plugin '{plugin_name}'"
1131 f" {handler!r} didn't respond with string. Returned '{resp!r}' to prompt {prompt!r}",
1132 )
1133 else:
1134 raise err.OperationalError(
1135 CR.CR_AUTH_PLUGIN_CANNOT_LOAD,
1136 f"Authentication plugin '{plugin_name}' not configured",
1137 )
1138 pkt = self._read_packet()
1139 pkt.check_error()
1140 if pkt.is_ok_packet() or last:
1141 break
1142 return pkt
1143 else:
1144 raise err.OperationalError(
1145 CR.CR_AUTH_PLUGIN_CANNOT_LOAD,
1146 "Authentication plugin '%s' not configured" % plugin_name,
1147 )
1149 self.write_packet(data)
1150 pkt = self._read_packet()
1151 pkt.check_error()
1152 return pkt
1154 def _get_auth_plugin_handler(self, plugin_name):
1155 plugin_class = self._auth_plugin_map.get(plugin_name)
1156 if not plugin_class and isinstance(plugin_name, bytes):
1157 plugin_class = self._auth_plugin_map.get(plugin_name.decode("ascii"))
1158 if plugin_class:
1159 try:
1160 handler = plugin_class(self)
1161 except TypeError:
1162 raise err.OperationalError(
1163 CR.CR_AUTH_PLUGIN_CANNOT_LOAD,
1164 f"Authentication plugin '{plugin_name}'"
1165 f" not loaded: - {plugin_class!r} cannot be constructed with connection object",
1166 )
1167 else:
1168 handler = None
1169 return handler
1171 # _mysql support
1172 def thread_id(self):
1173 return self.server_thread_id[0]
1175 def character_set_name(self):
1176 return self.charset
1178 def get_host_info(self):
1179 return self.host_info
1181 def get_proto_info(self):
1182 return self.protocol_version
1184 def _get_server_information(self):
1185 i = 0
1186 packet = self._read_packet()
1187 data = packet.get_all_data()
1189 self.protocol_version = data[i]
1190 i += 1
1192 server_end = data.find(b"\0", i)
1193 self.server_version = data[i:server_end].decode("latin1")
1194 i = server_end + 1
1196 self.server_thread_id = struct.unpack("<I", data[i : i + 4])
1197 i += 4
1199 self.salt = data[i : i + 8]
1200 i += 9 # 8 + 1(filler)
1202 self.server_capabilities = struct.unpack("<H", data[i : i + 2])[0]
1203 i += 2
1205 if len(data) >= i + 6:
1206 lang, stat, cap_h, salt_len = struct.unpack("<BHHB", data[i : i + 6])
1207 i += 6
1208 # TODO: deprecate server_language and server_charset.
1209 # mysqlclient-python doesn't provide it.
1210 self.server_language = lang
1211 try:
1212 self.server_charset = charset_by_id(lang).name
1213 except KeyError:
1214 # unknown collation
1215 self.server_charset = None
1217 self.server_status = stat
1218 if DEBUG:
1219 print("server_status: %x" % stat)
1221 self.server_capabilities |= cap_h << 16
1222 if DEBUG:
1223 print("salt_len:", salt_len)
1224 salt_len = max(12, salt_len - 9)
1226 # reserved
1227 i += 10
1229 if len(data) >= i + salt_len:
1230 # salt_len includes auth_plugin_data_part_1 and filler
1231 self.salt += data[i : i + salt_len]
1232 i += salt_len
1234 i += 1
1235 # AUTH PLUGIN NAME may appear here.
1236 if self.server_capabilities & CLIENT.PLUGIN_AUTH and len(data) >= i:
1237 # Due to Bug#59453 the auth-plugin-name is missing the terminating
1238 # NUL-char in versions prior to 5.5.10 and 5.6.2.
1239 # ref: https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::Handshake
1240 # didn't use version checks as mariadb is corrected and reports
1241 # earlier than those two.
1242 server_end = data.find(b"\0", i)
1243 if server_end < 0: # pragma: no cover - very specific upstream bug
1244 # not found \0 and last field so take it all
1245 self._auth_plugin_name = data[i:].decode("utf-8")
1246 else:
1247 self._auth_plugin_name = data[i:server_end].decode("utf-8")
1249 if _DEFAULT_AUTH_PLUGIN is not None: # for tests
1250 self._auth_plugin_name = _DEFAULT_AUTH_PLUGIN
1252 def get_server_info(self):
1253 return self.server_version
1255 Warning = err.Warning
1256 Error = err.Error
1257 InterfaceError = err.InterfaceError
1258 DatabaseError = err.DatabaseError
1259 DataError = err.DataError
1260 OperationalError = err.OperationalError
1261 IntegrityError = err.IntegrityError
1262 InternalError = err.InternalError
1263 ProgrammingError = err.ProgrammingError
1264 NotSupportedError = err.NotSupportedError
1267class MySQLResult:
1268 def __init__(self, connection):
1269 """
1270 :type connection: Connection
1271 """
1272 self.connection = connection
1273 self.affected_rows = None
1274 self.insert_id = None
1275 self.server_status = None
1276 self.warning_count = 0
1277 self.message = None
1278 self.field_count = 0
1279 self.description = None
1280 self.rows = None
1281 self.has_next = None
1282 self.unbuffered_active = False
1284 def __del__(self):
1285 if self.unbuffered_active:
1286 self._finish_unbuffered_query()
1288 def read(self):
1289 try:
1290 first_packet = self.connection._read_packet()
1292 if first_packet.is_ok_packet():
1293 self._read_ok_packet(first_packet)
1294 elif first_packet.is_load_local_packet():
1295 self._read_load_local_packet(first_packet)
1296 else:
1297 self._read_result_packet(first_packet)
1298 finally:
1299 self.connection = None
1301 def init_unbuffered_query(self):
1302 """
1303 :raise OperationalError: If the connection to the MySQL server is lost.
1304 :raise InternalError:
1305 """
1306 first_packet = self.connection._read_packet()
1308 if first_packet.is_ok_packet():
1309 self.connection = None
1310 self._read_ok_packet(first_packet)
1311 elif first_packet.is_load_local_packet():
1312 try:
1313 self._read_load_local_packet(first_packet)
1314 finally:
1315 self.connection = None
1316 else:
1317 self.field_count = first_packet.read_length_encoded_integer()
1318 self._get_descriptions()
1320 # Apparently, MySQLdb picks this number because it's the maximum
1321 # value of a 64bit unsigned integer. Since we're emulating MySQLdb,
1322 # we set it to this instead of None, which would be preferred.
1323 self.affected_rows = 18446744073709551615
1324 self.unbuffered_active = True
1326 def _read_ok_packet(self, packet):
1327 ok_packet = OKPacketWrapper(packet)
1328 self.affected_rows = ok_packet.affected_rows
1329 self.insert_id = ok_packet.insert_id
1330 self.server_status = ok_packet.server_status
1331 self.warning_count = ok_packet.warning_count
1332 self.message = ok_packet.message
1333 self.has_next = ok_packet.has_next
1335 def _read_load_local_packet(self, first_packet):
1336 conn: Connection = self.connection
1337 if not conn._local_infile:
1338 raise RuntimeError(
1339 "**WARN**: Received LOAD_LOCAL packet but local_infile option is false."
1340 )
1341 load_packet = LoadLocalPacketWrapper(first_packet)
1342 try:
1343 _send_local_file(load_packet.filename, conn)
1344 finally:
1345 # send the empty packet to signify we are done sending data
1346 conn.write_packet(b"")
1347 ok_packet = conn._read_packet()
1348 # If an error occurs while sending the file, exit here without handling
1349 # the OK packet.
1351 if not ok_packet.is_ok_packet():
1352 raise err.OperationalError(
1353 CR.CR_COMMANDS_OUT_OF_SYNC, "Commands Out of Sync"
1354 )
1355 self._read_ok_packet(ok_packet)
1357 def _check_packet_is_eof(self, packet):
1358 if not packet.is_eof_packet():
1359 return False
1360 # TODO: Support CLIENT.DEPRECATE_EOF
1361 # 1) Add DEPRECATE_EOF to CAPABILITIES
1362 # 2) Mask CAPABILITIES with server_capabilities
1363 # 3) if server_capabilities & CLIENT.DEPRECATE_EOF:
1364 # use OKPacketWrapper instead of EOFPacketWrapper
1365 wp = EOFPacketWrapper(packet)
1366 self.warning_count = wp.warning_count
1367 self.has_next = wp.has_next
1368 return True
1370 def _read_result_packet(self, first_packet):
1371 self.field_count = first_packet.read_length_encoded_integer()
1372 self._get_descriptions()
1373 self._read_rowdata_packet()
1375 def _read_rowdata_packet_unbuffered(self):
1376 # Check if in an active query
1377 if not self.unbuffered_active:
1378 return
1380 # EOF
1381 packet = self.connection._read_packet()
1382 if self._check_packet_is_eof(packet):
1383 self.unbuffered_active = False
1384 self.connection = None
1385 self.rows = None
1386 return
1388 row = self._read_row_from_packet(packet)
1389 self.affected_rows = 1
1390 self.rows = (row,) # rows should tuple of row for MySQL-python compatibility.
1391 return row
1393 def _finish_unbuffered_query(self):
1394 # After much reading on the MySQL protocol, it appears that there is,
1395 # in fact, no way to stop MySQL from sending all the data after
1396 # executing a query, so we just spin, and wait for an EOF packet.
1397 while self.unbuffered_active:
1398 try:
1399 packet = self.connection._read_packet()
1400 except err.OperationalError as e:
1401 if e.args[0] in (
1402 ER.QUERY_TIMEOUT,
1403 ER.STATEMENT_TIMEOUT,
1404 ):
1405 # if the query timed out we can simply ignore this error
1406 self.unbuffered_active = False
1407 self.connection = None
1408 return
1410 raise
1412 if self._check_packet_is_eof(packet):
1413 self.unbuffered_active = False
1414 self.connection = None # release reference to kill cyclic reference.
1416 def _read_rowdata_packet(self):
1417 """Read a rowdata packet for each data row in the result set."""
1418 rows = []
1419 while True:
1420 packet = self.connection._read_packet()
1421 if self._check_packet_is_eof(packet):
1422 self.connection = None # release reference to kill cyclic reference.
1423 break
1424 rows.append(self._read_row_from_packet(packet))
1426 self.affected_rows = len(rows)
1427 self.rows = tuple(rows)
1429 def _read_row_from_packet(self, packet):
1430 row = []
1431 for encoding, converter in self.converters:
1432 try:
1433 data = packet.read_length_coded_string()
1434 except IndexError:
1435 # No more columns in this row
1436 # See https://github.com/PyMySQL/PyMySQL/pull/434
1437 break
1438 if data is not None:
1439 if encoding is not None:
1440 data = data.decode(encoding)
1441 if DEBUG:
1442 print("DEBUG: DATA = ", data)
1443 if converter is not None:
1444 data = converter(data)
1445 row.append(data)
1446 return tuple(row)
1448 def _get_descriptions(self):
1449 """Read a column descriptor packet for each column in the result."""
1450 self.fields = []
1451 self.converters = []
1452 use_unicode = self.connection.use_unicode
1453 conn_encoding = self.connection.encoding
1454 description = []
1456 for i in range(self.field_count):
1457 field = self.connection._read_packet(FieldDescriptorPacket)
1458 self.fields.append(field)
1459 description.append(field.description())
1460 field_type = field.type_code
1461 if use_unicode:
1462 if field_type == FIELD_TYPE.JSON:
1463 # When SELECT from JSON column: charset = binary
1464 # When SELECT CAST(... AS JSON): charset = connection encoding
1465 # This behavior is different from TEXT / BLOB.
1466 # We should decode result by connection encoding regardless charsetnr.
1467 # See https://github.com/PyMySQL/PyMySQL/issues/488
1468 encoding = conn_encoding # SELECT CAST(... AS JSON)
1469 elif field_type in TEXT_TYPES:
1470 if field.charsetnr == 63: # binary
1471 # TEXTs with charset=binary means BINARY types.
1472 encoding = None
1473 else:
1474 encoding = conn_encoding
1475 else:
1476 # Integers, Dates and Times, and other basic data is encoded in ascii
1477 encoding = "ascii"
1478 else:
1479 encoding = None
1480 converter = self.connection.decoders.get(field_type)
1481 if converter is converters.through:
1482 converter = None
1483 if DEBUG:
1484 print(f"DEBUG: field={field}, converter={converter}")
1485 self.converters.append((encoding, converter))
1487 eof_packet = self.connection._read_packet()
1488 assert eof_packet.is_eof_packet(), "Protocol error, expecting EOF"
1489 self.description = tuple(description)
1492def _send_local_file(filename: str, conn: Connection):
1493 """Send data packets from the local file to the server"""
1494 packet_size = min(conn.max_allowed_packet, 16 * 1024)
1496 try:
1497 with open(filename, "rb") as file:
1498 # 16KB is efficient enough
1499 while True:
1500 chunk = file.read(packet_size)
1501 if not chunk:
1502 break
1503 conn.write_packet(chunk)
1504 except OSError as e:
1505 raise err.OperationalError(
1506 ER.FILE_NOT_FOUND,
1507 f"Can't open file '{filename}': {e}",
1508 )