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

867 statements  

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 

6 

7import contextlib 

8import errno 

9import os 

10import socket 

11import struct 

12import sys 

13import traceback 

14import warnings 

15 

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) 

29 

30try: 

31 import ssl 

32 

33 SSL_ENABLED = True 

34except ImportError: 

35 ssl = None 

36 SSL_ENABLED = False 

37 

38try: 

39 import getpass 

40 

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 

48 

49DEBUG = False 

50_DEFAULT_AUTH_PLUGIN = None # if this is not None, use it instead of server's default. 

51 

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} 

63 

64 

65DEFAULT_CHARSET = "utf8mb4" 

66 

67MAX_PACKET_LEN = 2**24 - 1 

68 

69 

70def _pack_int24(n): 

71 return struct.pack("<I", n)[:3] 

72 

73 

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 ) 

92 

93 

94class Connection: 

95 """ 

96 Representation of a socket with a mysql server. 

97 

98 The proper way to get an instance of this class is to call 

99 connect(). 

100 

101 Establish a connection to the MySQL database. Accepts several 

102 arguments: 

103 

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. 

166 

167 See `Connection <https://www.python.org/dev/peps/pep-0249/#connection-objects>`_ in the 

168 specification. 

169 """ 

170 

171 _sock = None 

172 _rfile = None 

173 _auth_plugin_name = "" 

174 _closed = False 

175 _secure = False 

176 

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 

229 

230 if compress or named_pipe: 

231 raise NotImplementedError( 

232 "compress and named_pipe arguments are not supported" 

233 ) 

234 

235 self._local_infile = bool(local_infile) 

236 if self._local_infile: 

237 client_flag |= CLIENT.LOCAL_FILES 

238 

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" 

244 

245 if read_default_file: 

246 if not read_default_group: 

247 read_default_group = "client" 

248 

249 cfg = Parser() 

250 cfg.read(os.path.expanduser(read_default_file)) 

251 

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 

259 

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 

275 

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({}) 

306 

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 

327 

328 self.charset = charset or DEFAULT_CHARSET 

329 self.collation = collation 

330 self.use_unicode = use_unicode 

331 

332 self.encoding = charset_by_name(self.charset).encoding 

333 

334 client_flag |= CLIENT.CAPABILITIES 

335 if self.db: 

336 client_flag |= CLIENT.CONNECT_WITH_DB 

337 

338 self.client_flag = client_flag 

339 

340 self.cursorclass = cursorclass 

341 

342 self._result = None 

343 self._affected_rows = 0 

344 self.host_info = "Not connected" 

345 

346 # specified autocommit mode. None means use server default. 

347 self.autocommit_mode = autocommit 

348 

349 if conv is None: 

350 conv = converters.conversions 

351 

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 

360 

361 self._connect_attrs = { 

362 "_client_name": "pymysql", 

363 "_client_version": VERSION_STRING, 

364 "_pid": str(os.getpid()), 

365 } 

366 

367 if program_name: 

368 self._connect_attrs["program_name"] = program_name 

369 

370 if defer_connect: 

371 self._sock = None 

372 else: 

373 self.connect() 

374 

375 def __enter__(self): 

376 return self 

377 

378 def __exit__(self, *exc_info): 

379 del exc_info 

380 self.close() 

381 

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) 

389 

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 

394 

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 

421 

422 def close(self): 

423 """ 

424 Send the quit message and close the socket. 

425 

426 See `Connection.close() <https://www.python.org/dev/peps/pep-0249/#Connection.close>`_ 

427 in the specification. 

428 

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() 

442 

443 @property 

444 def open(self): 

445 """Return True if the connection is open.""" 

446 return self._sock is not None 

447 

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 

459 

460 __del__ = _force_close 

461 

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() 

467 

468 def get_autocommit(self): 

469 return bool(self.server_status & SERVER_STATUS.SERVER_STATUS_AUTOCOMMIT) 

470 

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 

481 

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() 

488 

489 def begin(self): 

490 """Begin transaction.""" 

491 self._execute_command(COMMAND.COM_QUERY, "BEGIN") 

492 self._read_ok_packet() 

493 

494 def commit(self): 

495 """ 

496 Commit changes to stable storage. 

497 

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() 

503 

504 def rollback(self): 

505 """ 

506 Roll back the current transaction. 

507 

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() 

513 

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 

520 

521 def select_db(self, db): 

522 """ 

523 Set current db. 

524 

525 :param db: The name of the db. 

526 """ 

527 self._execute_command(COMMAND.COM_INIT_DB, db) 

528 self._read_ok_packet() 

529 

530 def escape(self, obj, mapping=None) -> str: 

531 """Escape whatever value is passed. 

532 

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)}'" 

537 

538 if isinstance(obj, (bytes, bytearray)): 

539 return f"X'{obj.hex()}'" 

540 

541 if mapping is None: 

542 mapping = self.encoders 

543 return converters.escape_item(obj, self.encoding, mapping=mapping) 

544 

545 def literal(self, obj) -> str: 

546 """Alias for escape(). 

547 

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) 

556 

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) 

561 

562 def cursor(self, cursor=None): 

563 """ 

564 Create a new cursor to execute queries with. 

565 

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) 

573 

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 

583 

584 def next_result(self, unbuffered=False): 

585 self._affected_rows = self._read_query_result(unbuffered=unbuffered) 

586 return self._affected_rows 

587 

588 def affected_rows(self): 

589 return self._affected_rows 

590 

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}") 

595 

596 def ping(self, reconnect=False): 

597 """ 

598 Check if the server is alive. 

599 

600 `reconnect` is deprecated. Create a new connection if you want to reconnect. 

601 

602 :param reconnect: If the connection is closed, reconnect. 

603 :type reconnect: boolean 

604 

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 

629 

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) 

642 

643 def set_character_set(self, charset, collation=None): 

644 """ 

645 Set charset (and collation) 

646 

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 

652 

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 

662 

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) 

694 

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 

700 

701 self._get_server_information() 

702 self._request_authentication() 

703 

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) 

718 

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() 

723 

724 if self.init_command is not None: 

725 c = self.cursor() 

726 c.execute(self.init_command) 

727 c.close() 

728 

729 if self.autocommit_mode is not None: 

730 self.autocommit(self.autocommit_mode) 

731 except BaseException as e: 

732 self._force_close() 

733 

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 

745 

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 

750 

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 

762 

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. 

766 

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 

776 

777 buff = [] 

778 while True: 

779 packet_header = self._read_bytes(4) 

780 # if DEBUG: dump_packet(packet_header) 

781 

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 

797 

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 

805 

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 

812 

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 

838 

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 ) 

850 

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 

862 

863 def insert_id(self): 

864 if self._result: 

865 return self._result.insert_id 

866 else: 

867 return 0 

868 

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, "") 

876 

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 

886 

887 if isinstance(sql, str): 

888 sql = sql.encode(self.encoding) 

889 

890 packet_size = min(MAX_PACKET_LEN, len(sql) + 1) # +1 is for command 

891 

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 

900 

901 if packet_size < MAX_PACKET_LEN: 

902 return 

903 

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 

911 

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 

916 

917 if self.user is None: 

918 raise ValueError("Did not specify a username") 

919 

920 charset_id = charset_by_name(self.charset).id 

921 if isinstance(self.user, str): 

922 self.user = self.user.encode(self.encoding) 

923 

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 

945 

946 data_init = struct.pack( 

947 "<iIB23s", client_flags, MAX_PACKET_LEN, charset_id, b"" 

948 ) 

949 

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 

955 

956 data = data_init + self.user + b"\0" 

957 

958 authresp = b"" 

959 plugin_name = None 

960 

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 

984 

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" 

991 

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" 

996 

997 if self.server_capabilities & CLIENT.PLUGIN_AUTH: 

998 data += (plugin_name or b"") + b"\0" 

999 

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 

1008 

1009 self.write_packet(data) 

1010 auth_packet = self._read_packet() 

1011 

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 

1022 

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 

1028 

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") 

1051 

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 ) 

1071 

1072 if auth_packet.is_ok_packet(): 

1073 break 

1074 raise err.OperationalError("unexpected packet during authentication") 

1075 

1076 if DEBUG: 

1077 print("Succeed to auth") 

1078 

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() 

1113 

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 ) 

1148 

1149 self.write_packet(data) 

1150 pkt = self._read_packet() 

1151 pkt.check_error() 

1152 return pkt 

1153 

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 

1170 

1171 # _mysql support 

1172 def thread_id(self): 

1173 return self.server_thread_id[0] 

1174 

1175 def character_set_name(self): 

1176 return self.charset 

1177 

1178 def get_host_info(self): 

1179 return self.host_info 

1180 

1181 def get_proto_info(self): 

1182 return self.protocol_version 

1183 

1184 def _get_server_information(self): 

1185 i = 0 

1186 packet = self._read_packet() 

1187 data = packet.get_all_data() 

1188 

1189 self.protocol_version = data[i] 

1190 i += 1 

1191 

1192 server_end = data.find(b"\0", i) 

1193 self.server_version = data[i:server_end].decode("latin1") 

1194 i = server_end + 1 

1195 

1196 self.server_thread_id = struct.unpack("<I", data[i : i + 4]) 

1197 i += 4 

1198 

1199 self.salt = data[i : i + 8] 

1200 i += 9 # 8 + 1(filler) 

1201 

1202 self.server_capabilities = struct.unpack("<H", data[i : i + 2])[0] 

1203 i += 2 

1204 

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 

1216 

1217 self.server_status = stat 

1218 if DEBUG: 

1219 print("server_status: %x" % stat) 

1220 

1221 self.server_capabilities |= cap_h << 16 

1222 if DEBUG: 

1223 print("salt_len:", salt_len) 

1224 salt_len = max(12, salt_len - 9) 

1225 

1226 # reserved 

1227 i += 10 

1228 

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 

1233 

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") 

1248 

1249 if _DEFAULT_AUTH_PLUGIN is not None: # for tests 

1250 self._auth_plugin_name = _DEFAULT_AUTH_PLUGIN 

1251 

1252 def get_server_info(self): 

1253 return self.server_version 

1254 

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 

1265 

1266 

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 

1283 

1284 def __del__(self): 

1285 if self.unbuffered_active: 

1286 self._finish_unbuffered_query() 

1287 

1288 def read(self): 

1289 try: 

1290 first_packet = self.connection._read_packet() 

1291 

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 

1300 

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() 

1307 

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() 

1319 

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 

1325 

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 

1334 

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. 

1350 

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) 

1356 

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 

1369 

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() 

1374 

1375 def _read_rowdata_packet_unbuffered(self): 

1376 # Check if in an active query 

1377 if not self.unbuffered_active: 

1378 return 

1379 

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 

1387 

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 

1392 

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 

1409 

1410 raise 

1411 

1412 if self._check_packet_is_eof(packet): 

1413 self.unbuffered_active = False 

1414 self.connection = None # release reference to kill cyclic reference. 

1415 

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)) 

1425 

1426 self.affected_rows = len(rows) 

1427 self.rows = tuple(rows) 

1428 

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) 

1447 

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 = [] 

1455 

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)) 

1486 

1487 eof_packet = self.connection._read_packet() 

1488 assert eof_packet.is_eof_packet(), "Protocol error, expecting EOF" 

1489 self.description = tuple(description) 

1490 

1491 

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) 

1495 

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 )