Skip to main content

mysql_async/conn/
mod.rs

1// Copyright (c) 2016 Anatoly Ikorsky
2//
3// Licensed under the Apache License, Version 2.0
4// <LICENSE-APACHE or http://www.apache.org/licenses/LICENSE-2.0> or the MIT
5// license <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
6// option. All files in the project carrying such notice may not be copied,
7// modified, or distributed except according to those terms.
8
9use 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
62/// Helper that asynchronously disconnects the given connection on the default tokio executor.
63fn disconnect(mut conn: Conn) {
64    let disconnected = conn.inner.disconnected;
65
66    // Mark conn as disconnected.
67    conn.inner.disconnected = true;
68
69    if !disconnected {
70        // We shouldn't call tokio::spawn if unwinding
71        if std::thread::panicking() {
72            return;
73        }
74
75        // Server will report broken connection if spawn fails.
76        // this might fail if, say, the runtime is shutting down, but we've done what we could
77        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/// Pending result set.
88#[derive(Debug, Clone)]
89pub(crate) enum PendingResult {
90    /// There is a pending result set.
91    Pending(ResultSetMeta),
92    /// Result set metadata was taken but not yet consumed.
93    Taken(Arc<ResultSetMeta>),
94}
95
96/// Mysql connection
97struct 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    /// Connection is already disconnected.
124    pub(crate) disconnected: bool,
125    /// One-time connection-level infile handler.
126    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    /// Constructs an empty connection.
150    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    /// Returns mutable reference to a connection stream.
184    ///
185    /// Returns `DriverError::ConnectionClosed` if there is no stream.
186    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/// MySql server connection.
200#[derive(Debug)]
201pub struct Conn {
202    inner: Box<ConnInner>,
203}
204
205impl Conn {
206    /// Returns connection identifier.
207    pub fn id(&self) -> u32 {
208        self.inner.id
209    }
210
211    /// Returns the disconnected state of the connection.
212    pub fn is_disconnected(&self) -> bool {
213        self.inner.disconnected
214    }
215
216    /// Returns the ID generated by a query (usually `INSERT`) on a table with a column having the
217    /// `AUTO_INCREMENT` attribute. Returns `None` if there was no previous query on the connection
218    /// or if the query did not update an AUTO_INCREMENT value.
219    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    /// Returns the number of rows affected by the last `INSERT`, `UPDATE`, `REPLACE` or `DELETE`
227    /// query.
228    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    /// Text information, as reported by the server in the last OK packet, or an empty string.
237    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    /// Number of warnings, as reported by the server in the last OK packet, or `0`.
246    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    /// Returns a reference to the last OK packet.
255    pub fn last_ok_packet(&self) -> Option<&OkPacket<'static>> {
256        self.inner.last_ok_packet.as_ref()
257    }
258
259    /// Turns on/off automatic connection reset (see [`crate::PoolOpts::with_reset_connection`]).
260    ///
261    /// Only makes sense for pooled connections.
262    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    /// Will update last IO time for this connection.
287    pub(crate) fn touch(&mut self) {
288        self.inner.last_io = Instant::now();
289    }
290
291    /// Will set packet sequence id to `0`.
292    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    /// Will synchronize sequence ids between compressed and uncompressed codecs.
299    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    /// Handles OK packet.
306    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    /// Handles ERR packet.
313    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    /// Returns the current transaction status.
326    pub(crate) fn get_tx_status(&self) -> TxStatus {
327        self.inner.tx_status
328    }
329
330    /// Sets the given transaction status for this connection.
331    pub(crate) fn set_tx_status(&mut self, tx_status: TxStatus) {
332        self.inner.tx_status = tx_status;
333    }
334
335    /// Returns pending result metadata, if any.
336    ///
337    /// If `Some(_)`, then result is not yet consumed.
338    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    /// Sets the given pending result metadata for this connection. Returns the previous value.
361    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    /// Gives the currently pending result to a caller for consumption.
379    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    /// Returns current status flags.
397    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                // either OK or non-fatal error
410                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    /// Returns server version.
423    pub fn server_version(&self) -> (u16, u16, u16) {
424        self.inner.version
425    }
426
427    /// Returns connection options.
428    pub fn opts(&self) -> &Opts {
429        &self.inner.opts
430    }
431
432    /// Setup _local_ `LOCAL INFILE` handler (see ["LOCAL INFILE Handlers"][2] section
433    /// of the crate-level docs).
434    ///
435    /// It'll overwrite existing _local_ handler, if any.
436    ///
437    /// [2]: ../mysql_async/#local-infile-handlers
438    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    /// Disconnects this connection from server.
451    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    /// Closes the connection.
462    async fn close_conn(mut self) -> Result<()> {
463        self = self.cleanup_for_pool().await?;
464        self.disconnect().await
465    }
466
467    /// Returns true if io stream is encrypted.
468    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    /// Returns true if io stream is socket.
483    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    /// Hacky way to move connection through &mut. `self` becomes unusable.
498    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    /// Set `io::Stream` options as defined in the `Opts` of the connection.
509    ///
510    /// Requires that self.inner.stream is Some
511    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        // Handshake scramble is always 21 bytes length (20 + zero terminator)
524        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            // Trim zero terminator. Fill with zeroes if nonce
528            // is somehow smaller than 20 bytes (this matches the server behavior).
529            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 we have a MariaDB server version, we are using mariadb extended capabilities from the handshake packet.
543        // MariaDB does not set the first standard capability flag bit to indicate that it supports extended capabilities.
544        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        // Allow only CachingSha2Password and MysqlNativePassword here
550        // because sha256_password is deprecated and other plugins won't
551        // appear here.
552        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        // Serialize here to satisfy borrow checker.
623        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                // For parsec at this point we need to send an empty packet first
672                _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        // NOTE: we need to box this since it may recurse
692        // see https://github.com/rust-lang/rust/issues/46415#issuecomment-528099782
693        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 packet for empty password
743                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        // Normally we need to skip escaping 0x01 byte. But in first parsec implementations, server did not send it.
759        let mut payload: &[u8] = &packet;
760        if packet.first() == Some(&0x01) {
761            payload = &packet[1..];
762        }
763        // At this point in future, when it will be possible for parsec to be default authentication method,
764        // we can have authentication switch request. The other possible option here(and for now the only option) -
765        // ext-salt packet.
766        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            // Letting parser function decide if all is fine with the packet
771            self.inner
772                .auth_plugin
773                .read_add_data(payload)
774                .ok_or_else(|| DriverError::InvalidParsecSalt)?;
775            // Now generating response.
776            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            // After client response, server will send either ok or error.
784            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 packet for empty password
800                Ok(())
801            }
802            Some(0x01) => match packet.get(1) {
803                Some(0x03) => {
804                    // auth ok
805                    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                    // map OldAuthSwitch to AuthSwitch with mysql_old_password plugin
859                    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    /// Returns `true` for ProgressReport packet.
874    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            // If we haven't completed the handshake the server will not be aware of our
895            // capabilities and so it will behave as if we have none. In particular, the error
896            // packet will not contain a SQL State field even if our capabilities do contain the
897            // `CLIENT_PROTOCOL_41` flag. Therefore it is necessary to parse an incoming packet
898            // with no capability assumptions if we have not completed the handshake.
899            //
900            // https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_connection_phase.html
901            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                // ignore progress report
927                continue;
928            } else {
929                return Ok(packet);
930            }
931        }
932    }
933
934    /// Returns future that reads packets from a server.
935    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    /// Writes bytes to a server.
954    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    /// Sends a serializable structure to a server.
960    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    /// Sends a command to a server.
967    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    /// Returns future that sends full command body to a server.
974    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    /// Returns future that writes command to a server.
982    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    /// Returns a future that resolves to [`Conn`].
1024    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    /// Returns a future that resolves to [`Conn`].
1059    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    /// Will try to reconnect via socket using socket address in `self.inner.socket`.
1064    ///
1065    /// Won't try to reconnect if socket connection is already enforced in [`Opts`].
1066    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                    // tidy up the old connection
1074                    old_conn.close_conn().await?;
1075                }
1076            }
1077        }
1078        Ok(())
1079    }
1080
1081    /// Configures the connection based on server settings. In particular:
1082    ///
1083    /// * It reads and stores socket address inside the connection unless if socket address is
1084    ///   already in [`Opts`] or if `prefer_socket` is `false`.
1085    ///
1086    /// * It reads and stores `max_allowed_packet` in the connection unless it's already in [`Opts`]
1087    ///
1088    /// * It reads and stores `wait_timeout` in the connection unless it's already in [`Opts`]
1089    ///
1090    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    /// Returns true if time since last IO exceeds `wait_timeout`
1211    /// (or `conn_ttl` if specified in opts).
1212    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    /// Returns duration since last IO.
1227    fn idling(&self) -> Duration {
1228        self.inner.last_io.elapsed()
1229    }
1230
1231    /// Executes [`COM_RESET_CONNECTION`][1].
1232    ///
1233    /// Returns `false` if command is not supported (requires MySql >5.7.2, MariaDb >10.2.3).
1234    /// For older versions consider using [`Conn::change_user`].
1235    ///
1236    /// [1]: https://dev.mysql.com/doc/c-api/5.7/en/mysql-reset-connection.html
1237    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            // assuming mysql
1242            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    /// Executes [`COM_CHANGE_USER`][1].
1256    ///
1257    /// This might be used as an older and slower alternative to `COM_RESET_CONNECTION` that
1258    /// works on MySql prior to 5.7.3 (MariaDb prior ot 10.2.4).
1259    ///
1260    /// ## Note
1261    ///
1262    /// * Using non-default `opts` for a pooled connection is discouraging.
1263    /// * Connection options will be permanently updated.
1264    ///
1265    /// [1]: https://dev.mysql.com/doc/c-api/5.7/en/mysql-change-user.html
1266    pub async fn change_user(&mut self, opts: ChangeUserOpts) -> Result<()> {
1267        // We'll kick this connection from a pool if opts are changed.
1268        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    /// Resets the connection upon returning it to a pool.
1296    ///
1297    /// Will invoke `COM_CHANGE_USER` if `COM_RESET_CONNECTION` is not supported.
1298    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    /// Requires that `self.inner.tx_status != TxStatus::None`
1306    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            // Reset tx_status if ROLLBACK fails so that the connection is considered dirty
1311            self.inner.tx_status = tx_status;
1312            return Err(e);
1313        }
1314        Ok(())
1315    }
1316
1317    /// Returns `true` if `SERVER_MORE_RESULTS_EXISTS` flag is contained
1318    /// in status flags of the connection.
1319    pub(crate) fn more_results_exists(&self) -> bool {
1320        self.status()
1321            .contains(StatusFlags::SERVER_MORE_RESULTS_EXISTS)
1322    }
1323
1324    /// The purpose of this function is to cleanup a pending result set
1325    /// for prematurely dropped connection or query result.
1326    ///
1327    /// Requires that there are no other references to the pending result.
1328    pub(crate) async fn drop_result(&mut self) -> Result<()> {
1329        // Map everything into `PendingResult::Pending`
1330        let meta = match self.set_pending_result(None)? {
1331            Some(PendingResult::Pending(meta)) => Some(meta),
1332            Some(PendingResult::Taken(meta)) => {
1333                // This also asserts that there is only one reference left to the taken ResultSetMeta,
1334                // therefore this result set must be dropped here since it won't be dropped anywhere else.
1335                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((/* this case does not require an action */)),
1354            Ok(Some(PendingResult::Taken(_))) | Err(_) => {
1355                unreachable!("this case must be handled earlier in this function")
1356            }
1357        }
1358    }
1359
1360    /// This function will drop pending result and rollback a transaction, if needed.
1361    ///
1362    /// The purpose of this function, is to cleanup the connection while returning it to a [`Pool`].
1363    async fn cleanup_for_pool(mut self) -> Result<Self> {
1364        loop {
1365            if self.has_pending_result() {
1366                // The connection was dropped and we assume that it was dropped intentionally,
1367                // so we'll ignore non-fatal errors during cleanup (also there is no direct caller
1368                // to return this error to).
1369                if let Err(err) = self.drop_result().await {
1370                    if err.is_fatal() {
1371                        // This means that connection is completely broken
1372                        // and shouldn't return to a pool.
1373                        return Err(err);
1374                    }
1375                }
1376            } else if self.inner.tx_status != TxStatus::None {
1377                // If an error occurs during rollback, don't reuse the connection.
1378                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        // Inserted one row, affected should be one.
1417        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        // The query doesn't affect any rows, but due to us wanting FOUND rows,
1424        // this has to return one.
1425        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        // Inserted one row, affected should be one.
1443        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        // The query doesn't affect any rows.
1450        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        // no database name
1465        let mut conn: Conn = Conn::new(get_opts().db_name(None::<String>)).await?;
1466        conn.ping().await?;
1467        conn.disconnect().await?;
1468
1469        // empty database name
1470        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        // dropped query:
1485        conn.query_iter("SELECT 1").await?;
1486        conn.ping().await?;
1487
1488        // dropped query in dropped transaction:
1489        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        // Should connect with any combination of supported plugin and empty-nonempty password.
1517        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        // connection will drop here
1585    }
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        // initial run
1632        let mut result: Vec<(u8, String)> = conn.query("SELECT @a, @b").await?;
1633        assert_eq!(result, vec![(42, "foo".into())]);
1634
1635        // after reset
1636        if conn.reset().await? {
1637            result = conn.query("SELECT @a, @b").await?;
1638            assert_eq!(result, vec![(42, "foo".into())]);
1639        }
1640
1641        // after change user
1642        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        /// Whether particular authentication plugin should be tested on the current database.
1685        type ShouldRunFn = fn(bool, (u16, u16, u16)) -> bool;
1686        /// Generates `CREATE USER` and `SET PASSWORD` statements
1687        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                    // IF EXISTS is not supported on 5.6 so the query will fail on the first iteration
1810                    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    // Test for exec_batch method.
1846    #[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        // Populating table with some data to verify that data still fetched correctly with cached metadata use
1860        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; // 1K
1904        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        // Calculate a row size that will make the total packet size > max_allowed_packet
1914        // Packet will have 4 byte header and 7 bytes COM_STMT_BULK_EXECUTE fields + 4 bytes for parameter types
1915        // 8 bytes per row for id + 2 bytes for indicators + 3 bytes for length encoding of val.
1916        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        // MySql compression does not respect max_allowed_packet, so we're going to query
1941        // one by one
1942        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        // Some corner cases: single row exceeding max_allowed_packet and
1972        // empty batch
1973        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); // implicit rollback
2273        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        // if error is in the first result set, then query should return it immediately.
2287        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        // first result set will contain one row
2293        let result_set: Vec<u8> = result.collect().await.unwrap();
2294        assert_eq!(result_set, vec![1]);
2295
2296        // second result set will contain an error.
2297        let result_set: super::Result<Vec<u8>> = result.collect().await;
2298        assert!(matches!(result_set, Err(Error::Server(_))));
2299
2300        // there will be no third result set
2301        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        // if error is in the first result set, then query should return it immediately.
2327        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        // first result set will contain one row
2333        let result_set: Vec<u8> = result.collect().await.unwrap();
2334        assert_eq!(result_set, vec![1]);
2335
2336        // second result set will contain an error.
2337        let result_set: super::Result<Vec<u8>> = result.collect().await;
2338        assert!(matches!(result_set, Err(Error::Server(_))));
2339
2340        // there will be no third result set
2341        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        // LOCAL INFILE in the middle of a multi-result set should not break anything.
2362        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                    // The used command is not allowed with this MySQL version
2389                    no_local_infile = true;
2390                    break;
2391                }
2392                Err(Error::Server(ref err)) if err.code == 3948 => {
2393                    // Loading local data is disabled;
2394                    // this must be enabled on both the client and server sides
2395                    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                // The used command is not allowed with this MySQL version
2505                return Ok(());
2506            }
2507            Err(super::Error::Server(ref err)) if err.code == 3948 => {
2508                // Loading local data is disabled;
2509                // this must be enabled on both the client and server sides
2510                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                // The used command is not allowed with this MySQL version
2551                return Ok(());
2552            }
2553            Err(super::Error::Server(ref err)) if err.code == 3948 => {
2554                // Loading local data is disabled;
2555                // this must be enabled on both the client and server sides
2556                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, // packet_length
2574            0x00, // sequence
2575            0xff, // error_header
2576            0x69, 0x04, // error_code
2577        ];
2578        let error_message = "Host '172.17.0.1' is blocked because of many connection errors; unblock with 'mysqladmin flush-hosts'";
2579
2580        // Create a fake MySQL server that immediately replies with an error packet.
2581        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            // Creating random password so in case of test failure user won't have
2618            // known password left behind.
2619            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    // This test verifies that the metadata is correct with or without metadata caching, and that protocol
2656    // is not broken afterwards and data is read correctly. It doesn't test that the metadata is really cached
2657    // (if that is possible) and not received twice.
2658    #[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        // Populating table with some data to verify that data still fetched correctly with cached metadata use
2675        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        // Comparing and verifying metadata.
2700        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        // Doing the same for exec_first. Technically it's internally the same as exec_iter,
2717        // but the test isn't supposed to know that and to test it
2718        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        // Testing the case when metadata is changed after execution
2736        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        // First query — server sends metadata because the type has changed
2745        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        // Second query — server skips metadata packets
2754        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        // Third query — server sends metadata because the type has changed
2763        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        // Comparing and verifying metadata.
2772        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}