1use futures_util::FutureExt;
10
11use mysql_common::{
12 constants::{
13 MariadbCapabilities, DEFAULT_MAX_ALLOWED_PACKET, UTF8MB4_GENERAL_CI, UTF8_GENERAL_CI,
14 },
15 crypto,
16 io::ParseBuf,
17 packets::{
18 AuthPlugin, AuthSwitchRequest, CommonOkPacket, ErrPacket, HandshakePacket,
19 HandshakeResponse, OkPacket, OkPacketDeserializer, OldAuthSwitchRequest, OldEofPacket,
20 ResultSetTerminator, SslRequest,
21 },
22 proto::MySerialize,
23 row::Row,
24};
25
26use std::{
27 borrow::Cow,
28 fmt,
29 future::Future,
30 mem::{self, replace},
31 pin::Pin,
32 str::FromStr,
33 sync::Arc,
34 time::{Duration, Instant},
35};
36
37use crate::{
38 buffer_pool::PooledBuf,
39 conn::{pool::Pool, stmt_cache::StmtCache},
40 consts::{CapabilityFlags, Command, StatusFlags},
41 error::*,
42 io::Stream,
43 opts::Opts,
44 queryable::{
45 query_result::{QueryResult, ResultSetMeta},
46 transaction::TxStatus,
47 BinaryProtocol, Queryable, TextProtocol,
48 },
49 ChangeUserOpts, InfileData, OptsBuilder,
50};
51
52use self::routines::Routine;
53
54#[cfg(feature = "binlog")]
55pub mod binlog_stream;
56pub mod pool;
57pub mod routines;
58pub mod stmt_cache;
59
60const DEFAULT_WAIT_TIMEOUT: usize = 28800;
61
62fn disconnect(mut conn: Conn) {
64 let disconnected = conn.inner.disconnected;
65
66 conn.inner.disconnected = true;
68
69 if !disconnected {
70 if std::thread::panicking() {
72 return;
73 }
74
75 if let Ok(handle) = tokio::runtime::Handle::try_current() {
78 handle.spawn(async move {
79 if let Ok(conn) = conn.cleanup_for_pool().await {
80 let _ = conn.disconnect().await;
81 }
82 });
83 }
84 }
85}
86
87#[derive(Debug, Clone)]
89pub(crate) enum PendingResult {
90 Pending(ResultSetMeta),
92 Taken(Arc<ResultSetMeta>),
94}
95
96struct ConnInner {
98 stream: Option<Stream>,
99 id: u32,
100 is_mariadb: bool,
101 version: (u16, u16, u16),
102 socket: Option<String>,
103 capabilities: CapabilityFlags,
104 mariadb_capabilities: MariadbCapabilities,
105 status: StatusFlags,
106 last_ok_packet: Option<OkPacket<'static>>,
107 last_err_packet: Option<mysql_common::packets::ServerError<'static>>,
108 handshake_complete: bool,
109 pool: Option<Pool>,
110 pending_result: std::result::Result<Option<PendingResult>, ServerError>,
111 tx_status: TxStatus,
112 reset_upon_returning_to_a_pool: bool,
113 opts: Opts,
114 ttl_deadline: Option<Instant>,
115 last_io: Instant,
116 wait_timeout: Duration,
117 stmt_cache: StmtCache,
118 nonce: Vec<u8>,
119 auth_plugin: AuthPlugin<'static>,
120 auth_switched: bool,
121 server_key: Option<Vec<u8>>,
122 active_since: Instant,
123 pub(crate) disconnected: bool,
125 infile_handler:
127 Option<Pin<Box<dyn Future<Output = crate::Result<InfileData>> + Send + Sync + 'static>>>,
128}
129
130impl fmt::Debug for ConnInner {
131 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
132 f.debug_struct("Conn")
133 .field("connection id", &self.id)
134 .field("server version", &self.version)
135 .field("pool", &self.pool)
136 .field("pending_result", &self.pending_result)
137 .field("tx_status", &self.tx_status)
138 .field("stream", &self.stream)
139 .field("options", &self.opts)
140 .field("server_key", &self.server_key)
141 .field("auth_plugin", &self.auth_plugin)
142 .field("capabilities", &self.capabilities)
143 .field("mariadb_capabilities", &self.mariadb_capabilities)
144 .finish()
145 }
146}
147
148impl ConnInner {
149 fn empty(opts: Opts) -> ConnInner {
151 let ttl_deadline = opts.pool_opts().new_connection_ttl_deadline();
152 ConnInner {
153 capabilities: CapabilityFlags::empty(),
154 mariadb_capabilities: MariadbCapabilities::empty(),
155 status: StatusFlags::empty(),
156 last_ok_packet: None,
157 last_err_packet: None,
158 handshake_complete: false,
159 stream: None,
160 is_mariadb: false,
161 version: (0, 0, 0),
162 id: 0,
163 pending_result: Ok(None),
164 pool: None,
165 tx_status: TxStatus::None,
166 last_io: Instant::now(),
167 wait_timeout: Duration::from_secs(0),
168 stmt_cache: StmtCache::new(opts.stmt_cache_size()),
169 socket: opts.socket().map(Into::into),
170 opts,
171 ttl_deadline,
172 nonce: Vec::default(),
173 auth_plugin: AuthPlugin::MysqlNativePassword,
174 auth_switched: false,
175 disconnected: false,
176 server_key: None,
177 infile_handler: None,
178 reset_upon_returning_to_a_pool: false,
179 active_since: Instant::now(),
180 }
181 }
182
183 fn stream_mut(&mut self) -> Result<&mut Stream> {
187 self.stream
188 .as_mut()
189 .ok_or_else(|| DriverError::ConnectionClosed.into())
190 }
191
192 fn stream_ref(&self) -> Result<&Stream> {
193 self.stream
194 .as_ref()
195 .ok_or_else(|| DriverError::ConnectionClosed.into())
196 }
197}
198
199#[derive(Debug)]
201pub struct Conn {
202 inner: Box<ConnInner>,
203}
204
205impl Conn {
206 pub fn id(&self) -> u32 {
208 self.inner.id
209 }
210
211 pub fn is_disconnected(&self) -> bool {
213 self.inner.disconnected
214 }
215
216 pub fn last_insert_id(&self) -> Option<u64> {
220 self.inner
221 .last_ok_packet
222 .as_ref()
223 .and_then(|ok| ok.last_insert_id())
224 }
225
226 pub fn affected_rows(&self) -> u64 {
229 self.inner
230 .last_ok_packet
231 .as_ref()
232 .map(|ok| ok.affected_rows())
233 .unwrap_or_default()
234 }
235
236 pub fn info(&self) -> Cow<'_, str> {
238 self.inner
239 .last_ok_packet
240 .as_ref()
241 .and_then(|ok| ok.info_str())
242 .unwrap_or_else(|| "".into())
243 }
244
245 pub fn get_warnings(&self) -> u16 {
247 self.inner
248 .last_ok_packet
249 .as_ref()
250 .map(|ok| ok.warnings())
251 .unwrap_or_default()
252 }
253
254 pub fn last_ok_packet(&self) -> Option<&OkPacket<'static>> {
256 self.inner.last_ok_packet.as_ref()
257 }
258
259 pub fn reset_connection(&mut self, reset_connection: bool) {
263 self.inner.reset_upon_returning_to_a_pool = reset_connection;
264 }
265
266 pub(crate) fn stream_ref(&self) -> Result<&Stream> {
267 self.inner.stream_ref()
268 }
269
270 pub(crate) fn stream_mut(&mut self) -> Result<&mut Stream> {
271 self.inner.stream_mut()
272 }
273
274 pub(crate) fn capabilities(&self) -> CapabilityFlags {
275 self.inner.capabilities
276 }
277
278 pub(crate) fn has_capabilities(&self, c: CapabilityFlags) -> bool {
279 self.inner.capabilities.contains(c)
280 }
281
282 pub(crate) fn has_mariadb_capabilities(&self, c: MariadbCapabilities) -> bool {
283 self.inner.mariadb_capabilities.contains(c)
284 }
285
286 pub(crate) fn touch(&mut self) {
288 self.inner.last_io = Instant::now();
289 }
290
291 pub(crate) fn reset_seq_id(&mut self) {
293 if let Some(stream) = self.inner.stream.as_mut() {
294 stream.reset_seq_id();
295 }
296 }
297
298 pub(crate) fn sync_seq_id(&mut self) {
300 if let Some(stream) = self.inner.stream.as_mut() {
301 stream.sync_seq_id();
302 }
303 }
304
305 pub(crate) fn handle_ok(&mut self, ok_packet: OkPacket<'static>) {
307 self.inner.status = ok_packet.status_flags();
308 self.inner.last_err_packet = None;
309 self.inner.last_ok_packet = Some(ok_packet);
310 }
311
312 pub(crate) fn handle_err(&mut self, err_packet: ErrPacket<'_>) -> Result<()> {
314 match err_packet {
315 ErrPacket::Error(err) => {
316 self.inner.status = StatusFlags::empty();
317 self.inner.last_ok_packet = None;
318 self.inner.last_err_packet = Some(err.clone().into_owned());
319 Err(Error::from(err))
320 }
321 ErrPacket::Progress(_) => Ok(()),
322 }
323 }
324
325 pub(crate) fn get_tx_status(&self) -> TxStatus {
327 self.inner.tx_status
328 }
329
330 pub(crate) fn set_tx_status(&mut self, tx_status: TxStatus) {
332 self.inner.tx_status = tx_status;
333 }
334
335 pub(crate) fn use_pending_result(
339 &mut self,
340 ) -> std::result::Result<Option<&PendingResult>, ServerError> {
341 if let Err(ref e) = self.inner.pending_result {
342 let e = e.clone();
343 self.inner.pending_result = Ok(None);
344 Err(e)
345 } else {
346 Ok(self.inner.pending_result.as_ref().unwrap().as_ref())
347 }
348 }
349
350 pub(crate) fn get_pending_result(
351 &self,
352 ) -> std::result::Result<Option<&PendingResult>, &ServerError> {
353 self.inner.pending_result.as_ref().map(|x| x.as_ref())
354 }
355
356 pub(crate) fn has_pending_result(&self) -> bool {
357 self.inner.pending_result.is_err() || matches!(self.inner.pending_result, Ok(Some(_)))
358 }
359
360 pub(crate) fn set_pending_result(
362 &mut self,
363 meta: Option<ResultSetMeta>,
364 ) -> std::result::Result<Option<PendingResult>, ServerError> {
365 replace(
366 &mut self.inner.pending_result,
367 Ok(meta.map(PendingResult::Pending)),
368 )
369 }
370
371 pub(crate) fn set_pending_result_error(
372 &mut self,
373 error: ServerError,
374 ) -> std::result::Result<Option<PendingResult>, ServerError> {
375 replace(&mut self.inner.pending_result, Err(error))
376 }
377
378 pub(crate) fn take_pending_result(
380 &mut self,
381 ) -> std::result::Result<Option<Arc<ResultSetMeta>>, ServerError> {
382 let mut output = None;
383
384 self.inner.pending_result = match replace(&mut self.inner.pending_result, Ok(None))? {
385 Some(PendingResult::Pending(x)) => {
386 let meta = Arc::new(x);
387 output = Some(meta.clone());
388 Ok(Some(PendingResult::Taken(meta)))
389 }
390 x => Ok(x),
391 };
392
393 Ok(output)
394 }
395
396 pub(crate) fn status(&self) -> StatusFlags {
398 self.inner.status
399 }
400
401 pub(crate) async fn routine<'a, F, T>(&mut self, f: F) -> crate::Result<T>
402 where
403 F: Routine<T> + 'a,
404 {
405 self.inner.disconnected = true;
406 let result = f.call(&mut *self).await;
407 match result {
408 result @ Ok(_) | result @ Err(crate::Error::Server(_)) => {
409 self.inner.disconnected = false;
411 result
412 }
413 Err(err) => {
414 if self.inner.stream.is_some() {
415 self.take_stream().close().await?;
416 }
417 Err(err)
418 }
419 }
420 }
421
422 pub fn server_version(&self) -> (u16, u16, u16) {
424 self.inner.version
425 }
426
427 pub fn opts(&self) -> &Opts {
429 &self.inner.opts
430 }
431
432 pub fn set_infile_handler<T>(&mut self, handler: T)
439 where
440 T: Future<Output = crate::Result<InfileData>>,
441 T: Send + Sync + 'static,
442 {
443 self.inner.infile_handler = Some(Box::pin(handler));
444 }
445
446 fn take_stream(&mut self) -> Stream {
447 self.inner.stream.take().unwrap()
448 }
449
450 pub async fn disconnect(mut self) -> Result<()> {
452 if !self.inner.disconnected {
453 self.inner.disconnected = true;
454 self.write_command_data(Command::COM_QUIT, &[]).await?;
455 let stream = self.take_stream();
456 stream.close().await?;
457 }
458 Ok(())
459 }
460
461 async fn close_conn(mut self) -> Result<()> {
463 self = self.cleanup_for_pool().await?;
464 self.disconnect().await
465 }
466
467 fn is_secure(&self) -> bool {
469 #[cfg(any(feature = "native-tls-tls", feature = "rustls-tls"))]
470 {
471 self.inner
472 .stream
473 .as_ref()
474 .map(|x| x.is_secure())
475 .unwrap_or_default()
476 }
477
478 #[cfg(not(any(feature = "native-tls-tls", feature = "rustls-tls")))]
479 false
480 }
481
482 fn is_socket(&self) -> bool {
484 #[cfg(unix)]
485 {
486 self.inner
487 .stream
488 .as_ref()
489 .map(|x| x.is_socket())
490 .unwrap_or_default()
491 }
492
493 #[cfg(not(unix))]
494 false
495 }
496
497 fn take(&mut self) -> Conn {
499 mem::replace(self, Conn::empty(Default::default()))
500 }
501
502 fn empty(opts: Opts) -> Self {
503 Self {
504 inner: Box::new(ConnInner::empty(opts)),
505 }
506 }
507
508 fn setup_stream(&mut self) -> Result<()> {
512 debug_assert!(self.inner.stream.is_some());
513 if let Some(stream) = self.inner.stream.as_mut() {
514 stream.set_tcp_nodelay(self.inner.opts.tcp_nodelay())?;
515 }
516 Ok(())
517 }
518
519 async fn handle_handshake(&mut self) -> Result<()> {
520 let packet = self.read_packet().await?;
521 let handshake = ParseBuf(&packet).parse::<HandshakePacket>(())?;
522
523 self.inner.nonce = {
525 let mut nonce = Vec::from(handshake.scramble_1_ref());
526 nonce.extend_from_slice(handshake.scramble_2_ref().unwrap_or(&[][..]));
527 nonce.resize(20, 0);
530 nonce
531 };
532
533 self.inner.capabilities = handshake.capabilities() & self.inner.opts.get_capabilities();
534 self.inner.version = handshake
535 .maria_db_server_version_parsed()
536 .inspect(|_| self.inner.is_mariadb = true)
537 .or_else(|| handshake.server_version_parsed())
538 .unwrap_or((0, 0, 0));
539 self.inner.id = handshake.connection_id();
540 self.inner.status = handshake.status_flags();
541
542 if self.inner.is_mariadb && !self.has_capabilities(CapabilityFlags::CLIENT_LONG_PASSWORD) {
545 self.inner.mariadb_capabilities =
546 handshake.mariadb_ext_capabilities() & self.inner.opts.get_mariadb_capabilities();
547 }
548
549 self.inner.auth_plugin = match handshake.auth_plugin() {
553 Some(AuthPlugin::CachingSha2Password) => AuthPlugin::CachingSha2Password,
554 _ => AuthPlugin::MysqlNativePassword,
555 };
556
557 Ok(())
558 }
559
560 async fn switch_to_ssl_if_needed(&mut self) -> Result<()> {
561 if self
562 .inner
563 .opts
564 .get_capabilities()
565 .contains(CapabilityFlags::CLIENT_SSL)
566 {
567 if !self.has_capabilities(CapabilityFlags::CLIENT_SSL) {
568 return Err(DriverError::NoClientSslFlagFromServer.into());
569 }
570
571 let collation = if self.inner.version >= (5, 5, 3) {
572 UTF8MB4_GENERAL_CI
573 } else {
574 UTF8_GENERAL_CI
575 };
576
577 let ssl_request = SslRequest::new(
578 self.inner.capabilities,
579 DEFAULT_MAX_ALLOWED_PACKET as u32,
580 collation as u8,
581 )
582 .with_mariadb_capabilities(self.inner.mariadb_capabilities);
583 self.write_struct(&ssl_request).await?;
584 let conn = self;
585 let ssl_opts = conn.opts().ssl_opts_and_connector().expect("unreachable");
586 let domain = ssl_opts
587 .ssl_opts()
588 .tls_hostname_override()
589 .unwrap_or_else(|| conn.opts().ip_or_hostname())
590 .into();
591 let tls_connector = ssl_opts.build_tls_connector().await?;
592 conn.stream_mut()?
593 .make_secure(domain, &tls_connector)
594 .await?;
595 Ok(())
596 } else {
597 Ok(())
598 }
599 }
600
601 async fn do_handshake_response(&mut self) -> Result<()> {
602 let auth_data = self
603 .inner
604 .auth_plugin
605 .gen_data(self.inner.opts.pass(), &self.inner.nonce);
606
607 let handshake_response = HandshakeResponse::new(
608 auth_data.as_deref(),
609 self.inner.version,
610 self.inner.opts.user().map(|x| x.as_bytes()),
611 self.inner.opts.db_name().map(|x| x.as_bytes()),
612 Some(self.inner.auth_plugin.borrow()),
613 self.capabilities(),
614 self.inner.opts.connect_attributes().cloned(),
615 self.inner
616 .opts
617 .max_allowed_packet()
618 .unwrap_or(DEFAULT_MAX_ALLOWED_PACKET) as u32,
619 )
620 .with_mariadb_ext_capabilities(self.inner.mariadb_capabilities);
621
622 let mut buf = crate::buffer_pool().get();
624 handshake_response.serialize(buf.as_mut());
625
626 self.write_packet(buf).await?;
627 self.inner.handshake_complete = true;
628 Ok(())
629 }
630
631 async fn perform_auth_switch(
632 &mut self,
633 auth_switch_request: AuthSwitchRequest<'_>,
634 ) -> Result<()> {
635 if !self.inner.auth_switched {
636 self.inner.auth_switched = true;
637 self.inner.nonce = auth_switch_request.plugin_data().to_vec();
638
639 if matches!(
640 auth_switch_request.auth_plugin(),
641 AuthPlugin::MysqlOldPassword
642 ) && self.inner.opts.secure_auth()
643 {
644 return Err(DriverError::MysqlOldPasswordDisabled.into());
645 }
646
647 self.inner.auth_plugin = auth_switch_request.auth_plugin().clone().into_owned();
648
649 let plugin_data = match &self.inner.auth_plugin {
650 x @ AuthPlugin::CachingSha2Password => {
651 x.gen_data(self.inner.opts.pass(), &self.inner.nonce)
652 }
653 x @ AuthPlugin::MysqlNativePassword => {
654 x.gen_data(self.inner.opts.pass(), &self.inner.nonce)
655 }
656 x @ AuthPlugin::MysqlOldPassword => {
657 if self.inner.opts.secure_auth() {
658 return Err(DriverError::MysqlOldPasswordDisabled.into());
659 } else {
660 x.gen_data(self.inner.opts.pass(), &self.inner.nonce)
661 }
662 }
663 x @ AuthPlugin::MysqlClearPassword => {
664 if self.inner.opts.enable_cleartext_plugin() {
665 x.gen_data(self.inner.opts.pass(), &self.inner.nonce)
666 } else {
667 return Err(DriverError::CleartextPluginDisabled.into());
668 }
669 }
670 x @ AuthPlugin::Ed25519 => x.gen_data(self.inner.opts.pass(), &self.inner.nonce),
671 _x @ AuthPlugin::MariadbParsec { .. } => None,
673 x @ AuthPlugin::Other(_) => x.gen_data(self.inner.opts.pass(), &self.inner.nonce),
674 };
675
676 if let Some(plugin_data) = plugin_data {
677 self.write_struct(&plugin_data.into_owned()).await?;
678 } else {
679 self.write_packet(crate::buffer_pool().get()).await?;
680 }
681
682 self.continue_auth().await?;
683
684 Ok(())
685 } else {
686 unreachable!("auth_switched flag should be checked by caller")
687 }
688 }
689
690 fn continue_auth(&mut self) -> Pin<Box<dyn Future<Output = Result<()>> + Send + '_>> {
691 Box::pin(async move {
694 match self.inner.auth_plugin {
695 AuthPlugin::MysqlNativePassword | AuthPlugin::MysqlOldPassword => {
696 self.continue_mysql_native_password_auth().await?;
697 Ok(())
698 }
699 AuthPlugin::CachingSha2Password => {
700 self.continue_caching_sha2_password_auth().await?;
701 Ok(())
702 }
703 AuthPlugin::MysqlClearPassword => {
704 if self.inner.opts.enable_cleartext_plugin() {
705 self.continue_mysql_native_password_auth().await?;
706 Ok(())
707 } else {
708 Err(DriverError::CleartextPluginDisabled.into())
709 }
710 }
711 AuthPlugin::Ed25519 => {
712 self.continue_ed25519_auth().await?;
713 Ok(())
714 }
715 AuthPlugin::MariadbParsec { .. } => {
716 self.continue_parsec_auth().await?;
717 Ok(())
718 }
719 AuthPlugin::Other(ref name) => Err(DriverError::UnknownAuthPlugin {
720 name: String::from_utf8_lossy(name.as_ref()).to_string(),
721 }
722 .into()),
723 }
724 })
725 }
726
727 fn switch_to_compression(&mut self) -> Result<()> {
728 if self.has_capabilities(CapabilityFlags::CLIENT_COMPRESS) {
729 if let Some(compression) = self.inner.opts.compression() {
730 if let Some(stream) = self.inner.stream.as_mut() {
731 stream.compress(compression);
732 }
733 }
734 }
735 Ok(())
736 }
737
738 async fn continue_ed25519_auth(&mut self) -> Result<()> {
739 let packet = self.read_packet().await?;
740 match packet.first() {
741 Some(0x00) => {
742 Ok(())
744 }
745 Some(0xfe) if !self.inner.auth_switched => {
746 let auth_switch_request = ParseBuf(&packet).parse::<AuthSwitchRequest>(())?;
747 self.perform_auth_switch(auth_switch_request).await
748 }
749 _ => Err(DriverError::UnexpectedPacket {
750 payload: packet.to_vec(),
751 }
752 .into()),
753 }
754 }
755
756 async fn continue_parsec_auth(&mut self) -> Result<()> {
757 let packet = self.read_packet().await?;
758 let mut payload: &[u8] = &packet;
760 if packet.first() == Some(&0x01) {
761 payload = &packet[1..];
762 }
763 if payload.first() == Some(&0xfe) {
767 let auth_switch_request = ParseBuf(payload).parse(())?;
768 self.perform_auth_switch(auth_switch_request).await
769 } else {
770 self.inner
772 .auth_plugin
773 .read_add_data(payload)
774 .ok_or_else(|| DriverError::InvalidParsecSalt)?;
775 let plugin_data = self
777 .inner
778 .auth_plugin
779 .gen_data(self.inner.opts.pass(), &self.inner.nonce)
780 .unwrap();
781
782 self.write_struct(&plugin_data.into_owned()).await?;
783 let payload = self.read_packet().await?;
785 match payload.first() {
786 Some(0x00) => Ok(()),
787 _ => Err(DriverError::UnexpectedPacket {
788 payload: payload.to_vec(),
789 }
790 .into()),
791 }
792 }
793 }
794
795 async fn continue_caching_sha2_password_auth(&mut self) -> Result<()> {
796 let packet = self.read_packet().await?;
797 match packet.first() {
798 Some(0x00) => {
799 Ok(())
801 }
802 Some(0x01) => match packet.get(1) {
803 Some(0x03) => {
804 self.drop_packet().await
806 }
807 Some(0x04) => {
808 let pass = self.inner.opts.pass().unwrap_or_default();
809 let mut pass = crate::buffer_pool().get_with(pass.as_bytes());
810 pass.as_mut().push(0);
811
812 if self.is_secure() || self.is_socket() {
813 self.write_packet(pass).await?;
814 } else {
815 if self.inner.server_key.is_none() {
816 self.write_bytes(&[0x02][..]).await?;
817 let packet = self.read_packet().await?;
818 self.inner.server_key = Some(packet[1..].to_vec());
819 }
820 for (i, byte) in pass.as_mut().iter_mut().enumerate() {
821 *byte ^= self.inner.nonce[i % self.inner.nonce.len()];
822 }
823 let encrypted_pass = crypto::encrypt(
824 &pass,
825 self.inner.server_key.as_deref().expect("unreachable"),
826 );
827 self.write_bytes(&encrypted_pass).await?;
828 };
829 self.drop_packet().await?;
830 Ok(())
831 }
832 _ => Err(DriverError::UnexpectedPacket {
833 payload: packet.to_vec(),
834 }
835 .into()),
836 },
837 Some(0xfe) if !self.inner.auth_switched => {
838 let auth_switch_request = ParseBuf(&packet).parse::<AuthSwitchRequest>(())?;
839 self.perform_auth_switch(auth_switch_request).await?;
840 Ok(())
841 }
842 _ => Err(DriverError::UnexpectedPacket {
843 payload: packet.to_vec(),
844 }
845 .into()),
846 }
847 }
848
849 async fn continue_mysql_native_password_auth(&mut self) -> Result<()> {
850 let packet = self.read_packet().await?;
851 match packet.first() {
852 Some(0x00) => Ok(()),
853 Some(0xfe) if !self.inner.auth_switched => {
854 let auth_switch = if packet.len() > 1 {
855 ParseBuf(&packet).parse(())?
856 } else {
857 let _ = ParseBuf(&packet).parse::<OldAuthSwitchRequest>(())?;
858 AuthSwitchRequest::new(
860 "mysql_old_password".as_bytes(),
861 self.inner.nonce.clone(),
862 )
863 };
864 self.perform_auth_switch(auth_switch).await
865 }
866 _ => Err(DriverError::UnexpectedPacket {
867 payload: packet.to_vec(),
868 }
869 .into()),
870 }
871 }
872
873 fn handle_packet(&mut self, packet: &PooledBuf) -> Result<bool> {
875 let ok_packet = if self.has_pending_result() {
876 if self.has_capabilities(CapabilityFlags::CLIENT_DEPRECATE_EOF) {
877 ParseBuf(packet)
878 .parse::<OkPacketDeserializer<ResultSetTerminator>>(self.capabilities())
879 .map(|x| x.into_inner())
880 } else {
881 ParseBuf(packet)
882 .parse::<OkPacketDeserializer<OldEofPacket>>(self.capabilities())
883 .map(|x| x.into_inner())
884 }
885 } else {
886 ParseBuf(packet)
887 .parse::<OkPacketDeserializer<CommonOkPacket>>(self.capabilities())
888 .map(|x| x.into_inner())
889 };
890
891 if let Ok(ok_packet) = ok_packet {
892 self.handle_ok(ok_packet.into_owned());
893 } else {
894 let capabilities = if self.inner.handshake_complete {
902 self.capabilities()
903 } else {
904 CapabilityFlags::empty()
905 };
906 let err_packet = ParseBuf(packet).parse::<ErrPacket>(capabilities);
907 if let Ok(err_packet) = err_packet {
908 self.handle_err(err_packet)?;
909 return Ok(true);
910 }
911 }
912
913 Ok(false)
914 }
915
916 pub(crate) async fn read_packet(&mut self) -> Result<PooledBuf> {
917 loop {
918 let packet = crate::io::ReadPacket::new(&mut *self)
919 .await
920 .map_err(|io_err| {
921 self.inner.stream.take();
922 self.inner.disconnected = true;
923 Error::from(io_err)
924 })?;
925 if self.handle_packet(&packet)? {
926 continue;
928 } else {
929 return Ok(packet);
930 }
931 }
932 }
933
934 pub(crate) async fn read_packets(&mut self, n: usize) -> Result<Vec<PooledBuf>> {
936 let mut packets = Vec::with_capacity(n);
937 for _ in 0..n {
938 packets.push(self.read_packet().await?);
939 }
940 Ok(packets)
941 }
942
943 pub(crate) async fn write_packet(&mut self, data: PooledBuf) -> Result<()> {
944 crate::io::WritePacket::new(&mut *self, data)
945 .await
946 .map_err(|io_err| {
947 self.inner.stream.take();
948 self.inner.disconnected = true;
949 From::from(io_err)
950 })
951 }
952
953 pub(crate) async fn write_bytes(&mut self, bytes: &[u8]) -> Result<()> {
955 let buf = crate::buffer_pool().get_with(bytes);
956 self.write_packet(buf).await
957 }
958
959 pub(crate) async fn write_struct<T: MySerialize>(&mut self, x: &T) -> Result<()> {
961 let mut buf = crate::buffer_pool().get();
962 x.serialize(buf.as_mut());
963 self.write_packet(buf).await
964 }
965
966 pub(crate) async fn write_command<T: MySerialize>(&mut self, cmd: &T) -> Result<()> {
968 self.clean_dirty().await?;
969 self.reset_seq_id();
970 self.write_struct(cmd).await
971 }
972
973 pub(crate) async fn write_command_raw(&mut self, body: PooledBuf) -> Result<()> {
975 debug_assert!(!body.is_empty());
976 self.clean_dirty().await?;
977 self.reset_seq_id();
978 self.write_packet(body).await
979 }
980
981 pub(crate) async fn write_command_data<T>(&mut self, cmd: Command, cmd_data: T) -> Result<()>
983 where
984 T: AsRef<[u8]>,
985 {
986 let cmd_data = cmd_data.as_ref();
987 let mut buf = crate::buffer_pool().get();
988 let body = buf.as_mut();
989 body.push(cmd as u8);
990 body.extend_from_slice(cmd_data);
991 self.write_command_raw(buf).await
992 }
993
994 async fn drop_packet(&mut self) -> Result<()> {
995 self.read_packet().await?;
996 Ok(())
997 }
998
999 async fn run_init_commands(&mut self) -> Result<()> {
1000 if let Some(callback) = self.inner.opts.after_connect() {
1001 callback(self).await?;
1002 }
1003
1004 let mut init = self.inner.opts.init().to_vec();
1005
1006 while let Some(query) = init.pop() {
1007 self.query_drop(query).await?;
1008 }
1009
1010 Ok(())
1011 }
1012
1013 async fn run_setup_commands(&mut self) -> Result<()> {
1014 let mut setup = self.inner.opts.setup().to_vec();
1015
1016 while let Some(query) = setup.pop() {
1017 self.query_drop(query).await?;
1018 }
1019
1020 Ok(())
1021 }
1022
1023 pub fn new<T: Into<Opts>>(opts: T) -> crate::BoxFuture<'static, Conn> {
1025 let opts = opts.into();
1026 async move {
1027 let mut conn = Conn::empty(opts.clone());
1028
1029 let stream = if let Some(_path) = opts.socket() {
1030 #[cfg(unix)]
1031 {
1032 Stream::connect_socket(_path.to_owned()).await?
1033 }
1034 #[cfg(not(unix))]
1035 return Err(crate::DriverError::NamedPipesDisabled.into());
1036 } else {
1037 let keepalive = opts.tcp_keepalive();
1038 Stream::connect_tcp(opts.hostport_or_url(), keepalive).await?
1039 };
1040
1041 conn.inner.stream = Some(stream);
1042 conn.setup_stream()?;
1043 conn.handle_handshake().await?;
1044 conn.switch_to_ssl_if_needed().await?;
1045 conn.do_handshake_response().await?;
1046 conn.continue_auth().await?;
1047 conn.switch_to_compression()?;
1048 conn.read_settings().await?;
1049 conn.reconnect_via_socket_if_needed().await?;
1050 conn.run_init_commands().await?;
1051 conn.run_setup_commands().await?;
1052
1053 Ok(conn)
1054 }
1055 .boxed()
1056 }
1057
1058 pub async fn from_url<T: AsRef<str>>(url: T) -> Result<Conn> {
1060 Conn::new(Opts::from_str(url.as_ref())?).await
1061 }
1062
1063 async fn reconnect_via_socket_if_needed(&mut self) -> Result<()> {
1067 if let Some(socket) = self.inner.socket.as_ref() {
1068 let opts = self.inner.opts.clone();
1069 if opts.socket().is_none() {
1070 let opts = OptsBuilder::from_opts(opts).socket(Some(&**socket));
1071 if let Ok(conn) = Conn::new(opts).await {
1072 let old_conn = std::mem::replace(self, conn);
1073 old_conn.close_conn().await?;
1075 }
1076 }
1077 }
1078 Ok(())
1079 }
1080
1081 async fn read_settings(&mut self) -> Result<()> {
1091 enum Action {
1092 Load(Cfg),
1093 Apply(CfgData),
1094 }
1095
1096 enum CfgData {
1097 MaxAllowedPacket(usize),
1098 WaitTimeout(usize),
1099 }
1100
1101 impl CfgData {
1102 fn apply(&self, conn: &mut Conn) {
1103 match self {
1104 Self::MaxAllowedPacket(value) => {
1105 if let Some(stream) = conn.inner.stream.as_mut() {
1106 stream.set_max_allowed_packet(*value);
1107 }
1108 }
1109 Self::WaitTimeout(value) => {
1110 conn.inner.wait_timeout = Duration::from_secs(*value as u64);
1111 }
1112 }
1113 }
1114 }
1115
1116 enum Cfg {
1117 Socket,
1118 MaxAllowedPacket,
1119 WaitTimeout,
1120 }
1121
1122 impl Cfg {
1123 const fn name(&self) -> &'static str {
1124 match self {
1125 Self::Socket => "@@socket",
1126 Self::MaxAllowedPacket => "@@max_allowed_packet",
1127 Self::WaitTimeout => "@@wait_timeout",
1128 }
1129 }
1130
1131 fn apply(&self, conn: &mut Conn, value: Option<crate::Value>) {
1132 match self {
1133 Cfg::Socket => {
1134 conn.inner.socket = value.and_then(crate::from_value);
1135 }
1136 Cfg::MaxAllowedPacket => {
1137 if let Some(stream) = conn.inner.stream.as_mut() {
1138 stream.set_max_allowed_packet(
1139 value
1140 .and_then(crate::from_value)
1141 .unwrap_or(DEFAULT_MAX_ALLOWED_PACKET),
1142 );
1143 }
1144 }
1145 Cfg::WaitTimeout => {
1146 conn.inner.wait_timeout = Duration::from_secs(
1147 value
1148 .and_then(crate::from_value)
1149 .unwrap_or(DEFAULT_WAIT_TIMEOUT) as u64,
1150 );
1151 }
1152 }
1153 }
1154 }
1155
1156 let mut actions = vec![
1157 if let Some(x) = self.opts().max_allowed_packet() {
1158 Action::Apply(CfgData::MaxAllowedPacket(x))
1159 } else {
1160 Action::Load(Cfg::MaxAllowedPacket)
1161 },
1162 if let Some(x) = self.opts().wait_timeout() {
1163 Action::Apply(CfgData::WaitTimeout(x))
1164 } else {
1165 Action::Load(Cfg::WaitTimeout)
1166 },
1167 ];
1168
1169 if self.inner.opts.prefer_socket() && self.inner.socket.is_none() {
1170 actions.push(Action::Load(Cfg::Socket))
1171 }
1172
1173 let loads = actions
1174 .iter()
1175 .filter_map(|x| match x {
1176 Action::Load(x) => Some(x),
1177 Action::Apply(_) => None,
1178 })
1179 .collect::<Vec<_>>();
1180
1181 let loaded = if !loads.is_empty() {
1182 let query = loads
1183 .iter()
1184 .zip(std::iter::once(' ').chain(std::iter::repeat(',')))
1185 .fold("SELECT".to_owned(), |mut acc, (cfg, prefix)| {
1186 acc.push(prefix);
1187 acc.push_str(cfg.name());
1188 acc
1189 });
1190
1191 self.query_internal::<Row, String>(query)
1192 .await?
1193 .map(|row| row.unwrap())
1194 .unwrap_or_else(|| vec![crate::Value::NULL; loads.len()])
1195 } else {
1196 vec![]
1197 };
1198 let mut loaded = loaded.into_iter();
1199
1200 for action in actions {
1201 match action {
1202 Action::Load(cfg) => cfg.apply(self, loaded.next()),
1203 Action::Apply(cfg) => cfg.apply(self),
1204 }
1205 }
1206
1207 Ok(())
1208 }
1209
1210 fn expired(&self) -> bool {
1213 if let Some(deadline) = self.inner.ttl_deadline {
1214 if Instant::now() > deadline {
1215 return true;
1216 }
1217 }
1218 let ttl = self
1219 .inner
1220 .opts
1221 .conn_ttl()
1222 .unwrap_or(self.inner.wait_timeout);
1223 !ttl.is_zero() && self.idling() > ttl
1224 }
1225
1226 fn idling(&self) -> Duration {
1228 self.inner.last_io.elapsed()
1229 }
1230
1231 pub async fn reset(&mut self) -> Result<bool> {
1238 let supports_com_reset_connection = if self.inner.is_mariadb {
1239 self.inner.version >= (10, 2, 4)
1240 } else {
1241 self.inner.version > (5, 7, 2)
1243 };
1244
1245 if supports_com_reset_connection {
1246 self.routine(routines::ResetRoutine).await?;
1247 self.inner.stmt_cache.clear();
1248 self.inner.infile_handler = None;
1249 self.run_setup_commands().await?;
1250 }
1251
1252 Ok(supports_com_reset_connection)
1253 }
1254
1255 pub async fn change_user(&mut self, opts: ChangeUserOpts) -> Result<()> {
1267 if opts != ChangeUserOpts::default() {
1269 let mut opts_changed = false;
1270 if let Some(user) = opts.user() {
1271 opts_changed |= user != self.opts().user()
1272 };
1273 if let Some(pass) = opts.pass() {
1274 opts_changed |= pass != self.opts().pass()
1275 };
1276 if let Some(db_name) = opts.db_name() {
1277 opts_changed |= db_name != self.opts().db_name()
1278 };
1279 if opts_changed {
1280 if let Some(pool) = self.inner.pool.take() {
1281 pool.cancel_connection();
1282 }
1283 }
1284 }
1285
1286 let conn_opts = &mut self.inner.opts;
1287 opts.update_opts(conn_opts);
1288 self.routine(routines::ChangeUser).await?;
1289 self.inner.stmt_cache.clear();
1290 self.inner.infile_handler = None;
1291 self.run_setup_commands().await?;
1292 Ok(())
1293 }
1294
1295 async fn reset_for_pool(mut self) -> Result<Self> {
1299 if !self.reset().await? {
1300 self.change_user(Default::default()).await?;
1301 }
1302 Ok(self)
1303 }
1304
1305 pub(crate) async fn rollback_transaction(&mut self) -> Result<()> {
1307 debug_assert_ne!(self.inner.tx_status, TxStatus::None);
1308 let tx_status = mem::replace(&mut self.inner.tx_status, TxStatus::None);
1309 if let Err(e) = self.query_drop("ROLLBACK").await {
1310 self.inner.tx_status = tx_status;
1312 return Err(e);
1313 }
1314 Ok(())
1315 }
1316
1317 pub(crate) fn more_results_exists(&self) -> bool {
1320 self.status()
1321 .contains(StatusFlags::SERVER_MORE_RESULTS_EXISTS)
1322 }
1323
1324 pub(crate) async fn drop_result(&mut self) -> Result<()> {
1329 let meta = match self.set_pending_result(None)? {
1331 Some(PendingResult::Pending(meta)) => Some(meta),
1332 Some(PendingResult::Taken(meta)) => {
1333 Some(Arc::try_unwrap(meta).expect("Conn::drop_result call on a pending result that may still be dropped by someone else"))
1336 }
1337 None => None,
1338 };
1339
1340 let _ = self.set_pending_result(meta);
1341
1342 match self.use_pending_result() {
1343 Ok(Some(PendingResult::Pending(ResultSetMeta::Text(_)))) => {
1344 QueryResult::<'_, '_, TextProtocol>::new(self)
1345 .drop_result()
1346 .await
1347 }
1348 Ok(Some(PendingResult::Pending(ResultSetMeta::Binary(_)))) => {
1349 QueryResult::<'_, '_, BinaryProtocol>::new(self)
1350 .drop_result()
1351 .await
1352 }
1353 Ok(None) => Ok(()),
1354 Ok(Some(PendingResult::Taken(_))) | Err(_) => {
1355 unreachable!("this case must be handled earlier in this function")
1356 }
1357 }
1358 }
1359
1360 async fn cleanup_for_pool(mut self) -> Result<Self> {
1364 loop {
1365 if self.has_pending_result() {
1366 if let Err(err) = self.drop_result().await {
1370 if err.is_fatal() {
1371 return Err(err);
1374 }
1375 }
1376 } else if self.inner.tx_status != TxStatus::None {
1377 self.rollback_transaction().await?;
1379 } else {
1380 break;
1381 }
1382 }
1383 Ok(self)
1384 }
1385}
1386
1387#[cfg(test)]
1388mod test {
1389 use bytes::Bytes;
1390 use futures_util::{
1391 stream::{self, StreamExt},
1392 FutureExt,
1393 };
1394 use mysql_common::constants::{MariadbCapabilities, MAX_PAYLOAD_LEN};
1395 use rand::RngExt as _;
1396 use tokio::{io::AsyncWriteExt, net::TcpListener};
1397
1398 use crate::{
1399 from_row, params, prelude::*, test_misc::get_opts, ChangeUserOpts, Conn, Error,
1400 OptsBuilder, Pool, ServerError, Value, WhiteListFsHandler,
1401 };
1402
1403 #[tokio::test]
1404 async fn should_return_found_rows_if_flag_is_set() -> super::Result<()> {
1405 let opts = get_opts().client_found_rows(true);
1406 let mut conn = Conn::new(opts).await.unwrap();
1407
1408 "CREATE TEMPORARY TABLE mysql.found_rows (id INT PRIMARY KEY AUTO_INCREMENT, val INT)"
1409 .ignore(&mut conn)
1410 .await?;
1411
1412 "INSERT INTO mysql.found_rows (val) VALUES (1)"
1413 .ignore(&mut conn)
1414 .await?;
1415
1416 assert_eq!(conn.affected_rows(), 1);
1418
1419 "UPDATE mysql.found_rows SET val = 1 WHERE val = 1"
1420 .ignore(&mut conn)
1421 .await?;
1422
1423 assert_eq!(conn.affected_rows(), 1);
1426
1427 Ok(())
1428 }
1429
1430 #[tokio::test]
1431 async fn should_not_return_found_rows_if_flag_is_not_set() -> super::Result<()> {
1432 let mut conn = Conn::new(get_opts()).await.unwrap();
1433
1434 "CREATE TEMPORARY TABLE mysql.found_rows (id INT PRIMARY KEY AUTO_INCREMENT, val INT)"
1435 .ignore(&mut conn)
1436 .await?;
1437
1438 "INSERT INTO mysql.found_rows (val) VALUES (1)"
1439 .ignore(&mut conn)
1440 .await?;
1441
1442 assert_eq!(conn.affected_rows(), 1);
1444
1445 "UPDATE mysql.found_rows SET val = 1 WHERE val = 1"
1446 .ignore(&mut conn)
1447 .await?;
1448
1449 assert_eq!(conn.affected_rows(), 0);
1451
1452 Ok(())
1453 }
1454
1455 #[test]
1456 fn opts_should_satisfy_send_and_sync() {
1457 struct A<T: Sync + Send>(T);
1458 #[allow(clippy::unnecessary_operation)]
1459 A(get_opts());
1460 }
1461
1462 #[tokio::test]
1463 async fn should_connect_without_database() -> super::Result<()> {
1464 let mut conn: Conn = Conn::new(get_opts().db_name(None::<String>)).await?;
1466 conn.ping().await?;
1467 conn.disconnect().await?;
1468
1469 let mut conn: Conn = Conn::new(get_opts().db_name(Some(""))).await?;
1471 conn.ping().await?;
1472 conn.disconnect().await?;
1473
1474 Ok(())
1475 }
1476
1477 #[tokio::test]
1478 async fn should_clean_state_if_wrapper_is_dropped() -> super::Result<()> {
1479 let mut conn: Conn = Conn::new(get_opts()).await?;
1480
1481 conn.query_drop("CREATE TEMPORARY TABLE mysql.foo (id SERIAL)")
1482 .await?;
1483
1484 conn.query_iter("SELECT 1").await?;
1486 conn.ping().await?;
1487
1488 let mut tx = conn.start_transaction(Default::default()).await?;
1490 tx.query_drop("INSERT INTO mysql.foo (id) VALUES (42)")
1491 .await?;
1492 tx.exec_iter("SELECT COUNT(*) FROM mysql.foo", ()).await?;
1493 drop(tx);
1494 conn.ping().await?;
1495
1496 let count: u8 = conn
1497 .query_first("SELECT COUNT(*) FROM mysql.foo")
1498 .await?
1499 .unwrap_or_default();
1500
1501 assert_eq!(count, 0);
1502
1503 Ok(())
1504 }
1505
1506 #[tokio::test]
1507 async fn should_connect() -> super::Result<()> {
1508 let mut conn: Conn = Conn::new(get_opts()).await?;
1509 conn.ping().await?;
1510 let plugins: Vec<String> = conn
1511 .query_map("SHOW PLUGINS", |mut row: crate::Row| {
1512 row.take("Name").unwrap()
1513 })
1514 .await?;
1515
1516 let variants = vec![
1518 ("caching_sha2_password", 2_u8, "non-empty"),
1519 ("caching_sha2_password", 2_u8, ""),
1520 ("mysql_native_password", 0_u8, "non-empty"),
1521 ("mysql_native_password", 0_u8, ""),
1522 ]
1523 .into_iter()
1524 .filter(|variant| plugins.iter().any(|p| p == variant.0));
1525
1526 for (plug, val, pass) in variants {
1527 dbg!((plug, val, pass, conn.inner.version));
1528
1529 if plug == "mysql_native_password" && conn.inner.version >= (8, 4, 0) {
1530 continue;
1531 }
1532
1533 let _ = conn.query_drop("DROP USER 'test_user'@'%'").await;
1534
1535 let query = format!("CREATE USER 'test_user'@'%' IDENTIFIED WITH {}", plug);
1536 conn.query_drop(query).await.unwrap();
1537
1538 if conn.inner.version < (8, 0, 11) {
1539 conn.query_drop(format!("SET old_passwords = {}", val))
1540 .await
1541 .unwrap();
1542 conn.query_drop(format!(
1543 "SET PASSWORD FOR 'test_user'@'%' = PASSWORD('{}')",
1544 pass
1545 ))
1546 .await
1547 .unwrap();
1548 } else {
1549 conn.query_drop(format!("SET PASSWORD FOR 'test_user'@'%' = '{}'", pass))
1550 .await
1551 .unwrap();
1552 };
1553
1554 let opts = get_opts()
1555 .user(Some("test_user"))
1556 .pass(Some(pass))
1557 .db_name(None::<String>);
1558 let result = Conn::new(opts).await;
1559
1560 conn.query_drop("DROP USER 'test_user'@'%'").await.unwrap();
1561
1562 result?.disconnect().await?;
1563 }
1564
1565 if crate::test_misc::test_compression() {
1566 assert!(format!("{:?}", conn).contains("Compression"));
1567 }
1568
1569 if crate::test_misc::test_ssl() {
1570 assert!(format!("{:?}", conn).contains("Tls"));
1571 }
1572
1573 conn.disconnect().await?;
1574 Ok(())
1575 }
1576
1577 #[test]
1578 fn should_not_panic_if_dropped_without_tokio_runtime() {
1579 let fut = Conn::new(get_opts());
1580 let runtime = tokio::runtime::Runtime::new().unwrap();
1581 runtime.block_on(async {
1582 fut.await.unwrap();
1583 });
1584 }
1586
1587 #[tokio::test]
1588 async fn should_execute_init_queries_on_new_connection() -> super::Result<()> {
1589 let opts = OptsBuilder::from_opts(get_opts()).init(vec!["SET @a = 42", "SET @b = 'foo'"]);
1590 let mut conn = Conn::new(opts).await?;
1591 let result: Vec<(u8, String)> = conn.query("SELECT @a, @b").await?;
1592 conn.disconnect().await?;
1593 assert_eq!(result, vec![(42, "foo".into())]);
1594 Ok(())
1595 }
1596
1597 #[tokio::test]
1598 async fn should_execute_after_connect_callback_on_new_connection() -> super::Result<()> {
1599 let opts = OptsBuilder::from_opts(get_opts()).after_connect(|conn| {
1600 async move {
1601 conn.query_drop("SET @a = 42").await?;
1602 conn.query_drop("SET @b = 'foo'").await?;
1603 Ok(())
1604 }
1605 .boxed()
1606 });
1607 let mut conn = Conn::new(opts).await?;
1608 let result: Vec<(u8, String)> = conn.query("SELECT @a, @b").await?;
1609 conn.disconnect().await?;
1610 assert_eq!(result, vec![(42, "foo".into())]);
1611 Ok(())
1612 }
1613
1614 #[tokio::test]
1615 async fn should_propagate_after_connect_callback_error() -> super::Result<()> {
1616 let opts = OptsBuilder::from_opts(get_opts())
1617 .after_connect(|_conn| async move { Err(Error::Other("rejected".into())) }.boxed());
1618 let e = Conn::new(opts).await.unwrap_err();
1619 match e {
1620 Error::Other(e) => assert_eq!(e.to_string(), "rejected"),
1621 e => panic!("expected error from after_connect(), got {e:?}"),
1622 }
1623 Ok(())
1624 }
1625
1626 #[tokio::test]
1627 async fn should_execute_setup_queries_on_reset() -> super::Result<()> {
1628 let opts = OptsBuilder::from_opts(get_opts()).setup(vec!["SET @a = 42", "SET @b = 'foo'"]);
1629 let mut conn = Conn::new(opts).await?;
1630
1631 let mut result: Vec<(u8, String)> = conn.query("SELECT @a, @b").await?;
1633 assert_eq!(result, vec![(42, "foo".into())]);
1634
1635 if conn.reset().await? {
1637 result = conn.query("SELECT @a, @b").await?;
1638 assert_eq!(result, vec![(42, "foo".into())]);
1639 }
1640
1641 conn.change_user(Default::default()).await?;
1643 result = conn.query("SELECT @a, @b").await?;
1644 assert_eq!(result, vec![(42, "foo".into())]);
1645
1646 conn.disconnect().await?;
1647 Ok(())
1648 }
1649
1650 #[tokio::test]
1651 async fn should_reset_the_connection() -> super::Result<()> {
1652 let mut conn = Conn::new(get_opts()).await?;
1653
1654 assert_eq!(
1655 conn.query_first::<Value, _>("SELECT @foo").await?.unwrap(),
1656 Value::NULL
1657 );
1658
1659 conn.query_drop("SET @foo = 'foo'").await?;
1660
1661 assert_eq!(
1662 conn.query_first::<String, _>("SELECT @foo").await?.unwrap(),
1663 "foo",
1664 );
1665
1666 if conn.reset().await? {
1667 assert_eq!(
1668 conn.query_first::<Value, _>("SELECT @foo").await?.unwrap(),
1669 Value::NULL
1670 );
1671 } else {
1672 assert_eq!(
1673 conn.query_first::<String, _>("SELECT @foo").await?.unwrap(),
1674 "foo",
1675 );
1676 }
1677
1678 conn.disconnect().await?;
1679 Ok(())
1680 }
1681
1682 #[tokio::test]
1683 async fn should_change_user() -> super::Result<()> {
1684 type ShouldRunFn = fn(bool, (u16, u16, u16)) -> bool;
1686 type CreateUserFn = fn(bool, (u16, u16, u16), &str) -> Vec<String>;
1688
1689 #[allow(clippy::type_complexity)]
1690 const TEST_MATRIX: [(&str, ShouldRunFn, CreateUserFn); 5] = [
1691 (
1692 "mysql_old_password",
1693 |is_mariadb, version| is_mariadb || version < (5, 7, 0),
1694 |is_mariadb, version, pass| {
1695 if is_mariadb {
1696 vec![
1697 "CREATE USER '__mats'@'%' IDENTIFIED WITH mysql_old_password".into(),
1698 "SET old_passwords=1".into(),
1699 format!("ALTER USER '__mats'@'%' IDENTIFIED BY '{pass}'"),
1700 "SET old_passwords=0".into(),
1701 ]
1702 } else if matches!(version, (5, 6, _)) {
1703 vec![
1704 "CREATE USER '__mats'@'%' IDENTIFIED WITH mysql_old_password".into(),
1705 format!("SET PASSWORD FOR '__mats'@'%' = OLD_PASSWORD('{pass}')"),
1706 ]
1707 } else {
1708 vec![
1709 "CREATE USER '__mats'@'%'".into(),
1710 format!("SET PASSWORD FOR '__mats'@'%' = PASSWORD('{pass}')"),
1711 ]
1712 }
1713 },
1714 ),
1715 (
1716 "mysql_native_password",
1717 |is_mariadb, version| is_mariadb || version < (8, 4, 0),
1718 |is_mariadb, version, pass| {
1719 if is_mariadb {
1720 vec![
1721 format!("CREATE USER '__mats'@'%' IDENTIFIED WITH mysql_native_password AS PASSWORD('{pass}')")
1722 ]
1723 } else if version < (8, 0, 0) {
1724 vec![
1725 format!(
1726 "CREATE USER '__mats'@'%' IDENTIFIED WITH mysql_native_password"
1727 ),
1728 format!("SET old_passwords = 0"),
1729 format!("SET PASSWORD FOR '__mats'@'%' = PASSWORD('{pass}')"),
1730 ]
1731 } else {
1732 vec![
1733 format!("CREATE USER '__mats'@'%' IDENTIFIED WITH mysql_native_password BY '{pass}'")
1734 ]
1735 }
1736 },
1737 ),
1738 (
1739 "caching_sha2_password",
1740 |is_mariadb, version| !is_mariadb && version >= (5, 8, 0),
1741 |_is_mariadb, _version, pass| {
1742 vec![
1743 format!("CREATE USER '__mats'@'%' IDENTIFIED WITH caching_sha2_password BY '{pass}'")
1744 ]
1745 },
1746 ),
1747 (
1748 "client_ed25519",
1749 |is_mariadb, version| is_mariadb && version >= (11, 6, 2),
1750 |_is_mariadb, _version, pass| {
1751 vec![format!(
1752 "CREATE USER '__mats'@'%' IDENTIFIED WITH ed25519 AS PASSWORD('{pass}')"
1753 )]
1754 },
1755 ),
1756 (
1757 "parsec",
1758 |is_mariadb, version| is_mariadb && version >= (11, 4, 1),
1759 |_is_mariadb, _version, pass| {
1760 vec![format!(
1761 "CREATE USER '__mats'@'%' IDENTIFIED WITH parsec AS PASSWORD('{pass}')"
1762 )]
1763 },
1764 ),
1765 ];
1766
1767 fn random_pass() -> String {
1768 let mut rng = rand::rng();
1769 let pass: [u8; 10] = rng.random();
1770
1771 IntoIterator::into_iter(pass)
1772 .map(|x| ((x % (123 - 97)) + 97) as char)
1773 .collect()
1774 }
1775
1776 let mut conn = Conn::new(get_opts()).await?;
1777
1778 assert_eq!(
1779 conn.query_first::<Value, _>("SELECT @foo").await?.unwrap(),
1780 Value::NULL
1781 );
1782
1783 conn.query_drop("SET @foo = 'foo'").await?;
1784
1785 assert_eq!(
1786 conn.query_first::<String, _>("SELECT @foo").await?.unwrap(),
1787 "foo",
1788 );
1789
1790 conn.change_user(Default::default()).await?;
1791 assert_eq!(
1792 conn.query_first::<Value, _>("SELECT @foo").await?.unwrap(),
1793 Value::NULL
1794 );
1795
1796 for (i, (plugin, should_run, create_statements)) in TEST_MATRIX.iter().enumerate() {
1797 dbg!(plugin);
1798 let is_mariadb = conn.inner.is_mariadb;
1799 let version = conn.server_version();
1800
1801 if should_run(is_mariadb, version) {
1802 let pass = random_pass();
1803
1804 let result = conn
1805 .query_drop("DROP USER /*!50700 IF EXISTS */ /*M!100103 IF EXISTS */ __mats")
1806 .await;
1807
1808 if matches!(version, (5, 6, _)) && i == 0 {
1809 drop(result);
1811 } else {
1812 result.unwrap();
1813 }
1814
1815 for statement in create_statements(is_mariadb, version, &pass) {
1816 conn.query_drop(dbg!(statement)).await.unwrap();
1817 }
1818
1819 let mut conn2 = Conn::new(get_opts().secure_auth(false)).await.unwrap();
1820 conn2
1821 .change_user(
1822 ChangeUserOpts::default()
1823 .with_db_name(None)
1824 .with_user(Some("__mats".into()))
1825 .with_pass(Some(pass)),
1826 )
1827 .await
1828 .unwrap();
1829
1830 let (db, user) = conn2
1831 .query_first::<(Option<String>, String), _>("SELECT DATABASE(), USER();")
1832 .await
1833 .unwrap()
1834 .unwrap();
1835 assert_eq!(db, None);
1836 assert!(user.starts_with("__mats"));
1837
1838 conn2.disconnect().await.unwrap();
1839 }
1840 }
1841
1842 Ok(())
1843 }
1844
1845 #[tokio::test]
1847 async fn test_exec_batch() {
1848 let mut conn = Conn::new(get_opts()).await.unwrap();
1849
1850 conn.query_drop(
1851 "CREATE TEMPORARY TABLE t_exec_batch (\
1852 id INT NOT NULL PRIMARY KEY,\
1853 val VARCHAR(32),\
1854 num BIGINT UNSIGNED)",
1855 )
1856 .await
1857 .unwrap();
1858
1859 let insert_stmt = "INSERT INTO t_exec_batch (id,val,num) VALUES (?,?,?)";
1861
1862 let params = [
1863 (1, Some("First"), None),
1864 (3, None, Some(1)),
1865 (4, Some("Third"), Some(u64::MAX)),
1866 ];
1867
1868 conn.exec_batch(insert_stmt, params.iter().copied())
1869 .await
1870 .unwrap();
1871
1872 conn.exec_batch(insert_stmt, [(8, None::<String>, None::<u64>)])
1873 .await
1874 .unwrap();
1875
1876 let fetched_rows: Vec<(i32, Option<String>, Option<u64>)> = conn
1877 .query_iter("SELECT id, val, num FROM t_exec_batch")
1878 .await
1879 .unwrap()
1880 .collect()
1881 .await
1882 .unwrap();
1883 let expected_rows: Vec<(i32, Option<String>, Option<u64>)> = vec![
1884 (1, Some("First".to_string()), None),
1885 (3, None, Some(1)),
1886 (4, Some("Third".to_string()), Some(u64::MAX)),
1887 (8, None, None),
1888 ];
1889 assert_eq!(fetched_rows, expected_rows);
1890
1891 if conn.has_mariadb_capabilities(MariadbCapabilities::MARIADB_CLIENT_STMT_BULK_OPERATIONS) {
1892 let select_stmt = "SELECT ?";
1893 let err = conn
1894 .exec_batch(select_stmt, [(1_u64,), (2_u64,), (3_u64,)])
1895 .await
1896 .unwrap_err();
1897 assert!(matches!(err, crate::Error::Server(e) if e.code == 1295));
1898 }
1899 }
1900
1901 #[tokio::test]
1902 async fn test_exec_batch_large() {
1903 const CLIENT_MAX_PACKET_SIZE: usize = 1024; let opts = get_opts().max_allowed_packet(Some(CLIENT_MAX_PACKET_SIZE));
1905 let mut conn = Conn::new(opts).await.unwrap();
1906 conn.query_drop(
1907 "CREATE TEMPORARY TABLE t_large_batch (id BIGINT NOT NULL PRIMARY KEY,
1908 val VARCHAR(1024) NOT NULL)",
1909 )
1910 .await
1911 .unwrap();
1912
1913 let num_rows = 3;
1917 let row_chunk_size = CLIENT_MAX_PACKET_SIZE / num_rows;
1918 let mut row_data_1 = "a".repeat(row_chunk_size);
1919 let row_data_2 = "b".repeat(row_chunk_size);
1920 let row_data_3 = "c".repeat(row_chunk_size);
1921
1922 let evaluated_packet_len = 4 + 7 + 4 + (1 + 8 + 1 + 3 + row_chunk_size * num_rows);
1923
1924 assert!(
1925 evaluated_packet_len > CLIENT_MAX_PACKET_SIZE,
1926 "Data size must be greater than max packet size"
1927 );
1928
1929 let params: Vec<(u64, &str)> = vec![
1930 (1, &row_data_1[..]),
1931 (7, &row_data_2[..]),
1932 (22, &row_data_3[..]),
1933 ];
1934
1935 let query = "INSERT INTO t_large_batch (id, val) VALUES (?,?)";
1936 conn.exec_batch(query, params)
1937 .await
1938 .expect("Batch execution should succeed");
1939
1940 let mut inserted_rows: Vec<(u64, String)> = vec![];
1943 inserted_rows.extend(
1944 conn.query("SELECT id, val FROM t_large_batch ORDER BY id LIMIT 1 OFFSET 0")
1945 .await
1946 .unwrap(),
1947 );
1948 inserted_rows.extend(
1949 conn.query("SELECT id, val FROM t_large_batch ORDER BY id LIMIT 1 OFFSET 1")
1950 .await
1951 .unwrap(),
1952 );
1953 inserted_rows.extend(
1954 conn.query("SELECT id, val FROM t_large_batch ORDER BY id LIMIT 65536 OFFSET 2")
1955 .await
1956 .unwrap(),
1957 );
1958
1959 assert_eq!(
1960 inserted_rows.len(),
1961 num_rows,
1962 "The number of inserted rows ({}) does not match the expected number ({})",
1963 inserted_rows.len(),
1964 num_rows
1965 );
1966
1967 assert_eq!(inserted_rows[0], (1, row_data_1));
1968 assert_eq!(inserted_rows[1], (7, row_data_2));
1969 assert_eq!(inserted_rows[2], (22, row_data_3));
1970
1971 row_data_1 = "x".repeat(CLIENT_MAX_PACKET_SIZE);
1974 let params: Vec<(u64, &str)> = vec![(33, &row_data_1[..])];
1975 let result = conn.exec_batch(query, params).await;
1976 assert!(
1977 result.is_err(),
1978 "Batch execution should fail due to packet size exceeding max_allowed_packet"
1979 );
1980 }
1981
1982 #[tokio::test]
1983 async fn test_exec_batch_no_params() -> crate::Result<()> {
1984 let mut conn = Conn::new(get_opts()).await.unwrap();
1985 conn.query_drop("CREATE TEMPORARY TABLE t_counter (counter INTEGER NOT NULL)")
1986 .await
1987 .unwrap();
1988 conn.query_drop("INSERT INTO t_counter (counter) VALUES (0)")
1989 .await
1990 .unwrap();
1991
1992 const COUNT: usize = 10;
1993 conn.exec_batch("UPDATE t_counter SET counter = counter+1", vec![(); COUNT])
1994 .await
1995 .expect("Batch execution should succeed");
1996 let rows: Vec<(usize,)> = conn.query("SELECT counter FROM t_counter").await.unwrap();
1997 assert_eq!(rows, vec![(COUNT,)]);
1998 Ok(())
1999 }
2000
2001 #[tokio::test]
2002 async fn should_not_cache_statements_if_stmt_cache_size_is_zero() -> super::Result<()> {
2003 let opts = OptsBuilder::from_opts(get_opts()).stmt_cache_size(0);
2004
2005 let mut conn = Conn::new(opts).await?;
2006 conn.exec_drop("DO ?", (1_u8,)).await?;
2007
2008 let stmt = conn.prep("DO 2").await?;
2009 conn.exec_drop(&stmt, ()).await?;
2010 conn.exec_drop(&stmt, ()).await?;
2011 conn.close(stmt).await?;
2012
2013 conn.exec_drop("DO 3", ()).await?;
2014 conn.exec_batch("DO 4", vec![(), ()]).await?;
2015 conn.exec_first::<u8, _, _>("DO 5", ()).await?;
2016 let row: Option<(crate::Value, usize)> = conn
2017 .query_first("SHOW SESSION STATUS LIKE 'Com_stmt_close';")
2018 .await?;
2019
2020 assert_eq!(row.unwrap().1, 1);
2021 assert_eq!(conn.inner.stmt_cache.len(), 0);
2022
2023 conn.disconnect().await?;
2024
2025 Ok(())
2026 }
2027
2028 #[tokio::test]
2029 async fn should_hold_stmt_cache_size_bound() -> super::Result<()> {
2030 let opts = OptsBuilder::from_opts(get_opts()).stmt_cache_size(3);
2031 let mut conn = Conn::new(opts).await?;
2032 conn.exec_drop("DO 1", ()).await?;
2033 conn.exec_drop("DO 2", ()).await?;
2034 conn.exec_drop("DO 3", ()).await?;
2035 conn.exec_drop("DO 1", ()).await?;
2036 conn.exec_drop("DO 4", ()).await?;
2037 conn.exec_drop("DO 3", ()).await?;
2038 conn.exec_drop("DO 5", ()).await?;
2039 conn.exec_drop("DO 6", ()).await?;
2040 let row_opt = conn
2041 .query_first("SHOW SESSION STATUS LIKE 'Com_stmt_close';")
2042 .await?;
2043 let (_, count): (String, usize) = row_opt.unwrap();
2044 assert_eq!(count, 3);
2045 let order = conn
2046 .stmt_cache_ref()
2047 .iter()
2048 .map(|item| item.1.query.0.as_ref())
2049 .collect::<Vec<&[u8]>>();
2050 assert_eq!(order, &[b"DO 6", b"DO 5", b"DO 3"]);
2051 conn.disconnect().await?;
2052 Ok(())
2053 }
2054
2055 #[tokio::test]
2056 async fn should_perform_queries() -> super::Result<()> {
2057 let mut conn = Conn::new(get_opts()).await?;
2058 for x in (MAX_PAYLOAD_LEN - 2)..=(MAX_PAYLOAD_LEN + 2) {
2059 let long_string = "A".repeat(x);
2060 let result: Vec<(String, u8)> = conn
2061 .query(format!(r"SELECT '{}', 231", long_string))
2062 .await?;
2063 assert_eq!((long_string, 231_u8), result[0]);
2064 }
2065 conn.disconnect().await?;
2066 Ok(())
2067 }
2068
2069 #[tokio::test]
2070 async fn should_query_drop() -> super::Result<()> {
2071 let mut conn = Conn::new(get_opts()).await?;
2072 conn.query_drop("CREATE TEMPORARY TABLE tmp (id int DEFAULT 10, name text)")
2073 .await?;
2074 conn.query_drop("INSERT INTO tmp VALUES (1, 'foo')").await?;
2075 let result: Option<u8> = conn.query_first("SELECT COUNT(*) FROM tmp").await?;
2076 conn.disconnect().await?;
2077 assert_eq!(result, Some(1_u8));
2078 Ok(())
2079 }
2080
2081 #[tokio::test]
2082 async fn should_prepare_statement() -> super::Result<()> {
2083 let mut conn = Conn::new(get_opts()).await?;
2084 let stmt = conn.prep(r"SELECT ?").await?;
2085 conn.close(stmt).await?;
2086 conn.disconnect().await?;
2087
2088 let mut conn = Conn::new(get_opts()).await?;
2089 let stmt = conn.prep(r"SELECT :foo").await?;
2090
2091 {
2092 let query = String::from("SELECT ?, ?");
2093 let stmt = conn.prep(&*query).await?;
2094 conn.close(stmt).await?;
2095 {
2096 let mut conn = Conn::new(get_opts()).await?;
2097 let stmt = conn.prep(&*query).await?;
2098 conn.close(stmt).await?;
2099 conn.disconnect().await?;
2100 }
2101 }
2102
2103 conn.close(stmt).await?;
2104 conn.disconnect().await?;
2105
2106 Ok(())
2107 }
2108
2109 #[tokio::test]
2110 async fn should_execute_statement() -> super::Result<()> {
2111 let long_string = "A".repeat(18 * 1024 * 1024);
2112 let mut conn = Conn::new(get_opts()).await?;
2113 let stmt = conn.prep(r"SELECT ?").await?;
2114 let result = conn.exec_iter(&stmt, (&long_string,)).await?;
2115 let mut mapped = result.map_and_drop(from_row::<(String,)>).await?;
2116 assert_eq!(mapped.len(), 1);
2117 assert_eq!(mapped.pop(), Some((long_string,)));
2118 let result = conn.exec_iter(&stmt, (42_u8,)).await?;
2119 let collected = result.collect_and_drop::<(u8,)>().await?;
2120 assert_eq!(collected, vec![(42u8,)]);
2121 let result = conn.exec_iter(&stmt, (8_u8,)).await?;
2122 let reduced = result
2123 .reduce_and_drop(2, |mut acc, row| {
2124 acc += from_row::<i32>(row);
2125 acc
2126 })
2127 .await?;
2128 conn.close(stmt).await?;
2129 conn.disconnect().await?;
2130 assert_eq!(reduced, 10);
2131
2132 let mut conn = Conn::new(get_opts()).await?;
2133 let stmt = conn.prep(r"SELECT :foo, :bar, :foo, 3").await?;
2134 let result = conn
2135 .exec_iter(&stmt, params! { "foo" => "quux", "bar" => "baz" })
2136 .await?;
2137 let mut mapped = result
2138 .map_and_drop(from_row::<(String, String, String, u8)>)
2139 .await?;
2140 assert_eq!(mapped.len(), 1);
2141 assert_eq!(
2142 mapped.pop(),
2143 Some(("quux".into(), "baz".into(), "quux".into(), 3))
2144 );
2145 let result = conn
2146 .exec_iter(&stmt, params! { "foo" => 2, "bar" => 3 })
2147 .await?;
2148 let collected = result.collect_and_drop::<(u8, u8, u8, u8)>().await?;
2149 assert_eq!(collected, vec![(2, 3, 2, 3)]);
2150 let result = conn
2151 .exec_iter(&stmt, params! { "foo" => 2, "bar" => 3 })
2152 .await?;
2153 let reduced = result
2154 .reduce_and_drop(0, |acc, row| {
2155 let (a, b, c, d): (u8, u8, u8, u8) = from_row(row);
2156 acc + a + b + c + d
2157 })
2158 .await?;
2159 conn.close(stmt).await?;
2160 conn.disconnect().await?;
2161 assert_eq!(reduced, 10);
2162 Ok(())
2163 }
2164
2165 #[tokio::test]
2166 async fn should_prep_exec_statement() -> super::Result<()> {
2167 let mut conn = Conn::new(get_opts()).await?;
2168 let result = conn
2169 .exec_iter(r"SELECT :a, :b, :a", params! { "a" => 2, "b" => 3 })
2170 .await?;
2171 let output = result
2172 .map_and_drop(|row| {
2173 let (a, b, c): (u8, u8, u8) = from_row(row);
2174 a * b * c
2175 })
2176 .await?;
2177 conn.disconnect().await?;
2178 assert_eq!(output[0], 12u8);
2179 Ok(())
2180 }
2181
2182 #[tokio::test]
2183 async fn should_first_exec_statement() -> super::Result<()> {
2184 let mut conn = Conn::new(get_opts()).await?;
2185 let output = conn
2186 .exec_first(
2187 r"SELECT :a UNION ALL SELECT :b",
2188 params! { "a" => 2, "b" => 3 },
2189 )
2190 .await?;
2191 conn.disconnect().await?;
2192 assert_eq!(output, Some(2u8));
2193 Ok(())
2194 }
2195
2196 #[tokio::test]
2197 async fn issue_107() -> super::Result<()> {
2198 let mut conn = Conn::new(get_opts()).await?;
2199 conn.query_drop(
2200 r"CREATE TEMPORARY TABLE mysql.issue (
2201 a BIGINT(20) UNSIGNED,
2202 b VARBINARY(16),
2203 c BINARY(32),
2204 d BIGINT(20) UNSIGNED,
2205 e BINARY(32)
2206 )",
2207 )
2208 .await?;
2209 conn.query_drop(
2210 r"INSERT INTO mysql.issue VALUES (
2211 0,
2212 0xC066F966B0860000,
2213 0x7939DA98E524C5F969FC2DE8D905FD9501EBC6F20001B0A9C941E0BE6D50CF44,
2214 0,
2215 ''
2216 ), (
2217 1,
2218 '',
2219 0x076311DF4D407B0854371BA13A5F3FB1A4555AC22B361375FD47B263F31822F2,
2220 0,
2221 ''
2222 )",
2223 )
2224 .await?;
2225
2226 let q = "SELECT b, c, d, e FROM mysql.issue";
2227 let result = conn.query_iter(q).await?;
2228
2229 let loaded_structs = result
2230 .map_and_drop(crate::from_row::<(Vec<u8>, Vec<u8>, u64, Vec<u8>)>)
2231 .await?;
2232
2233 conn.disconnect().await?;
2234
2235 assert_eq!(loaded_structs.len(), 2);
2236
2237 Ok(())
2238 }
2239
2240 #[tokio::test]
2241 async fn should_run_transactions() -> super::Result<()> {
2242 let mut conn = Conn::new(get_opts()).await?;
2243 conn.query_drop("CREATE TEMPORARY TABLE tmp (id INT, name TEXT)")
2244 .await?;
2245 let mut transaction = conn.start_transaction(Default::default()).await?;
2246 transaction
2247 .query_drop("INSERT INTO tmp VALUES (1, 'foo'), (2, 'bar')")
2248 .await?;
2249 assert_eq!(transaction.last_insert_id(), None);
2250 assert_eq!(transaction.affected_rows(), 2);
2251 assert_eq!(transaction.get_warnings(), 0);
2252 assert_eq!(transaction.info(), "Records: 2 Duplicates: 0 Warnings: 0");
2253 transaction.commit().await?;
2254 let output_opt = conn.query_first("SELECT COUNT(*) FROM tmp").await?;
2255 assert_eq!(output_opt, Some((2u8,)));
2256 let mut transaction = conn.start_transaction(Default::default()).await?;
2257 transaction
2258 .query_drop("INSERT INTO tmp VALUES (3, 'baz'), (4, 'quux')")
2259 .await?;
2260 let output_opt = transaction
2261 .exec_first("SELECT COUNT(*) FROM tmp", ())
2262 .await?;
2263 assert_eq!(output_opt, Some((4u8,)));
2264 transaction.rollback().await?;
2265 let output_opt = conn.query_first("SELECT COUNT(*) FROM tmp").await?;
2266 assert_eq!(output_opt, Some((2u8,)));
2267
2268 let mut transaction = conn.start_transaction(Default::default()).await?;
2269 transaction
2270 .query_drop("INSERT INTO tmp VALUES (3, 'baz')")
2271 .await?;
2272 drop(transaction); let output_opt = conn.query_first("SELECT COUNT(*) FROM tmp").await?;
2274 assert_eq!(output_opt, Some((2u8,)));
2275
2276 conn.disconnect().await?;
2277 Ok(())
2278 }
2279
2280 #[tokio::test]
2281 async fn should_handle_multiresult_set_with_error() -> super::Result<()> {
2282 const QUERY_FIRST: &str = "SELECT * FROM tmp; SELECT 1; SELECT 2;";
2283 const QUERY_MIDDLE: &str = "SELECT 1; SELECT * FROM tmp; SELECT 2";
2284 let mut conn = Conn::new(get_opts()).await.unwrap();
2285
2286 let result = QUERY_FIRST.run(&mut conn).await;
2288 assert!(matches!(result, Err(Error::Server(_))));
2289
2290 let mut result = QUERY_MIDDLE.run(&mut conn).await.unwrap();
2291
2292 let result_set: Vec<u8> = result.collect().await.unwrap();
2294 assert_eq!(result_set, vec![1]);
2295
2296 let result_set: super::Result<Vec<u8>> = result.collect().await;
2298 assert!(matches!(result_set, Err(Error::Server(_))));
2299
2300 assert!(result.is_empty());
2302
2303 conn.ping().await?;
2304 conn.disconnect().await?;
2305
2306 Ok(())
2307 }
2308
2309 #[tokio::test]
2310 async fn should_handle_binary_multiresult_set_with_error() -> super::Result<()> {
2311 const PROC_DEF_FIRST: &str =
2312 r#"CREATE PROCEDURE err_first() BEGIN SELECT * FROM tmp; SELECT 1; END"#;
2313 const PROC_DEF_MIDDLE: &str =
2314 r#"CREATE PROCEDURE err_middle() BEGIN SELECT 1; SELECT * FROM tmp; SELECT 2; END"#;
2315
2316 let mut conn = Conn::new(get_opts()).await.unwrap();
2317
2318 conn.query_drop("DROP PROCEDURE IF EXISTS err_first")
2319 .await?;
2320 conn.query_iter(PROC_DEF_FIRST).await?;
2321
2322 conn.query_drop("DROP PROCEDURE IF EXISTS err_middle")
2323 .await?;
2324 conn.query_iter(PROC_DEF_MIDDLE).await?;
2325
2326 let result = conn.query_iter("CALL err_first()").await;
2328 assert!(matches!(result, Err(Error::Server(_))));
2329
2330 let mut result = conn.query_iter("CALL err_middle()").await?;
2331
2332 let result_set: Vec<u8> = result.collect().await.unwrap();
2334 assert_eq!(result_set, vec![1]);
2335
2336 let result_set: super::Result<Vec<u8>> = result.collect().await;
2338 assert!(matches!(result_set, Err(Error::Server(_))));
2339
2340 assert!(result.is_empty());
2342
2343 conn.ping().await?;
2344 conn.disconnect().await?;
2345
2346 Ok(())
2347 }
2348
2349 #[tokio::test]
2350 async fn should_handle_multiresult_set_with_local_infile() -> super::Result<()> {
2351 use std::fs::write;
2352
2353 let file_path = tempfile::Builder::new().tempfile_in("").unwrap();
2354 let file_path = file_path.path();
2355 let file_name = file_path.file_name().unwrap();
2356
2357 write(file_name, b"AAAAAA\nBBBBBB\nCCCCCC\n")?;
2358
2359 let opts = get_opts().local_infile_handler(Some(WhiteListFsHandler::new(&[file_name][..])));
2360
2361 let mut conn = Conn::new(opts).await.unwrap();
2363 "CREATE TEMPORARY TABLE tmp (a TEXT)".run(&mut conn).await?;
2364
2365 let query = format!(
2366 r#"SELECT * FROM tmp;
2367 LOAD DATA LOCAL INFILE "{}" INTO TABLE tmp;
2368 LOAD DATA LOCAL INFILE "{}" INTO TABLE tmp;
2369 SELECT * FROM tmp"#,
2370 file_name.to_str().unwrap(),
2371 file_name.to_str().unwrap(),
2372 );
2373
2374 let mut result = query.run(&mut conn).await?;
2375
2376 let result_set = result.collect::<String>().await?;
2377 assert_eq!(result_set.len(), 0);
2378
2379 let mut no_local_infile = false;
2380
2381 for _ in 0..2 {
2382 match result.collect::<String>().await {
2383 Ok(result_set) => {
2384 assert_eq!(result.affected_rows(), 3);
2385 assert!(result_set.is_empty())
2386 }
2387 Err(Error::Server(ref err)) if err.code == 1148 => {
2388 no_local_infile = true;
2390 break;
2391 }
2392 Err(Error::Server(ref err)) if err.code == 3948 => {
2393 no_local_infile = true;
2396 break;
2397 }
2398 Err(err) => return Err(err),
2399 }
2400 }
2401
2402 if no_local_infile {
2403 assert!(result.is_empty());
2404 assert_eq!(result_set.len(), 0);
2405 } else {
2406 let result_set = result.collect::<String>().await?;
2407 assert_eq!(result_set.len(), 6);
2408 assert_eq!(result_set[0], "AAAAAA");
2409 assert_eq!(result_set[1], "BBBBBB");
2410 assert_eq!(result_set[2], "CCCCCC");
2411 assert_eq!(result_set[3], "AAAAAA");
2412 assert_eq!(result_set[4], "BBBBBB");
2413 assert_eq!(result_set[5], "CCCCCC");
2414 }
2415
2416 conn.ping().await?;
2417 conn.disconnect().await?;
2418
2419 Ok(())
2420 }
2421
2422 #[tokio::test]
2423 async fn should_provide_multiresult_set_metadata() -> super::Result<()> {
2424 let mut c = Conn::new(get_opts()).await?;
2425 c.query_drop("CREATE TEMPORARY TABLE tmp (id INT, foo TEXT)")
2426 .await?;
2427
2428 let mut result = c
2429 .query_iter("SELECT 1; SELECT id, foo FROM tmp WHERE 1 = 2; DO 42; SELECT 2;")
2430 .await?;
2431 assert_eq!(result.columns().map(|x| x.len()).unwrap_or_default(), 1);
2432
2433 result.for_each(drop).await?;
2434 assert_eq!(result.columns().map(|x| x.len()).unwrap_or_default(), 2);
2435
2436 result.for_each(drop).await?;
2437 assert_eq!(result.columns().map(|x| x.len()).unwrap_or_default(), 0);
2438
2439 result.for_each(drop).await?;
2440 assert_eq!(result.columns().map(|x| x.len()).unwrap_or_default(), 1);
2441
2442 c.disconnect().await?;
2443 Ok(())
2444 }
2445
2446 #[tokio::test]
2447 async fn should_expose_query_result_metadata() -> super::Result<()> {
2448 let pool = Pool::new(get_opts());
2449 let mut c = pool.get_conn().await?;
2450
2451 c.query_drop(
2452 r"
2453 CREATE TEMPORARY TABLE `foo`
2454 ( `id` SERIAL
2455 , `bar_id` varchar(36) NOT NULL
2456 , `baz_id` varchar(36) NOT NULL
2457 , `ctime` timestamp NOT NULL DEFAULT CURRENT_TIMESTAMP()
2458 , PRIMARY KEY (`id`)
2459 , KEY `bar_idx` (`bar_id`)
2460 , KEY `baz_idx` (`baz_id`)
2461 );",
2462 )
2463 .await?;
2464
2465 const QUERY: &str = "INSERT INTO foo (bar_id, baz_id) VALUES (?, ?)";
2466 let params = ("qwerty", "data.employee_id");
2467
2468 let query_result = c.exec_iter(QUERY, params).await?;
2469 assert_eq!(query_result.last_insert_id(), Some(1));
2470 query_result.drop_result().await?;
2471
2472 c.exec_drop(QUERY, params).await?;
2473 assert_eq!(c.last_insert_id(), Some(2));
2474
2475 let mut tx = c.start_transaction(Default::default()).await?;
2476
2477 tx.exec_drop(QUERY, params).await?;
2478 assert_eq!(tx.last_insert_id(), Some(3));
2479
2480 Ok(())
2481 }
2482
2483 #[tokio::test]
2484 async fn should_handle_local_infile_locally() -> super::Result<()> {
2485 let mut conn = Conn::new(get_opts()).await.unwrap();
2486 conn.query_drop("CREATE TEMPORARY TABLE tmp (a TEXT);")
2487 .await
2488 .unwrap();
2489
2490 conn.set_infile_handler(async move {
2491 Ok(
2492 stream::iter([Bytes::from("AAAAAA\n"), Bytes::from("BBBBBB\nCCCCCC\n")])
2493 .map(Ok)
2494 .boxed(),
2495 )
2496 });
2497
2498 match conn
2499 .query_drop(r#"LOAD DATA LOCAL INFILE "dummy" INTO TABLE tmp;"#)
2500 .await
2501 {
2502 Ok(_) => (),
2503 Err(super::Error::Server(ref err)) if err.code == 1148 => {
2504 return Ok(());
2506 }
2507 Err(super::Error::Server(ref err)) if err.code == 3948 => {
2508 return Ok(());
2511 }
2512 e @ Err(_) => e.unwrap(),
2513 };
2514
2515 let result: Vec<String> = conn.query("SELECT * FROM tmp").await?;
2516 assert_eq!(result.len(), 3);
2517 assert_eq!(result[0], "AAAAAA");
2518 assert_eq!(result[1], "BBBBBB");
2519 assert_eq!(result[2], "CCCCCC");
2520
2521 Ok(())
2522 }
2523
2524 #[tokio::test]
2525 async fn should_handle_local_infile_globally() -> super::Result<()> {
2526 use std::fs::write;
2527
2528 let file_path = tempfile::Builder::new().tempfile_in("").unwrap();
2529 let file_path = file_path.path();
2530 let file_name = file_path.file_name().unwrap();
2531
2532 write(file_name, b"AAAAAA\nBBBBBB\nCCCCCC\n")?;
2533
2534 let opts = get_opts().local_infile_handler(Some(WhiteListFsHandler::new(&[file_name][..])));
2535
2536 let mut conn = Conn::new(opts).await.unwrap();
2537 conn.query_drop("CREATE TEMPORARY TABLE tmp (a TEXT);")
2538 .await
2539 .unwrap();
2540
2541 match conn
2542 .query_drop(format!(
2543 r#"LOAD DATA LOCAL INFILE "{}" INTO TABLE tmp;"#,
2544 file_name.to_str().unwrap(),
2545 ))
2546 .await
2547 {
2548 Ok(_) => (),
2549 Err(super::Error::Server(ref err)) if err.code == 1148 => {
2550 return Ok(());
2552 }
2553 Err(super::Error::Server(ref err)) if err.code == 3948 => {
2554 return Ok(());
2557 }
2558 e @ Err(_) => e.unwrap(),
2559 };
2560
2561 let result: Vec<String> = conn.query("SELECT * FROM tmp").await?;
2562 assert_eq!(result.len(), 3);
2563 assert_eq!(result[0], "AAAAAA");
2564 assert_eq!(result[1], "BBBBBB");
2565 assert_eq!(result[2], "CCCCCC");
2566
2567 Ok(())
2568 }
2569
2570 #[tokio::test]
2571 async fn should_handle_initial_error_packet() {
2572 let header = [
2573 0x68, 0x00, 0x00, 0x00, 0xff, 0x69, 0x04, ];
2578 let error_message = "Host '172.17.0.1' is blocked because of many connection errors; unblock with 'mysqladmin flush-hosts'";
2579
2580 let listener = TcpListener::bind("127.0.0.1:0000").await.unwrap();
2582
2583 let listen_addr = listener.local_addr().unwrap();
2584
2585 tokio::task::spawn(async move {
2586 let (mut stream, _) = listener.accept().await.unwrap();
2587 stream.write_all(&header).await.unwrap();
2588 stream.write_all(error_message.as_bytes()).await.unwrap();
2589 stream.shutdown().await.unwrap();
2590 });
2591
2592 let opts = OptsBuilder::default()
2593 .ip_or_hostname(listen_addr.ip().to_string())
2594 .tcp_port(listen_addr.port());
2595 let server_err = match Conn::new(opts).await {
2596 Err(Error::Server(server_err)) => server_err,
2597 other => panic!("expected server error but got: {:?}", other),
2598 };
2599 assert_eq!(
2600 server_err,
2601 ServerError {
2602 code: 1129,
2603 state: "HY000".to_owned(),
2604 message: error_message.to_owned(),
2605 }
2606 );
2607 }
2608
2609 #[cfg_attr(docsrs, doc(cfg(feature = "client_parsec")))]
2610 #[tokio::test]
2611 #[cfg(feature = "client_parsec")]
2612 async fn parsec_connect() {
2613 let mut conn = Conn::new(get_opts()).await.unwrap();
2614 let is_mariadb = conn.inner.is_mariadb;
2615 let version = conn.server_version();
2616 if is_mariadb && version >= (11, 4, 1) {
2617 let mut rng = rand::rng();
2620 let mut pass_bytes = [0u8; 16];
2621 rng.fill(&mut pass_bytes);
2622 pass_bytes.iter_mut().for_each(|b| {
2623 *b = match *b % 3 {
2624 0 => b'A' + (*b % 26),
2625 1 => b'a' + (*b % 26),
2626 _ => b'0' + (*b % 10),
2627 }
2628 });
2629 let pass = String::from_utf8_lossy(&pass_bytes).to_string();
2630
2631 conn.query_drop("DROP USER IF EXISTS 'parsec_test_user'@'%'")
2632 .await
2633 .unwrap();
2634 let create_user_query = format!(
2635 "CREATE USER 'parsec_test_user'@'%' IDENTIFIED VIA 'parsec' USING PASSWORD('{}')",
2636 pass
2637 );
2638 conn.query_drop(create_user_query).await.unwrap();
2639 let mut conn_parsec = Conn::new(
2640 get_opts()
2641 .user(Some("parsec_test_user"))
2642 .pass(Some(pass))
2643 .db_name(None::<String>)
2644 .init(vec![] as Vec<String>),
2645 )
2646 .await
2647 .unwrap();
2648 assert!(conn_parsec.ping().await.is_ok());
2649 conn.query_drop("DROP USER 'parsec_test_user'@'%'")
2650 .await
2651 .unwrap();
2652 }
2653 }
2654
2655 #[tokio::test]
2659 async fn test_metadata_caching() {
2660 use crate::consts::ColumnType;
2661 let mut conn = Conn::new(get_opts()).await.unwrap();
2662 if !conn.inner.is_mariadb {
2663 return;
2664 }
2665
2666 conn.query_drop(
2667 r"CREATE TEMPORARY TABLE t_metadata_caching (
2668 id INT NOT NULL PRIMARY KEY AUTO_INCREMENT,
2669 val VARCHAR(32) NOT NULL)",
2670 )
2671 .await
2672 .unwrap();
2673
2674 let insert_stmt = conn
2676 .prep("INSERT INTO t_metadata_caching (val) VALUES (?)")
2677 .await
2678 .unwrap();
2679 let _ = conn.exec_drop(&insert_stmt, ("AAA",)).await;
2680 let _ = conn.exec_drop(&insert_stmt, ("BB",)).await;
2681 let mut ps = conn
2682 .prep("SELECT id, val FROM t_metadata_caching")
2683 .await
2684 .unwrap();
2685
2686 let mut columns_from_prep = ps.columns();
2687 let mut metadata_from_prep: Vec<(String, ColumnType)> = columns_from_prep
2688 .iter()
2689 .map(|column| (column.name_str().to_string(), column.column_type()))
2690 .collect();
2691
2692 let mut query_result = conn.exec_iter(&ps, ()).await.unwrap();
2693 let mut columns_from_exec1 = query_result.columns().unwrap();
2694 let mut metadata_from_exec1: Vec<(String, ColumnType)> = columns_from_exec1
2695 .iter()
2696 .map(|column| (column.name_str().to_string(), column.column_type()))
2697 .collect();
2698
2699 assert_eq!(metadata_from_prep, metadata_from_exec1);
2701 assert_eq!(metadata_from_prep.len(), 2);
2702 assert_eq!(metadata_from_prep[0].0, "id");
2703 assert_eq!(metadata_from_prep[0].1, ColumnType::MYSQL_TYPE_LONG);
2704 assert_eq!(metadata_from_prep[1].0, "val");
2705 assert_eq!(metadata_from_prep[1].1, ColumnType::MYSQL_TYPE_VAR_STRING);
2706
2707 let fetched_rows: Vec<(i32, String)> = query_result.collect().await.unwrap();
2708
2709 let expected_rows = [(1, "AAA".to_string()), (2, "BB".to_string())];
2710 assert_eq!(fetched_rows.len(), expected_rows.len());
2711
2712 for (fetched, expected) in fetched_rows.iter().zip(expected_rows.iter()) {
2713 assert_eq!(fetched, expected);
2714 }
2715
2716 ps = conn
2719 .prep("SELECT val FROM t_metadata_caching WHERE id = ?")
2720 .await
2721 .unwrap();
2722 columns_from_prep = ps.columns();
2723 metadata_from_prep = columns_from_prep
2724 .iter()
2725 .map(|column| (column.name_str().to_string(), column.column_type()))
2726 .collect();
2727 let single_row: Option<String> = conn.exec_first(&ps, (1,)).await.unwrap();
2728 if let Some(val) = single_row {
2729 assert_eq!(metadata_from_prep.len(), 1);
2730 assert_eq!(metadata_from_prep[0].0, "val");
2731 assert_eq!(metadata_from_prep[0].1, ColumnType::MYSQL_TYPE_VAR_STRING);
2732
2733 assert_eq!(val, "AAA".to_string());
2734 }
2735 ps = conn.prep("SELECT ?").await.unwrap();
2737
2738 columns_from_prep = ps.columns();
2739 metadata_from_prep = columns_from_prep
2740 .iter()
2741 .map(|column| (column.name_str().to_string(), column.column_type()))
2742 .collect();
2743
2744 query_result = conn.exec_iter(&ps, (12,)).await.unwrap();
2746 columns_from_exec1 = query_result.columns().unwrap();
2747 metadata_from_exec1 = columns_from_exec1
2748 .iter()
2749 .map(|column| (column.name_str().to_string(), column.column_type()))
2750 .collect();
2751 let fetched_rows: Vec<i32> = query_result.collect().await.unwrap();
2752
2753 query_result = conn.exec_iter(&ps, (42,)).await.unwrap();
2755 let columns_from_exec2 = query_result.columns().unwrap();
2756 let metadata_from_exec2 = columns_from_exec2
2757 .iter()
2758 .map(|column| (column.name_str().to_string(), column.column_type()))
2759 .collect::<Vec<_>>();
2760 let fetched_rows2: Vec<i32> = query_result.collect().await.unwrap();
2761
2762 query_result = conn.exec_iter(&ps, ("foo",)).await.unwrap();
2764 let columns_from_exec3 = query_result.columns().unwrap();
2765 let metadata_from_exec3 = columns_from_exec3
2766 .iter()
2767 .map(|column| (column.name_str().to_string(), column.column_type()))
2768 .collect::<Vec<_>>();
2769 let fetched_rows3: Vec<String> = query_result.collect().await.unwrap();
2770
2771 assert_eq!(metadata_from_exec1.len(), 1);
2773 assert_eq!(metadata_from_exec2.len(), 1);
2774 assert_eq!(metadata_from_exec3.len(), 1);
2775 assert_eq!(metadata_from_prep.len(), 1);
2776 assert_eq!(metadata_from_prep[0].0, "?");
2777 assert!(
2778 metadata_from_prep[0].1 == ColumnType::MYSQL_TYPE_NULL
2779 || metadata_from_prep[0].1 == ColumnType::MYSQL_TYPE_VAR_STRING,
2780 "Expected MYSQL_TYPE_NULL(MariaDB) or MYSQL_TYPE_VAR_STRING(MySQL), got {:?}",
2781 metadata_from_prep[0].1
2782 );
2783 assert_eq!(metadata_from_exec1[0].0, "?");
2784 assert_eq!(metadata_from_exec1[0].1, ColumnType::MYSQL_TYPE_LONGLONG);
2785 assert_eq!(metadata_from_exec2[0].0, "?");
2786 assert_eq!(metadata_from_exec2[0].1, ColumnType::MYSQL_TYPE_LONGLONG);
2787 assert_eq!(metadata_from_exec3[0].0, "?");
2788 assert_eq!(metadata_from_exec3[0].1, ColumnType::MYSQL_TYPE_VAR_STRING);
2789
2790 assert_eq!(fetched_rows[0], 12);
2791 assert_eq!(fetched_rows2[0], 42);
2792 assert_eq!(fetched_rows3[0], "foo".to_owned());
2793 }
2794
2795 #[cfg(feature = "nightly")]
2796 mod bench {
2797 use crate::{conn::Conn, queryable::Queryable, test_misc::get_opts};
2798
2799 #[bench]
2800 fn simple_exec(bencher: &mut test::Bencher) {
2801 let mut runtime = tokio::runtime::Runtime::new().unwrap();
2802 let mut conn = runtime.block_on(Conn::new(get_opts())).unwrap();
2803
2804 bencher.iter(|| {
2805 runtime.block_on(conn.query_drop("DO 1")).unwrap();
2806 });
2807
2808 runtime.block_on(conn.disconnect()).unwrap();
2809 }
2810
2811 #[bench]
2812 fn select_large_string(bencher: &mut test::Bencher) {
2813 let mut runtime = tokio::runtime::Runtime::new().unwrap();
2814 let mut conn = runtime.block_on(Conn::new(get_opts())).unwrap();
2815
2816 bencher.iter(|| {
2817 runtime
2818 .block_on(conn.query_drop("SELECT REPEAT('A', 10000)"))
2819 .unwrap();
2820 });
2821
2822 runtime.block_on(conn.disconnect()).unwrap();
2823 }
2824
2825 #[bench]
2826 fn prepared_exec(bencher: &mut test::Bencher) {
2827 let mut runtime = tokio::runtime::Runtime::new().unwrap();
2828 let mut conn = runtime.block_on(Conn::new(get_opts())).unwrap();
2829 let stmt = runtime.block_on(conn.prep("DO 1")).unwrap();
2830
2831 bencher.iter(|| {
2832 runtime.block_on(conn.exec_drop(&stmt, ())).unwrap();
2833 });
2834
2835 runtime.block_on(conn.close(stmt)).unwrap();
2836 runtime.block_on(conn.disconnect()).unwrap();
2837 }
2838
2839 #[bench]
2840 fn prepare_and_exec(bencher: &mut test::Bencher) {
2841 let mut runtime = tokio::runtime::Runtime::new().unwrap();
2842 let mut conn = runtime.block_on(Conn::new(get_opts())).unwrap();
2843
2844 bencher.iter(|| {
2845 runtime.block_on(conn.exec_drop("SELECT ?", (0,))).unwrap();
2846 });
2847
2848 runtime.block_on(conn.disconnect()).unwrap();
2849 }
2850 }
2851}