Skip to main content

mysql_common/packets/
mod.rs

1// Copyright (c) 2017 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 btoi::btoi;
10use bytes::BufMut;
11use regex::bytes::Regex;
12use uuid::Uuid;
13
14use std::{
15    borrow::Cow,
16    cmp::max,
17    collections::HashMap,
18    convert::TryFrom,
19    fmt, io,
20    marker::PhantomData,
21    mem,
22    str::FromStr,
23    sync::{Arc, LazyLock},
24};
25
26use crate::{
27    collations::CollationId,
28    constants::{
29        CapabilityFlags, ColumnFlags, ColumnType, Command, CursorType, MAX_PAYLOAD_LEN,
30        MariadbBulkIndicator, MariadbCapabilities, SessionStateType, StatusFlags,
31        StmtBulkExecuteFlags, StmtExecuteParamFlags, StmtExecuteParamsFlags,
32    },
33    io::{BufMutExt, ParseBuf},
34    misc::{
35        lenenc_str_len,
36        raw::{
37            Const, Either, RawBytes, RawConst, RawInt, Skip,
38            bytes::{
39                BareBytes, ConstBytes, ConstBytesValue, EofBytes, LenEnc, NullBytes, U8Bytes,
40                U32Bytes,
41            },
42            int::{ConstU8, ConstU32, LeU16, LeU24, LeU32, LeU32LowerHalf, LeU32UpperHalf, LeU64},
43            seq::{Seq, Unknown},
44        },
45        read_varlen_uint, unexpected_buf_eof, varlen_uint_size, write_varlen_uint,
46    },
47    params::{Params, ParamsError},
48    proto::{MyDeserialize, MySerialize},
49    scramble::create_response_for_ed25519,
50    scramble::create_response_for_parsec,
51    value::{BinValue, ClientSide, SerializationSide, Value, ValueDeserializer},
52};
53
54use self::session_state_change::SessionStateChange;
55
56static MARIADB_VERSION_RE: LazyLock<Regex> =
57    LazyLock::new(|| Regex::new(r"^(?:5.5.5-)?(\d{1,2})\.(\d{1,2})\.(\d{1,3})-MariaDB").unwrap());
58static VERSION_RE: LazyLock<Regex> =
59    LazyLock::new(|| Regex::new(r"^(\d{1,2})\.(\d{1,2})\.(\d{1,3})(.*)").unwrap());
60
61macro_rules! define_header {
62    ($name:ident, $err:ident($msg:literal), $val:literal) => {
63        #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
64        #[error($msg)]
65        pub struct $err;
66        pub type $name = crate::misc::raw::int::ConstU8<$err, $val>;
67    };
68    ($name:ident, $cmd:ident, $err:ident) => {
69        #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
70        #[error("Invalid header for {}", stringify!($cmd))]
71        pub struct $err;
72        pub type $name = crate::misc::raw::int::ConstU8<$err, { Command::$cmd as u8 }>;
73    };
74}
75
76macro_rules! define_const {
77    ($kind:ident, $name:ident, $err:ident($msg:literal), $val:literal) => {
78        #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
79        #[error($msg)]
80        pub struct $err;
81        pub type $name = $kind<$err, $val>;
82    };
83}
84
85macro_rules! define_const_bytes {
86    ($v_name:ident, $name:ident, $err:ident($msg:literal), $val:expr_2021, $len:literal) => {
87        #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
88        #[error($msg)]
89        pub struct $err;
90
91        #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash)]
92        pub struct $v_name;
93
94        impl ConstBytesValue<$len> for $v_name {
95            const VALUE: [u8; $len] = $val;
96            type Error = $err;
97        }
98
99        pub type $name = ConstBytes<$v_name, $len>;
100    };
101}
102
103pub mod binlog_request;
104pub mod caching_sha2_password;
105pub mod session_state_change;
106
107define_const_bytes!(
108    Catalog,
109    ColumnDefinitionCatalog,
110    InvalidCatalog("Invalid catalog value in the column definition"),
111    *b"\x03def",
112    4
113);
114
115define_const!(
116    ConstU8,
117    FixedLengthFieldsLen,
118    InvalidFixedLengthFieldsLen("Invalid fixed length field length in the column definition"),
119    0x0c
120);
121
122/// Dynamically-sized column metadata — a part of the [`Column`] packet.
123#[derive(Debug, Default, Clone, Eq, PartialEq)]
124struct ColumnMeta<'a> {
125    schema: RawBytes<'a, LenEnc>,
126    table: RawBytes<'a, LenEnc>,
127    org_table: RawBytes<'a, LenEnc>,
128    name: RawBytes<'a, LenEnc>,
129    org_name: RawBytes<'a, LenEnc>,
130}
131
132impl ColumnMeta<'_> {
133    pub fn into_owned(self) -> ColumnMeta<'static> {
134        ColumnMeta {
135            schema: self.schema.into_owned(),
136            table: self.table.into_owned(),
137            org_table: self.org_table.into_owned(),
138            name: self.name.into_owned(),
139            org_name: self.org_name.into_owned(),
140        }
141    }
142
143    /// Returns the value of the [`ColumnMeta::schema`] field as a byte slice.
144    pub fn schema_ref(&self) -> &[u8] {
145        self.schema.as_bytes()
146    }
147
148    /// Returns the value of the [`ColumnMeta::schema`] field as a string (lossy converted).
149    pub fn schema_str(&self) -> Cow<'_, str> {
150        String::from_utf8_lossy(self.schema_ref())
151    }
152
153    /// Returns the value of the [`ColumnMeta::table`] field as a byte slice.
154    pub fn table_ref(&self) -> &[u8] {
155        self.table.as_bytes()
156    }
157
158    /// Returns the value of the [`ColumnMeta::table`] field as a string (lossy converted).
159    pub fn table_str(&self) -> Cow<'_, str> {
160        String::from_utf8_lossy(self.table_ref())
161    }
162
163    /// Returns the value of the [`ColumnMeta::org_table`] field as a byte slice.
164    ///
165    /// "org_table" is for original table name.
166    pub fn org_table_ref(&self) -> &[u8] {
167        self.org_table.as_bytes()
168    }
169
170    /// Returns the value of the [`ColumnMeta::org_table`] field as a string (lossy converted).
171    pub fn org_table_str(&self) -> Cow<'_, str> {
172        String::from_utf8_lossy(self.org_table_ref())
173    }
174
175    /// Returns the value of the [`ColumnMeta::name`] field as a byte slice.
176    pub fn name_ref(&self) -> &[u8] {
177        self.name.as_bytes()
178    }
179
180    /// Returns the value of the [`ColumnMeta::name`] field as a string (lossy converted).
181    pub fn name_str(&self) -> Cow<'_, str> {
182        String::from_utf8_lossy(self.name_ref())
183    }
184
185    /// Returns the value of the [`ColumnMeta::org_name`] field as a byte slice.
186    ///
187    /// "org_name" is for original column name.
188    pub fn org_name_ref(&self) -> &[u8] {
189        self.org_name.as_bytes()
190    }
191
192    /// Returns value of the [`ColumnMeta::org_name`] field as a string (lossy converted).
193    pub fn org_name_str(&self) -> Cow<'_, str> {
194        String::from_utf8_lossy(self.org_name_ref())
195    }
196}
197
198impl<'de> MyDeserialize<'de> for ColumnMeta<'de> {
199    const SIZE: Option<usize> = None;
200    type Ctx = ();
201
202    fn deserialize(_ctx: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
203        Ok(Self {
204            schema: buf.parse_unchecked(())?,
205            table: buf.parse_unchecked(())?,
206            org_table: buf.parse_unchecked(())?,
207            name: buf.parse_unchecked(())?,
208            org_name: buf.parse_unchecked(())?,
209        })
210    }
211}
212
213impl MySerialize for ColumnMeta<'_> {
214    fn serialize(&self, buf: &mut Vec<u8>) {
215        self.schema.serialize(&mut *buf);
216        self.table.serialize(&mut *buf);
217        self.org_table.serialize(&mut *buf);
218        self.name.serialize(&mut *buf);
219        self.org_name.serialize(&mut *buf);
220    }
221}
222
223/// Represents MySql Column (column packet).
224#[derive(Debug, Clone, Eq, PartialEq)]
225pub struct Column {
226    catalog: ColumnDefinitionCatalog,
227    meta: Arc<ColumnMeta<'static>>,
228    fixed_length_fields_len: FixedLengthFieldsLen,
229    column_length: RawInt<LeU32>,
230    character_set: RawInt<LeU16>,
231    column_type: Const<ColumnType, u8>,
232    flags: Const<ColumnFlags, LeU16>,
233    decimals: RawInt<u8>,
234    __filler: Skip<2>,
235    // COM_FIELD_LIST is deprecated, so we won't support it
236}
237
238impl<'de> MyDeserialize<'de> for Column {
239    const SIZE: Option<usize> = None;
240    type Ctx = ();
241
242    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
243        let catalog = buf.parse(())?;
244        let meta = Arc::new(buf.parse::<ColumnMeta<'_>>(())?.into_owned());
245        let mut buf: ParseBuf<'_> = buf.parse(13)?;
246
247        Ok(Column {
248            catalog,
249            meta,
250            fixed_length_fields_len: buf.parse_unchecked(())?,
251            character_set: buf.parse_unchecked(())?,
252            column_length: buf.parse_unchecked(())?,
253            column_type: buf.parse_unchecked(())?,
254            flags: buf.parse_unchecked(())?,
255            decimals: buf.parse_unchecked(())?,
256            __filler: buf.parse_unchecked(())?,
257        })
258    }
259}
260
261impl MySerialize for Column {
262    fn serialize(&self, buf: &mut Vec<u8>) {
263        self.catalog.serialize(&mut *buf);
264        self.meta.serialize(&mut *buf);
265        self.fixed_length_fields_len.serialize(&mut *buf);
266        self.column_length.serialize(&mut *buf);
267        self.character_set.serialize(&mut *buf);
268        self.column_type.serialize(&mut *buf);
269        self.flags.serialize(&mut *buf);
270        self.decimals.serialize(&mut *buf);
271        self.__filler.serialize(&mut *buf);
272    }
273}
274
275impl Column {
276    pub fn new(column_type: ColumnType) -> Self {
277        Self {
278            catalog: Default::default(),
279            meta: Default::default(),
280            fixed_length_fields_len: Default::default(),
281            column_length: Default::default(),
282            character_set: Default::default(),
283            flags: Default::default(),
284            column_type: Const::new(column_type),
285            decimals: Default::default(),
286            __filler: Skip,
287        }
288    }
289
290    pub fn with_schema(mut self, schema: &[u8]) -> Self {
291        Arc::make_mut(&mut self.meta).schema = RawBytes::new(schema).into_owned();
292        self
293    }
294
295    pub fn with_table(mut self, table: &[u8]) -> Self {
296        Arc::make_mut(&mut self.meta).table = RawBytes::new(table).into_owned();
297        self
298    }
299
300    pub fn with_org_table(mut self, org_table: &[u8]) -> Self {
301        Arc::make_mut(&mut self.meta).org_table = RawBytes::new(org_table).into_owned();
302        self
303    }
304
305    pub fn with_name(mut self, name: &[u8]) -> Self {
306        Arc::make_mut(&mut self.meta).name = RawBytes::new(name).into_owned();
307        self
308    }
309
310    pub fn with_org_name(mut self, org_name: &[u8]) -> Self {
311        Arc::make_mut(&mut self.meta).org_name = RawBytes::new(org_name).into_owned();
312        self
313    }
314
315    pub fn with_flags(mut self, flags: ColumnFlags) -> Self {
316        self.flags = Const::new(flags);
317        self
318    }
319
320    pub fn with_column_length(mut self, column_length: u32) -> Self {
321        self.column_length = RawInt::new(column_length);
322        self
323    }
324
325    pub fn with_character_set(mut self, character_set: u16) -> Self {
326        self.character_set = RawInt::new(character_set);
327        self
328    }
329
330    pub fn with_decimals(mut self, decimals: u8) -> Self {
331        self.decimals = RawInt::new(decimals);
332        self
333    }
334
335    /// Returns value of the column_length field of a column packet.
336    ///
337    /// Can be used for text-output formatting.
338    pub fn column_length(&self) -> u32 {
339        *self.column_length
340    }
341
342    /// Returns value of the column_type field of a column packet.
343    pub fn column_type(&self) -> ColumnType {
344        *self.column_type
345    }
346
347    /// Returns value of the character_set field of a column packet.
348    pub fn character_set(&self) -> u16 {
349        *self.character_set
350    }
351
352    /// Returns value of the flags field of a column packet.
353    pub fn flags(&self) -> ColumnFlags {
354        *self.flags
355    }
356
357    /// Returns value of the decimals field of a column packet.
358    ///
359    /// Max shown decimal digits. Can be used for text-output formatting
360    ///
361    /// *   `0x00` for integers and static strings
362    /// *   `0x1f` for dynamic strings, double, float
363    /// *   `0x00..=0x51` for decimals
364    pub fn decimals(&self) -> u8 {
365        *self.decimals
366    }
367
368    /// Returns value of the schema field of a column packet as a byte slice.
369    #[inline(always)]
370    pub fn schema_ref(&self) -> &[u8] {
371        self.meta.schema_ref()
372    }
373
374    /// Returns value of the schema field of a column packet as a string (lossy converted).
375    #[inline(always)]
376    pub fn schema_str(&self) -> Cow<'_, str> {
377        self.meta.schema_str()
378    }
379
380    /// Returns value of the table field of a column packet as a byte slice.
381    #[inline(always)]
382    pub fn table_ref(&self) -> &[u8] {
383        self.meta.table_ref()
384    }
385
386    /// Returns value of the table field of a column packet as a string (lossy converted).
387    #[inline(always)]
388    pub fn table_str(&self) -> Cow<'_, str> {
389        self.meta.table_str()
390    }
391
392    /// Returns value of the org_table field of a column packet as a byte slice.
393    ///
394    /// "org_table" is for original table name.
395    #[inline(always)]
396    pub fn org_table_ref(&self) -> &[u8] {
397        self.meta.org_table_ref()
398    }
399
400    /// Returns value of the org_table field of a column packet as a string (lossy converted).
401    #[inline(always)]
402    pub fn org_table_str(&self) -> Cow<'_, str> {
403        self.meta.org_table_str()
404    }
405
406    /// Returns value of the name field of a column packet as a byte slice.
407    #[inline(always)]
408    pub fn name_ref(&self) -> &[u8] {
409        self.meta.name_ref()
410    }
411
412    /// Returns value of the name field of a column packet as a string (lossy converted).
413    #[inline(always)]
414    pub fn name_str(&self) -> Cow<'_, str> {
415        self.meta.name_str()
416    }
417
418    /// Returns value of the org_name field of a column packet as a byte slice.
419    ///
420    /// "org_name" is for original column name.
421    #[inline(always)]
422    pub fn org_name_ref(&self) -> &[u8] {
423        self.meta.org_name_ref()
424    }
425
426    /// Returns value of the org_name field of a column packet as a string (lossy converted).
427    #[inline(always)]
428    pub fn org_name_str(&self) -> Cow<'_, str> {
429        self.meta.org_name_str()
430    }
431}
432
433/// Represents change in session state (part of MySql's Ok packet).
434#[derive(Debug, Clone, Eq, PartialEq)]
435pub struct SessionStateInfo<'a> {
436    data_type: Const<SessionStateType, u8>,
437    data: RawBytes<'a, LenEnc>,
438}
439
440impl SessionStateInfo<'_> {
441    pub fn into_owned(self) -> SessionStateInfo<'static> {
442        let SessionStateInfo { data_type, data } = self;
443        SessionStateInfo {
444            data_type,
445            data: data.into_owned(),
446        }
447    }
448
449    pub fn data_type(&self) -> SessionStateType {
450        *self.data_type
451    }
452
453    /// Returns raw session state info data.
454    pub fn data_ref(&self) -> &[u8] {
455        self.data.as_bytes()
456    }
457
458    /// Tries to decode session state info data.
459    pub fn decode(&self) -> io::Result<SessionStateChange<'_>> {
460        ParseBuf(self.data.as_bytes()).parse_unchecked(*self.data_type)
461    }
462}
463
464impl<'de> MyDeserialize<'de> for SessionStateInfo<'de> {
465    const SIZE: Option<usize> = None;
466    type Ctx = ();
467
468    fn deserialize(_ctx: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
469        Ok(SessionStateInfo {
470            data_type: buf.parse(())?,
471            data: buf.parse(())?,
472        })
473    }
474}
475
476impl MySerialize for SessionStateInfo<'_> {
477    fn serialize(&self, buf: &mut Vec<u8>) {
478        self.data_type.serialize(&mut *buf);
479        self.data.serialize(buf);
480    }
481}
482
483/// Represents MySql's Ok packet.
484#[derive(Debug, Clone, Eq, PartialEq)]
485pub struct OkPacketBody<'a> {
486    affected_rows: RawInt<LenEnc>,
487    last_insert_id: RawInt<LenEnc>,
488    status_flags: Const<StatusFlags, LeU16>,
489    warnings: RawInt<LeU16>,
490    info: RawBytes<'a, LenEnc>,
491    session_state_info: RawBytes<'a, LenEnc>,
492}
493
494/// OK packet kind (see _OK packet identifier_ section of [WL#7766][1]).
495///
496/// [1]: https://dev.mysql.com/worklog/task/?id=7766
497pub trait OkPacketKind {
498    const HEADER: u8;
499
500    fn parse_body<'de>(
501        capabilities: CapabilityFlags,
502        buf: &mut ParseBuf<'de>,
503    ) -> io::Result<OkPacketBody<'de>>;
504}
505
506/// Ok packet that terminates a result set (text or binary).
507#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
508pub struct ResultSetTerminator;
509
510impl OkPacketKind for ResultSetTerminator {
511    const HEADER: u8 = 0xFE;
512
513    fn parse_body<'de>(
514        capabilities: CapabilityFlags,
515        buf: &mut ParseBuf<'de>,
516    ) -> io::Result<OkPacketBody<'de>> {
517        // We need to skip affected_rows and insert_id here
518        // because valid content of EOF packet includes
519        // packet marker, server status and warning count only.
520        // (see `read_ok_ex` in sql-common/client.cc)
521        buf.parse::<RawInt<LenEnc>>(())?;
522        buf.parse::<RawInt<LenEnc>>(())?;
523
524        // assume CLIENT_PROTOCOL_41 flag
525        let mut sbuf: ParseBuf<'_> = buf.parse(4)?;
526        let status_flags: Const<StatusFlags, LeU16> = sbuf.parse_unchecked(())?;
527        let warnings = sbuf.parse_unchecked(())?;
528
529        let (info, session_state_info) =
530            if capabilities.contains(CapabilityFlags::CLIENT_SESSION_TRACK) && !buf.is_empty() {
531                let info = buf.parse(())?;
532                let session_state_info =
533                    if status_flags.contains(StatusFlags::SERVER_SESSION_STATE_CHANGED) {
534                        buf.parse(())?
535                    } else {
536                        RawBytes::default()
537                    };
538                (info, session_state_info)
539            } else if !buf.is_empty() && buf.0[0] > 0 {
540                // The `info` field is a `string<EOF>` according to the MySQL Internals
541                // Manual, but actually it's a `string<lenenc>`.
542                // SEE: sql/protocol_classics.cc `net_send_ok`
543                let info = buf.parse(())?;
544                (info, RawBytes::default())
545            } else {
546                (RawBytes::default(), RawBytes::default())
547            };
548
549        Ok(OkPacketBody {
550            affected_rows: RawInt::new(0),
551            last_insert_id: RawInt::new(0),
552            status_flags,
553            warnings,
554            info,
555            session_state_info,
556        })
557    }
558}
559
560/// Old deprecated EOF packet.
561#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
562pub struct OldEofPacket;
563
564impl OkPacketKind for OldEofPacket {
565    const HEADER: u8 = 0xFE;
566
567    fn parse_body<'de>(
568        _: CapabilityFlags,
569        buf: &mut ParseBuf<'de>,
570    ) -> io::Result<OkPacketBody<'de>> {
571        // We assume that CLIENT_PROTOCOL_41 was set
572        let mut buf: ParseBuf<'_> = buf.parse(4)?;
573        let warnings = buf.parse_unchecked(())?;
574        let status_flags = buf.parse_unchecked(())?;
575
576        Ok(OkPacketBody {
577            affected_rows: RawInt::new(0),
578            last_insert_id: RawInt::new(0),
579            status_flags,
580            warnings,
581            info: RawBytes::new(&[][..]),
582            session_state_info: RawBytes::new(&[][..]),
583        })
584    }
585}
586
587/// This packet terminates a binlog network stream.
588#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
589pub struct NetworkStreamTerminator;
590
591impl OkPacketKind for NetworkStreamTerminator {
592    const HEADER: u8 = 0xFE;
593
594    fn parse_body<'de>(
595        flags: CapabilityFlags,
596        buf: &mut ParseBuf<'de>,
597    ) -> io::Result<OkPacketBody<'de>> {
598        OldEofPacket::parse_body(flags, buf)
599    }
600}
601
602/// Ok packet that is not a result set terminator.
603#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
604pub struct CommonOkPacket;
605
606impl OkPacketKind for CommonOkPacket {
607    const HEADER: u8 = 0x00;
608
609    fn parse_body<'de>(
610        capabilities: CapabilityFlags,
611        buf: &mut ParseBuf<'de>,
612    ) -> io::Result<OkPacketBody<'de>> {
613        let affected_rows = buf.parse(())?;
614        let last_insert_id = buf.parse(())?;
615
616        // We assume that CLIENT_PROTOCOL_41 was set
617        let mut sbuf: ParseBuf<'_> = buf.parse(4)?;
618        let status_flags: Const<StatusFlags, LeU16> = sbuf.parse_unchecked(())?;
619        let warnings = sbuf.parse_unchecked(())?;
620
621        let (info, session_state_info) =
622            if capabilities.contains(CapabilityFlags::CLIENT_SESSION_TRACK) && !buf.is_empty() {
623                let info = buf.parse(())?;
624                let session_state_info =
625                    if status_flags.contains(StatusFlags::SERVER_SESSION_STATE_CHANGED) {
626                        buf.parse(())?
627                    } else {
628                        RawBytes::default()
629                    };
630                (info, session_state_info)
631            } else if !buf.is_empty() && buf.0[0] > 0 {
632                // The `info` field is a `string<EOF>` according to the MySQL Internals
633                // Manual, but actually it's a `string<lenenc>`.
634                // SEE: sql/protocol_classics.cc `net_send_ok`
635                let info = buf.parse(())?;
636                (info, RawBytes::default())
637            } else {
638                (RawBytes::default(), RawBytes::default())
639            };
640
641        Ok(OkPacketBody {
642            affected_rows,
643            last_insert_id,
644            status_flags,
645            warnings,
646            info,
647            session_state_info,
648        })
649    }
650}
651
652impl<'a> TryFrom<OkPacketBody<'a>> for OkPacket<'a> {
653    type Error = io::Error;
654
655    fn try_from(body: OkPacketBody<'a>) -> io::Result<Self> {
656        Ok(OkPacket {
657            affected_rows: *body.affected_rows,
658            last_insert_id: if *body.last_insert_id == 0 {
659                None
660            } else {
661                Some(*body.last_insert_id)
662            },
663            status_flags: *body.status_flags,
664            warnings: *body.warnings,
665            info: if !body.info.is_empty() {
666                Some(body.info)
667            } else {
668                None
669            },
670            session_state_info: if !body.session_state_info.is_empty() {
671                Some(body.session_state_info)
672            } else {
673                None
674            },
675        })
676    }
677}
678
679/// Represents MySql's Ok packet.
680#[derive(Debug, Clone, Eq, PartialEq)]
681pub struct OkPacket<'a> {
682    affected_rows: u64,
683    last_insert_id: Option<u64>,
684    status_flags: StatusFlags,
685    warnings: u16,
686    info: Option<RawBytes<'a, LenEnc>>,
687    session_state_info: Option<RawBytes<'a, LenEnc>>,
688}
689
690impl OkPacket<'_> {
691    pub fn into_owned(self) -> OkPacket<'static> {
692        OkPacket {
693            affected_rows: self.affected_rows,
694            last_insert_id: self.last_insert_id,
695            status_flags: self.status_flags,
696            warnings: self.warnings,
697            info: self.info.map(|x| x.into_owned()),
698            session_state_info: self.session_state_info.map(|x| x.into_owned()),
699        }
700    }
701
702    /// Value of the affected_rows field of an Ok packet.
703    pub fn affected_rows(&self) -> u64 {
704        self.affected_rows
705    }
706
707    /// Value of the last_insert_id field of an Ok packet.
708    pub fn last_insert_id(&self) -> Option<u64> {
709        self.last_insert_id
710    }
711
712    /// Value of the status_flags field of an Ok packet.
713    pub fn status_flags(&self) -> StatusFlags {
714        self.status_flags
715    }
716
717    /// Value of the warnings field of an Ok packet.
718    pub fn warnings(&self) -> u16 {
719        self.warnings
720    }
721
722    /// Value of the info field of an Ok packet as a byte slice.
723    pub fn info_ref(&self) -> Option<&[u8]> {
724        self.info.as_ref().map(|x| x.as_bytes())
725    }
726
727    /// Value of the info field of an Ok packet as a string (lossy converted).
728    pub fn info_str(&self) -> Option<Cow<'_, str>> {
729        self.info.as_ref().map(|x| x.as_str())
730    }
731
732    /// Returns raw reference to a session state info.
733    pub fn session_state_info_ref(&self) -> Option<&[u8]> {
734        self.session_state_info.as_ref().map(|x| x.as_bytes())
735    }
736
737    /// Tries to parse session state info, if any.
738    pub fn session_state_info(&self) -> io::Result<Vec<SessionStateInfo<'_>>> {
739        self.session_state_info_ref()
740            .map(|data| {
741                let mut data = ParseBuf(data);
742                let mut entries = Vec::new();
743                while !data.is_empty() {
744                    entries.push(data.parse(())?);
745                }
746                Ok(entries)
747            })
748            .transpose()
749            .map(|x| x.unwrap_or_default())
750    }
751}
752
753#[derive(Debug, Clone, PartialEq, Eq)]
754pub struct OkPacketDeserializer<'de, T>(OkPacket<'de>, PhantomData<T>);
755
756impl<'de, T> OkPacketDeserializer<'de, T> {
757    pub fn into_inner(self) -> OkPacket<'de> {
758        self.0
759    }
760}
761
762impl<'de, T> From<OkPacketDeserializer<'de, T>> for OkPacket<'de> {
763    fn from(x: OkPacketDeserializer<'de, T>) -> Self {
764        x.0
765    }
766}
767
768#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
769#[error("Invalid OK packet header")]
770pub struct InvalidOkPacketHeader;
771
772impl<'de, T: OkPacketKind> MyDeserialize<'de> for OkPacketDeserializer<'de, T> {
773    const SIZE: Option<usize> = None;
774    type Ctx = CapabilityFlags;
775
776    fn deserialize(capabilities: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
777        if *buf.parse::<RawInt<u8>>(())? == T::HEADER {
778            let body = T::parse_body(capabilities, buf)?;
779            let ok = OkPacket::try_from(body)?;
780            Ok(Self(ok, PhantomData))
781        } else {
782            Err(io::Error::new(
783                io::ErrorKind::InvalidData,
784                InvalidOkPacketHeader,
785            ))
786        }
787    }
788}
789
790/// Progress report information (may be in an error packet of MariaDB server).
791#[derive(Debug, Clone, Eq, PartialEq)]
792pub struct ProgressReport<'a> {
793    stage: RawInt<u8>,
794    max_stage: RawInt<u8>,
795    progress: RawInt<LeU24>,
796    stage_info: RawBytes<'a, LenEnc>,
797}
798
799impl<'a> ProgressReport<'a> {
800    pub fn new(
801        stage: u8,
802        max_stage: u8,
803        progress: u32,
804        stage_info: impl Into<Cow<'a, [u8]>>,
805    ) -> ProgressReport<'a> {
806        ProgressReport {
807            stage: RawInt::new(stage),
808            max_stage: RawInt::new(max_stage),
809            progress: RawInt::new(progress),
810            stage_info: RawBytes::new(stage_info),
811        }
812    }
813
814    /// 1 to max_stage
815    pub fn stage(&self) -> u8 {
816        *self.stage
817    }
818
819    pub fn max_stage(&self) -> u8 {
820        *self.max_stage
821    }
822
823    /// Progress as '% * 1000'
824    pub fn progress(&self) -> u32 {
825        *self.progress
826    }
827
828    /// Status or state name as a byte slice.
829    pub fn stage_info_ref(&self) -> &[u8] {
830        self.stage_info.as_bytes()
831    }
832
833    /// Status or state name as a string (lossy converted).
834    pub fn stage_info_str(&self) -> Cow<'_, str> {
835        self.stage_info.as_str()
836    }
837
838    pub fn into_owned(self) -> ProgressReport<'static> {
839        ProgressReport {
840            stage: self.stage,
841            max_stage: self.max_stage,
842            progress: self.progress,
843            stage_info: self.stage_info.into_owned(),
844        }
845    }
846}
847
848impl<'de> MyDeserialize<'de> for ProgressReport<'de> {
849    const SIZE: Option<usize> = None;
850    type Ctx = ();
851
852    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
853        let mut sbuf: ParseBuf<'_> = buf.parse(6)?;
854
855        sbuf.skip(1); // Ignore number of strings.
856
857        Ok(ProgressReport {
858            stage: sbuf.parse_unchecked(())?,
859            max_stage: sbuf.parse_unchecked(())?,
860            progress: sbuf.parse_unchecked(())?,
861            stage_info: buf.parse(())?,
862        })
863    }
864}
865
866impl MySerialize for ProgressReport<'_> {
867    fn serialize(&self, buf: &mut Vec<u8>) {
868        buf.put_u8(1);
869        self.stage.serialize(&mut *buf);
870        self.max_stage.serialize(&mut *buf);
871        self.progress.serialize(&mut *buf);
872        self.stage_info.serialize(buf);
873    }
874}
875
876impl fmt::Display for ProgressReport<'_> {
877    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
878        write!(
879            f,
880            "Stage: {} of {} '{}'  {:.2}% of stage done",
881            self.stage(),
882            self.max_stage(),
883            self.progress(),
884            self.stage_info_str()
885        )
886    }
887}
888
889define_header!(
890    ErrPacketHeader,
891    InvalidErrPacketHeader("Invalid error packet header"),
892    0xFF
893);
894
895/// MySql error packet.
896///
897/// May hold an error or a progress report.
898#[derive(Debug, Clone, PartialEq)]
899pub enum ErrPacket<'a> {
900    Error(ServerError<'a>),
901    Progress(ProgressReport<'a>),
902}
903
904impl ErrPacket<'_> {
905    /// Returns false if this error packet contains progress report.
906    pub fn is_error(&self) -> bool {
907        matches!(self, ErrPacket::Error { .. })
908    }
909
910    /// Returns true if this error packet contains progress report.
911    pub fn is_progress_report(&self) -> bool {
912        !self.is_error()
913    }
914
915    /// Will panic if ErrPacket does not contains progress report
916    pub fn progress_report(&self) -> &ProgressReport<'_> {
917        match *self {
918            ErrPacket::Progress(ref progress_report) => progress_report,
919            _ => panic!("This ErrPacket does not contains progress report"),
920        }
921    }
922
923    /// Will panic if ErrPacket does not contains a `ServerError`.
924    pub fn server_error(&self) -> &ServerError<'_> {
925        match self {
926            ErrPacket::Error(error) => error,
927            ErrPacket::Progress(_) => panic!("This ErrPacket does not contain a ServerError"),
928        }
929    }
930}
931
932impl<'de> MyDeserialize<'de> for ErrPacket<'de> {
933    const SIZE: Option<usize> = None;
934    type Ctx = CapabilityFlags;
935
936    fn deserialize(capabilities: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
937        let mut sbuf: ParseBuf<'_> = buf.parse(3)?;
938        sbuf.parse_unchecked::<ErrPacketHeader>(())?;
939        let code: RawInt<LeU16> = sbuf.parse_unchecked(())?;
940
941        if *code == 0xFFFF && capabilities.contains(CapabilityFlags::CLIENT_PROGRESS_OBSOLETE) {
942            buf.parse(()).map(ErrPacket::Progress)
943        } else {
944            buf.parse((
945                *code,
946                capabilities.contains(CapabilityFlags::CLIENT_PROTOCOL_41),
947            ))
948            .map(ErrPacket::Error)
949        }
950    }
951}
952
953impl MySerialize for ErrPacket<'_> {
954    fn serialize(&self, buf: &mut Vec<u8>) {
955        ErrPacketHeader::new().serialize(&mut *buf);
956        match self {
957            ErrPacket::Error(server_error) => {
958                server_error.code.serialize(&mut *buf);
959                server_error.serialize(buf);
960            }
961            ErrPacket::Progress(progress_report) => {
962                RawInt::<LeU16>::new(0xFFFF).serialize(&mut *buf);
963                progress_report.serialize(buf);
964            }
965        }
966    }
967}
968
969impl fmt::Display for ErrPacket<'_> {
970    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
971        match self {
972            ErrPacket::Error(server_error) => write!(f, "{}", server_error),
973            ErrPacket::Progress(progress_report) => write!(f, "{}", progress_report),
974        }
975    }
976}
977
978define_header!(
979    SqlStateMarker,
980    InvalidSqlStateMarker("Invalid SqlStateMarker value"),
981    b'#'
982);
983
984/// MySql error state.
985#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
986pub struct SqlState {
987    __state_marker: SqlStateMarker,
988    state: [u8; 5],
989}
990
991impl SqlState {
992    /// Creates new sql state.
993    pub fn new(state: [u8; 5]) -> Self {
994        Self {
995            __state_marker: SqlStateMarker::new(),
996            state,
997        }
998    }
999
1000    /// Returns an sql state as bytes.
1001    pub fn as_bytes(&self) -> [u8; 5] {
1002        self.state
1003    }
1004
1005    /// Returns an sql state as a string (lossy converted).
1006    pub fn as_str(&self) -> Cow<'_, str> {
1007        String::from_utf8_lossy(&self.state)
1008    }
1009}
1010
1011impl<'de> MyDeserialize<'de> for SqlState {
1012    const SIZE: Option<usize> = Some(6);
1013    type Ctx = ();
1014
1015    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1016        Ok(Self {
1017            __state_marker: buf.parse(())?,
1018            state: buf.parse(())?,
1019        })
1020    }
1021}
1022
1023impl MySerialize for SqlState {
1024    fn serialize(&self, buf: &mut Vec<u8>) {
1025        self.__state_marker.serialize(buf);
1026        self.state.serialize(buf);
1027    }
1028}
1029
1030/// MySql error packet.
1031///
1032/// May hold an error or a progress report.
1033#[derive(Debug, Clone, PartialEq)]
1034pub struct ServerError<'a> {
1035    code: RawInt<LeU16>,
1036    state: Option<SqlState>,
1037    message: RawBytes<'a, EofBytes>,
1038}
1039
1040impl<'a> ServerError<'a> {
1041    pub fn new(code: u16, state: Option<SqlState>, msg: impl Into<Cow<'a, [u8]>>) -> Self {
1042        Self {
1043            code: RawInt::new(code),
1044            state,
1045            message: RawBytes::new(msg),
1046        }
1047    }
1048
1049    /// Returns an error code.
1050    pub fn error_code(&self) -> u16 {
1051        *self.code
1052    }
1053
1054    /// Returns an sql state.
1055    pub fn sql_state_ref(&self) -> Option<&SqlState> {
1056        self.state.as_ref()
1057    }
1058
1059    /// Returns an error message.
1060    pub fn message_ref(&self) -> &[u8] {
1061        self.message.as_bytes()
1062    }
1063
1064    /// Returns an error message as a string (lossy converted).
1065    pub fn message_str(&self) -> Cow<'_, str> {
1066        self.message.as_str()
1067    }
1068
1069    pub fn into_owned(self) -> ServerError<'static> {
1070        ServerError {
1071            code: self.code,
1072            state: self.state,
1073            message: self.message.into_owned(),
1074        }
1075    }
1076}
1077
1078impl<'de> MyDeserialize<'de> for ServerError<'de> {
1079    const SIZE: Option<usize> = None;
1080    /// An error packet error code + whether CLIENT_PROTOCOL_41 capability was negotiated.
1081    type Ctx = (u16, bool);
1082
1083    fn deserialize((code, protocol_41): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1084        let server_error = if protocol_41 {
1085            ServerError {
1086                code: RawInt::new(code),
1087                state: Some(buf.parse(())?),
1088                message: buf.parse(())?,
1089            }
1090        } else {
1091            ServerError {
1092                code: RawInt::new(code),
1093                state: None,
1094                message: buf.parse(())?,
1095            }
1096        };
1097        Ok(server_error)
1098    }
1099}
1100
1101impl MySerialize for ServerError<'_> {
1102    fn serialize(&self, buf: &mut Vec<u8>) {
1103        if let Some(state) = &self.state {
1104            state.serialize(buf);
1105        }
1106        self.message.serialize(buf);
1107    }
1108}
1109
1110impl fmt::Display for ServerError<'_> {
1111    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1112        let sql_state_str = self
1113            .sql_state_ref()
1114            .map(|s| format!(" ({})", s.as_str()))
1115            .unwrap_or_default();
1116
1117        write!(
1118            f,
1119            "ERROR {}{}: {}",
1120            self.error_code(),
1121            sql_state_str,
1122            self.message_str()
1123        )
1124    }
1125}
1126
1127define_header!(
1128    LocalInfileHeader,
1129    InvalidLocalInfileHeader("Invalid LOCAL_INFILE header"),
1130    0xFB
1131);
1132
1133/// Represents MySql's local infile packet.
1134#[derive(Debug, Clone, Eq, PartialEq)]
1135pub struct LocalInfilePacket<'a> {
1136    __header: LocalInfileHeader,
1137    file_name: RawBytes<'a, EofBytes>,
1138}
1139
1140impl<'a> LocalInfilePacket<'a> {
1141    pub fn new(file_name: impl Into<Cow<'a, [u8]>>) -> Self {
1142        Self {
1143            __header: LocalInfileHeader::new(),
1144            file_name: RawBytes::new(file_name),
1145        }
1146    }
1147
1148    /// Value of the file_name field of a local infile packet as a byte slice.
1149    pub fn file_name_ref(&self) -> &[u8] {
1150        self.file_name.as_bytes()
1151    }
1152
1153    /// Value of the file_name field of a local infile packet as a string (lossy converted).
1154    pub fn file_name_str(&self) -> Cow<'_, str> {
1155        self.file_name.as_str()
1156    }
1157
1158    pub fn into_owned(self) -> LocalInfilePacket<'static> {
1159        LocalInfilePacket {
1160            __header: self.__header,
1161            file_name: self.file_name.into_owned(),
1162        }
1163    }
1164}
1165
1166impl<'de> MyDeserialize<'de> for LocalInfilePacket<'de> {
1167    const SIZE: Option<usize> = None;
1168    type Ctx = ();
1169
1170    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1171        Ok(LocalInfilePacket {
1172            __header: buf.parse(())?,
1173            file_name: buf.parse(())?,
1174        })
1175    }
1176}
1177
1178impl MySerialize for LocalInfilePacket<'_> {
1179    fn serialize(&self, buf: &mut Vec<u8>) {
1180        self.__header.serialize(buf);
1181        self.file_name.serialize(buf);
1182    }
1183}
1184
1185const MYSQL_OLD_PASSWORD_PLUGIN_NAME: &[u8] = b"mysql_old_password";
1186const MYSQL_NATIVE_PASSWORD_PLUGIN_NAME: &[u8] = b"mysql_native_password";
1187const CACHING_SHA2_PASSWORD_PLUGIN_NAME: &[u8] = b"caching_sha2_password";
1188const MYSQL_CLEAR_PASSWORD_PLUGIN_NAME: &[u8] = b"mysql_clear_password";
1189const ED25519_PLUGIN_NAME: &[u8] = b"client_ed25519";
1190const PARSEC_PLUGIN_NAME: &[u8] = b"parsec";
1191
1192#[derive(Debug, Clone, PartialEq, Eq)]
1193pub enum AuthPluginData<'a> {
1194    /// Auth data for the `mysql_old_password` plugin.
1195    Old([u8; 8]),
1196    /// Auth data for the `mysql_native_password` plugin.
1197    Native([u8; 20]),
1198    /// Auth data for `sha2_password` and `caching_sha2_password` plugins.
1199    Sha2([u8; 32]),
1200    /// Clear password for `mysql_clear_password` plugin.
1201    Clear(Cow<'a, [u8]>),
1202    /// Auth data for MariaDB's `client_ed25519` plugin.
1203    ///
1204    /// This plugin is known to the library but the actual support is enabled
1205    /// by the `client_ed25519` feature.
1206    Ed25519([u8; 64]),
1207    /// Auth data for MariaDB's `parsec` plugin.
1208    ///
1209    /// Plugin support is enabled by the `client_parsec` feature.
1210    Parsec([u8; 96]),
1211}
1212
1213impl AuthPluginData<'_> {
1214    pub fn into_owned(self) -> AuthPluginData<'static> {
1215        match self {
1216            AuthPluginData::Old(x) => AuthPluginData::Old(x),
1217            AuthPluginData::Native(x) => AuthPluginData::Native(x),
1218            AuthPluginData::Sha2(x) => AuthPluginData::Sha2(x),
1219            AuthPluginData::Clear(x) => AuthPluginData::Clear(Cow::Owned(x.into_owned())),
1220            AuthPluginData::Ed25519(x) => AuthPluginData::Ed25519(x),
1221            AuthPluginData::Parsec(x) => AuthPluginData::Parsec(x),
1222        }
1223    }
1224}
1225
1226impl std::ops::Deref for AuthPluginData<'_> {
1227    type Target = [u8];
1228
1229    fn deref(&self) -> &Self::Target {
1230        match self {
1231            Self::Sha2(x) => &x[..],
1232            Self::Native(x) => &x[..],
1233            Self::Old(x) => &x[..],
1234            Self::Clear(x) => &x[..],
1235            Self::Ed25519(x) => &x[..],
1236            Self::Parsec(x) => &x[..],
1237        }
1238    }
1239}
1240
1241impl MySerialize for AuthPluginData<'_> {
1242    fn serialize(&self, buf: &mut Vec<u8>) {
1243        match self {
1244            Self::Sha2(x) => buf.put_slice(&x[..]),
1245            Self::Native(x) => buf.put_slice(&x[..]),
1246            Self::Old(x) => {
1247                buf.put_slice(&x[..]);
1248                buf.push(0);
1249            }
1250            Self::Clear(x) => {
1251                buf.put_slice(x);
1252                buf.push(0);
1253            }
1254            Self::Ed25519(x) => buf.put_slice(&x[..]),
1255            Self::Parsec(x) => buf.put_slice(&x[..]),
1256        }
1257    }
1258}
1259
1260/// Authentication plugin
1261#[derive(Debug, Clone, Eq, PartialEq, Hash)]
1262pub enum AuthPlugin<'a> {
1263    /// Old Password Authentication
1264    MysqlOldPassword,
1265    /// Client-Side Cleartext Pluggable Authentication
1266    MysqlClearPassword,
1267    /// Legacy authentication plugin
1268    MysqlNativePassword,
1269    /// Default since MySql v8.0.4
1270    CachingSha2Password,
1271    /// MariaDB's Ed25519 based authentication
1272    ///
1273    /// This plugin is known to the library but the actual support is enabled
1274    /// by the `client_ed25519` feature.
1275    Ed25519,
1276    /// MariaDB's Parsec(Password Authentication using Response Signed with Eliptic Curve) authentication
1277    ///
1278    /// Actual support is enabled by the `client_parsec` feature.
1279    MariadbParsec {
1280        iterations: u32,
1281        ext_salt: [u8; 18],
1282    },
1283    Other(Cow<'a, [u8]>),
1284}
1285
1286impl<'de> MyDeserialize<'de> for AuthPlugin<'de> {
1287    const SIZE: Option<usize> = None;
1288    type Ctx = ();
1289
1290    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1291        Ok(Self::from_bytes(buf.eat_all()))
1292    }
1293}
1294
1295impl MySerialize for AuthPlugin<'_> {
1296    fn serialize(&self, buf: &mut Vec<u8>) {
1297        buf.put_slice(self.as_bytes());
1298        buf.put_u8(0);
1299    }
1300}
1301
1302impl<'a> AuthPlugin<'a> {
1303    pub fn from_bytes(name: &'a [u8]) -> AuthPlugin<'a> {
1304        let name = if let [name @ .., 0] = name {
1305            name
1306        } else {
1307            name
1308        };
1309        match name {
1310            CACHING_SHA2_PASSWORD_PLUGIN_NAME => AuthPlugin::CachingSha2Password,
1311            MYSQL_NATIVE_PASSWORD_PLUGIN_NAME => AuthPlugin::MysqlNativePassword,
1312            MYSQL_OLD_PASSWORD_PLUGIN_NAME => AuthPlugin::MysqlOldPassword,
1313            MYSQL_CLEAR_PASSWORD_PLUGIN_NAME => AuthPlugin::MysqlClearPassword,
1314            ED25519_PLUGIN_NAME => AuthPlugin::Ed25519,
1315            PARSEC_PLUGIN_NAME => AuthPlugin::MariadbParsec {
1316                iterations: 0,
1317                ext_salt: [0; 18],
1318            },
1319            name => AuthPlugin::Other(Cow::Borrowed(name)),
1320        }
1321    }
1322
1323    pub fn as_bytes(&self) -> &[u8] {
1324        match self {
1325            AuthPlugin::CachingSha2Password => CACHING_SHA2_PASSWORD_PLUGIN_NAME,
1326            AuthPlugin::MysqlNativePassword => MYSQL_NATIVE_PASSWORD_PLUGIN_NAME,
1327            AuthPlugin::MysqlOldPassword => MYSQL_OLD_PASSWORD_PLUGIN_NAME,
1328            AuthPlugin::MysqlClearPassword => MYSQL_CLEAR_PASSWORD_PLUGIN_NAME,
1329            AuthPlugin::Ed25519 => ED25519_PLUGIN_NAME,
1330            AuthPlugin::MariadbParsec { .. } => PARSEC_PLUGIN_NAME,
1331            AuthPlugin::Other(name) => name,
1332        }
1333    }
1334
1335    pub fn into_owned(self) -> AuthPlugin<'static> {
1336        match self {
1337            AuthPlugin::CachingSha2Password => AuthPlugin::CachingSha2Password,
1338            AuthPlugin::MysqlNativePassword => AuthPlugin::MysqlNativePassword,
1339            AuthPlugin::MysqlOldPassword => AuthPlugin::MysqlOldPassword,
1340            AuthPlugin::MysqlClearPassword => AuthPlugin::MysqlClearPassword,
1341            AuthPlugin::Ed25519 => AuthPlugin::Ed25519,
1342            AuthPlugin::MariadbParsec {
1343                iterations,
1344                ext_salt,
1345            } => AuthPlugin::MariadbParsec {
1346                iterations,
1347                ext_salt,
1348            },
1349            AuthPlugin::Other(name) => AuthPlugin::Other(Cow::Owned(name.into_owned())),
1350        }
1351    }
1352
1353    pub fn borrow(&self) -> AuthPlugin<'_> {
1354        match self {
1355            AuthPlugin::CachingSha2Password => AuthPlugin::CachingSha2Password,
1356            AuthPlugin::MysqlNativePassword => AuthPlugin::MysqlNativePassword,
1357            AuthPlugin::MysqlOldPassword => AuthPlugin::MysqlOldPassword,
1358            AuthPlugin::MysqlClearPassword => AuthPlugin::MysqlClearPassword,
1359            AuthPlugin::Ed25519 => AuthPlugin::Ed25519,
1360            AuthPlugin::MariadbParsec {
1361                iterations,
1362                ext_salt,
1363            } => AuthPlugin::MariadbParsec {
1364                iterations: *iterations,
1365                ext_salt: *ext_salt,
1366            },
1367            AuthPlugin::Other(name) => AuthPlugin::Other(Cow::Borrowed(name.as_ref())),
1368        }
1369    }
1370
1371    /// Generates auth plugin data for this plugin.
1372    ///
1373    /// It'll generate `None` if password is `None` or empty.
1374    ///
1375    /// Note, that you should trim terminating null character from the `nonce`.
1376    ///
1377    /// # Panic
1378    ///
1379    /// * [`AuthPlugin::Ed25519`] will panic if `client_ed25519` feature is disabled.
1380    /// * [`AuthPlugin::MariadbParsec`] will panic if `client_parsec` feature is disabled.
1381    pub fn gen_data<'b>(&self, pass: Option<&'b str>, nonce: &[u8]) -> Option<AuthPluginData<'b>> {
1382        use super::scramble::{scramble_323, scramble_native, scramble_sha256};
1383
1384        match pass {
1385            Some(pass) if !pass.is_empty() => match self {
1386                AuthPlugin::CachingSha2Password => {
1387                    scramble_sha256(nonce, pass.as_bytes()).map(AuthPluginData::Sha2)
1388                }
1389                AuthPlugin::MysqlNativePassword => {
1390                    scramble_native(nonce, pass.as_bytes()).map(AuthPluginData::Native)
1391                }
1392                AuthPlugin::MysqlOldPassword => {
1393                    scramble_323(nonce.chunks(8).next().unwrap(), pass.as_bytes())
1394                        .map(AuthPluginData::Old)
1395                }
1396                AuthPlugin::MysqlClearPassword => {
1397                    Some(AuthPluginData::Clear(Cow::Borrowed(pass.as_bytes())))
1398                }
1399                AuthPlugin::Ed25519 => Some(AuthPluginData::Ed25519(create_response_for_ed25519(
1400                    pass.as_bytes(),
1401                    nonce,
1402                ))),
1403                AuthPlugin::MariadbParsec {
1404                    iterations,
1405                    ext_salt,
1406                } => match nonce.try_into() {
1407                    Ok(nonce_array) => Some(AuthPluginData::Parsec(create_response_for_parsec(
1408                        pass.as_bytes(),
1409                        nonce_array,
1410                        *iterations,
1411                        ext_salt,
1412                    ))),
1413                    Err(_) => None,
1414                },
1415                AuthPlugin::Other(_) => None,
1416            },
1417            _ => None,
1418        }
1419    }
1420
1421    /// Reads additional (packet) data required for authentication data generation.
1422    ///
1423    /// Currently only needed for Parseec plugin to read the result of additional packets exchange with the server.
1424    pub fn read_add_data(&mut self, payload: &[u8]) -> Option<()> {
1425        match self {
1426            AuthPlugin::MariadbParsec {
1427                iterations,
1428                ext_salt,
1429            } => {
1430                let result = parse_parsec_salt(payload)?;
1431                *iterations = result.0;
1432                *ext_salt = *result.1;
1433                Some(())
1434            }
1435            _ => None,
1436        }
1437    }
1438}
1439
1440define_header!(
1441    AuthMoreDataHeader,
1442    InvalidAuthMoreDataHeader("Invalid AuthMoreData header"),
1443    0x01
1444);
1445
1446/// Extra auth-data beyond the initial challenge.
1447#[derive(Debug, Clone, Eq, PartialEq)]
1448pub struct AuthMoreData<'a> {
1449    __header: AuthMoreDataHeader,
1450    data: RawBytes<'a, EofBytes>,
1451}
1452
1453impl<'a> AuthMoreData<'a> {
1454    pub fn new(data: impl Into<Cow<'a, [u8]>>) -> Self {
1455        Self {
1456            __header: AuthMoreDataHeader::new(),
1457            data: RawBytes::new(data),
1458        }
1459    }
1460
1461    pub fn data(&self) -> &[u8] {
1462        self.data.as_bytes()
1463    }
1464
1465    pub fn into_owned(self) -> AuthMoreData<'static> {
1466        AuthMoreData {
1467            __header: self.__header,
1468            data: self.data.into_owned(),
1469        }
1470    }
1471}
1472
1473impl<'de> MyDeserialize<'de> for AuthMoreData<'de> {
1474    const SIZE: Option<usize> = None;
1475    type Ctx = ();
1476
1477    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1478        Ok(Self {
1479            __header: buf.parse(())?,
1480            data: buf.parse(())?,
1481        })
1482    }
1483}
1484
1485impl MySerialize for AuthMoreData<'_> {
1486    fn serialize(&self, buf: &mut Vec<u8>) {
1487        self.__header.serialize(&mut *buf);
1488        self.data.serialize(buf);
1489    }
1490}
1491
1492define_header!(
1493    PublicKeyResponseHeader,
1494    InvalidPublicKeyResponse("Invalid PublicKeyResponse header"),
1495    0x01
1496);
1497
1498/// A server response to a [`PublicKeyRequest`] containing a public RSA key for authentication protection.
1499///
1500/// [`PublicKeyRequest`]: crate::packets::caching_sha2_password::PublicKeyRequest
1501#[derive(Debug, Clone, Eq, PartialEq)]
1502pub struct PublicKeyResponse<'a> {
1503    __header: PublicKeyResponseHeader,
1504    rsa_key: RawBytes<'a, EofBytes>,
1505}
1506
1507impl<'a> PublicKeyResponse<'a> {
1508    pub fn new(rsa_key: impl Into<Cow<'a, [u8]>>) -> Self {
1509        Self {
1510            __header: PublicKeyResponseHeader::new(),
1511            rsa_key: RawBytes::new(rsa_key),
1512        }
1513    }
1514
1515    /// The server's RSA public key in PEM format.
1516    pub fn rsa_key(&self) -> Cow<'_, str> {
1517        self.rsa_key.as_str()
1518    }
1519}
1520
1521impl<'de> MyDeserialize<'de> for PublicKeyResponse<'de> {
1522    const SIZE: Option<usize> = None;
1523    type Ctx = ();
1524
1525    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1526        Ok(Self {
1527            __header: buf.parse(())?,
1528            rsa_key: buf.parse(())?,
1529        })
1530    }
1531}
1532
1533impl MySerialize for PublicKeyResponse<'_> {
1534    fn serialize(&self, buf: &mut Vec<u8>) {
1535        self.__header.serialize(&mut *buf);
1536        self.rsa_key.serialize(buf);
1537    }
1538}
1539
1540define_header!(
1541    AuthSwitchRequestHeader,
1542    InvalidAuthSwitchRequestHeader("Invalid auth switch request header"),
1543    0xFE
1544);
1545
1546/// Old Authentication Method Switch Request Packet.
1547///
1548/// Used for It is sent by server to request client to switch to Old Password Authentication
1549/// if `CLIENT_PLUGIN_AUTH` capability is not supported (by either the client or the server).
1550#[derive(Debug, Clone, Eq, PartialEq)]
1551pub struct OldAuthSwitchRequest {
1552    __header: AuthSwitchRequestHeader,
1553}
1554
1555impl OldAuthSwitchRequest {
1556    pub fn new() -> Self {
1557        Self {
1558            __header: AuthSwitchRequestHeader::new(),
1559        }
1560    }
1561
1562    pub const fn auth_plugin(&self) -> AuthPlugin<'static> {
1563        AuthPlugin::MysqlOldPassword
1564    }
1565}
1566
1567impl Default for OldAuthSwitchRequest {
1568    fn default() -> Self {
1569        Self::new()
1570    }
1571}
1572
1573impl<'de> MyDeserialize<'de> for OldAuthSwitchRequest {
1574    const SIZE: Option<usize> = Some(1);
1575    type Ctx = ();
1576
1577    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1578        Ok(Self {
1579            __header: buf.parse(())?,
1580        })
1581    }
1582}
1583
1584impl MySerialize for OldAuthSwitchRequest {
1585    fn serialize(&self, buf: &mut Vec<u8>) {
1586        self.__header.serialize(&mut *buf);
1587    }
1588}
1589
1590/// Authentication Method Switch Request Packet.
1591///
1592/// If both server and client support `CLIENT_PLUGIN_AUTH` capability, server can send this packet
1593/// to ask client to use another authentication method.
1594#[derive(Debug, Clone, Eq, PartialEq)]
1595pub struct AuthSwitchRequest<'a> {
1596    __header: AuthSwitchRequestHeader,
1597    auth_plugin: RawBytes<'a, NullBytes>,
1598    plugin_data: RawBytes<'a, EofBytes>,
1599}
1600
1601impl<'a> AuthSwitchRequest<'a> {
1602    pub fn new(
1603        auth_plugin: impl Into<Cow<'a, [u8]>>,
1604        plugin_data: impl Into<Cow<'a, [u8]>>,
1605    ) -> Self {
1606        Self {
1607            __header: AuthSwitchRequestHeader::new(),
1608            auth_plugin: RawBytes::new(auth_plugin),
1609            plugin_data: RawBytes::new(plugin_data),
1610        }
1611    }
1612
1613    pub fn auth_plugin(&self) -> AuthPlugin<'_> {
1614        ParseBuf(self.auth_plugin.as_bytes())
1615            .parse(())
1616            .expect("infallible")
1617    }
1618
1619    pub fn plugin_data(&self) -> &[u8] {
1620        match self.plugin_data.as_bytes() {
1621            [head @ .., 0] => head,
1622            all => all,
1623        }
1624    }
1625
1626    pub fn into_owned(self) -> AuthSwitchRequest<'static> {
1627        AuthSwitchRequest {
1628            __header: self.__header,
1629            auth_plugin: self.auth_plugin.into_owned(),
1630            plugin_data: self.plugin_data.into_owned(),
1631        }
1632    }
1633}
1634
1635impl<'de> MyDeserialize<'de> for AuthSwitchRequest<'de> {
1636    const SIZE: Option<usize> = None;
1637    type Ctx = ();
1638
1639    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1640        Ok(Self {
1641            __header: buf.parse(())?,
1642            auth_plugin: buf.parse(())?,
1643            plugin_data: buf.parse(())?,
1644        })
1645    }
1646}
1647
1648impl MySerialize for AuthSwitchRequest<'_> {
1649    fn serialize(&self, buf: &mut Vec<u8>) {
1650        self.__header.serialize(&mut *buf);
1651        self.auth_plugin.serialize(&mut *buf);
1652        self.plugin_data.serialize(buf);
1653    }
1654}
1655
1656// Parses and verifies Parsec additional exchange reply packet, including the ext-salt and iteration count.
1657pub fn parse_parsec_salt(data: &[u8]) -> Option<(u32, &[u8; 18])> {
1658    if data.len() != 20 {
1659        return None;
1660    }
1661
1662    let ([algo, factor], rest) = data.split_first_chunk::<2>().expect("infallible");
1663    // It used to be hardocoded 3 in the server, but that is gonna change to be configurable(in the server).
1664    // The limit 20 is rather practical - it should be more than enough for good security and yet the
1665    // calculation won't take ages.
1666    if *algo != b'P' || *factor > 20 {
1667        return None;
1668    }
1669
1670    Some((1024 << factor, rest.try_into().expect("infallible")))
1671}
1672
1673/// Represents MySql's initial handshake packet.
1674#[derive(Debug, Clone, Eq, PartialEq)]
1675pub struct HandshakePacket<'a> {
1676    protocol_version: RawInt<u8>,
1677    server_version: RawBytes<'a, NullBytes>,
1678    connection_id: RawInt<LeU32>,
1679    scramble_1: [u8; 8],
1680    __filler: Skip<1>,
1681    // lower 16 bytes
1682    capabilities_1: Const<CapabilityFlags, LeU32LowerHalf>,
1683    default_collation: RawInt<u8>,
1684    status_flags: Const<StatusFlags, LeU16>,
1685    // upper 16 bytes
1686    capabilities_2: Const<CapabilityFlags, LeU32UpperHalf>,
1687    auth_plugin_data_len: RawInt<u8>,
1688    __reserved: Skip<6>,
1689    // MariaDB uses last 4 reserved bytes to pass its extended capabilities.
1690    mariadb_ext_capabilities: Const<MariadbCapabilities, LeU32>,
1691    scramble_2: Option<RawBytes<'a, BareBytes<{ (u8::MAX as usize) - 8 }>>>,
1692    auth_plugin_name: Option<RawBytes<'a, NullBytes>>,
1693}
1694
1695impl<'de> MyDeserialize<'de> for HandshakePacket<'de> {
1696    const SIZE: Option<usize> = None;
1697    type Ctx = ();
1698
1699    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
1700        let protocol_version = buf.parse(())?;
1701        let server_version = buf.parse(())?;
1702
1703        // includes trailing 10 bytes filler
1704        let mut sbuf: ParseBuf<'_> = buf.parse(31)?;
1705        let connection_id = sbuf.parse_unchecked(())?;
1706        let scramble_1 = sbuf.parse_unchecked(())?;
1707        let __filler = sbuf.parse_unchecked(())?;
1708        let capabilities_1: RawConst<LeU32LowerHalf, CapabilityFlags> = sbuf.parse_unchecked(())?;
1709        let default_collation = sbuf.parse_unchecked(())?;
1710        let status_flags = sbuf.parse_unchecked(())?;
1711        let capabilities_2: RawConst<LeU32UpperHalf, CapabilityFlags> = sbuf.parse_unchecked(())?;
1712        let auth_plugin_data_len: RawInt<u8> = sbuf.parse_unchecked(())?;
1713        let __reserved = sbuf.parse_unchecked(())?;
1714        // If the server is MariaDB, it will pass its extended capabilities
1715        // in the last 4 reserved bytes.
1716        let mariadb_capabilities: RawConst<LeU32, MariadbCapabilities> = sbuf.parse_unchecked(())?;
1717        let mut scramble_2 = None;
1718        if capabilities_1.0 & CapabilityFlags::CLIENT_SECURE_CONNECTION.bits() > 0 {
1719            let len = max(13, auth_plugin_data_len.0 as i8 - 8) as usize;
1720            scramble_2 = buf.parse(len).map(Some)?;
1721        }
1722        let mut auth_plugin_name = None;
1723        if capabilities_2.0 & CapabilityFlags::CLIENT_PLUGIN_AUTH.bits() > 0 {
1724            auth_plugin_name = match buf.eat_all() {
1725                [head @ .., 0] => Some(RawBytes::new(head)),
1726                // missing trailing `0` is a known bug in mysql
1727                all => Some(RawBytes::new(all)),
1728            }
1729        }
1730
1731        Ok(Self {
1732            protocol_version,
1733            server_version,
1734            connection_id,
1735            scramble_1,
1736            __filler,
1737            capabilities_1: Const::new(CapabilityFlags::from_bits_truncate(capabilities_1.0)),
1738            default_collation,
1739            status_flags,
1740            capabilities_2: Const::new(CapabilityFlags::from_bits_truncate(capabilities_2.0)),
1741            auth_plugin_data_len,
1742            __reserved,
1743            mariadb_ext_capabilities: Const::new(MariadbCapabilities::from_bits_truncate(
1744                mariadb_capabilities.0,
1745            )),
1746            scramble_2,
1747            auth_plugin_name,
1748        })
1749    }
1750}
1751
1752impl MySerialize for HandshakePacket<'_> {
1753    fn serialize(&self, buf: &mut Vec<u8>) {
1754        self.protocol_version.serialize(&mut *buf);
1755        self.server_version.serialize(&mut *buf);
1756        self.connection_id.serialize(&mut *buf);
1757        self.scramble_1.serialize(&mut *buf);
1758        buf.put_u8(0x00);
1759        self.capabilities_1.serialize(&mut *buf);
1760        self.default_collation.serialize(&mut *buf);
1761        self.status_flags.serialize(&mut *buf);
1762        self.capabilities_2.serialize(&mut *buf);
1763
1764        if self
1765            .capabilities_2
1766            .contains(CapabilityFlags::CLIENT_PLUGIN_AUTH)
1767        {
1768            buf.put_u8(
1769                self.scramble_2
1770                    .as_ref()
1771                    .map(|x| (x.len() + 8) as u8)
1772                    .unwrap_or_default(),
1773            );
1774        } else {
1775            buf.put_u8(0);
1776        }
1777
1778        self.__reserved.serialize(&mut *buf);
1779        self.mariadb_ext_capabilities.serialize(&mut *buf);
1780
1781        // Assume that the packet is well formed:
1782        // * the CLIENT_SECURE_CONNECTION is set.
1783        if let Some(scramble_2) = &self.scramble_2 {
1784            scramble_2.serialize(&mut *buf);
1785        }
1786
1787        // Assume that the packet is well formed:
1788        // * the CLIENT_PLUGIN_AUTH is set.
1789        if let Some(client_plugin_auth) = &self.auth_plugin_name {
1790            client_plugin_auth.serialize(buf);
1791        }
1792    }
1793}
1794
1795impl<'a> HandshakePacket<'a> {
1796    #[allow(clippy::too_many_arguments)]
1797    pub fn new(
1798        protocol_version: u8,
1799        server_version: impl Into<Cow<'a, [u8]>>,
1800        connection_id: u32,
1801        scramble_1: [u8; 8],
1802        scramble_2: Option<impl Into<Cow<'a, [u8]>>>,
1803        capabilities: CapabilityFlags,
1804        default_collation: u8,
1805        status_flags: StatusFlags,
1806        auth_plugin_name: Option<impl Into<Cow<'a, [u8]>>>,
1807    ) -> Self {
1808        // Safety:
1809        // * capabilities are given as a valid CapabilityFlags instance
1810        // * the BitAnd operation can't set new bits
1811        let (capabilities_1, capabilities_2) = (
1812            CapabilityFlags::from_bits_retain(capabilities.bits() & 0x0000_FFFF),
1813            CapabilityFlags::from_bits_retain(capabilities.bits() & 0xFFFF_0000),
1814        );
1815
1816        let scramble_2 = scramble_2.map(RawBytes::new);
1817
1818        HandshakePacket {
1819            protocol_version: RawInt::new(protocol_version),
1820            server_version: RawBytes::new(server_version),
1821            connection_id: RawInt::new(connection_id),
1822            scramble_1,
1823            __filler: Skip,
1824            capabilities_1: Const::new(capabilities_1),
1825            default_collation: RawInt::new(default_collation),
1826            status_flags: Const::new(status_flags),
1827            capabilities_2: Const::new(capabilities_2),
1828            auth_plugin_data_len: RawInt::new(
1829                scramble_2
1830                    .as_ref()
1831                    .map(|x| x.len() as u8)
1832                    .unwrap_or_default(),
1833            ),
1834            __reserved: Skip,
1835            mariadb_ext_capabilities: Const::new(MariadbCapabilities::empty()),
1836            scramble_2,
1837            auth_plugin_name: auth_plugin_name.map(RawBytes::new),
1838        }
1839    }
1840
1841    pub fn with_mariadb_ext_capabilities(mut self, flags: MariadbCapabilities) -> Self {
1842        self.mariadb_ext_capabilities = Const::new(flags);
1843        self
1844    }
1845
1846    pub fn into_owned(self) -> HandshakePacket<'static> {
1847        HandshakePacket {
1848            protocol_version: self.protocol_version,
1849            server_version: self.server_version.into_owned(),
1850            connection_id: self.connection_id,
1851            scramble_1: self.scramble_1,
1852            __filler: self.__filler,
1853            capabilities_1: self.capabilities_1,
1854            default_collation: self.default_collation,
1855            status_flags: self.status_flags,
1856            capabilities_2: self.capabilities_2,
1857            auth_plugin_data_len: self.auth_plugin_data_len,
1858            __reserved: self.__reserved,
1859            mariadb_ext_capabilities: self.mariadb_ext_capabilities,
1860            scramble_2: self.scramble_2.map(|x| x.into_owned()),
1861            auth_plugin_name: self.auth_plugin_name.map(RawBytes::into_owned),
1862        }
1863    }
1864
1865    /// Value of the protocol_version field of an initial handshake packet.
1866    pub fn protocol_version(&self) -> u8 {
1867        self.protocol_version.0
1868    }
1869
1870    /// Value of the server_version field of an initial handshake packet as a byte slice.
1871    pub fn server_version_ref(&self) -> &[u8] {
1872        self.server_version.as_bytes()
1873    }
1874
1875    /// Value of the server_version field of an initial handshake packet as a string
1876    /// (lossy converted).
1877    pub fn server_version_str(&self) -> Cow<'_, str> {
1878        self.server_version.as_str()
1879    }
1880
1881    /// Parsed server version.
1882    ///
1883    /// Will parse first \d+.\d+.\d+ of a server version string (if any).
1884    pub fn server_version_parsed(&self) -> Option<(u16, u16, u16)> {
1885        VERSION_RE
1886            .captures(self.server_version_ref())
1887            .map(|captures| {
1888                // Should not panic because validated with regex
1889                (
1890                    btoi::<u16>(captures.get(1).unwrap().as_bytes()).unwrap(),
1891                    btoi::<u16>(captures.get(2).unwrap().as_bytes()).unwrap(),
1892                    btoi::<u16>(captures.get(3).unwrap().as_bytes()).unwrap(),
1893                )
1894            })
1895    }
1896
1897    /// Parsed mariadb server version.
1898    pub fn maria_db_server_version_parsed(&self) -> Option<(u16, u16, u16)> {
1899        MARIADB_VERSION_RE
1900            .captures(self.server_version_ref())
1901            .map(|captures| {
1902                // Should not panic because validated with regex
1903                (
1904                    btoi::<u16>(captures.get(1).unwrap().as_bytes()).unwrap(),
1905                    btoi::<u16>(captures.get(2).unwrap().as_bytes()).unwrap(),
1906                    btoi::<u16>(captures.get(3).unwrap().as_bytes()).unwrap(),
1907                )
1908            })
1909    }
1910
1911    /// Value of the connection_id field of an initial handshake packet.
1912    pub fn connection_id(&self) -> u32 {
1913        self.connection_id.0
1914    }
1915
1916    /// Value of the scramble_1 field of an initial handshake packet as a byte slice.
1917    pub fn scramble_1_ref(&self) -> &[u8] {
1918        self.scramble_1.as_ref()
1919    }
1920
1921    /// Value of the scramble_2 field of an initial handshake packet as a byte slice.
1922    ///
1923    /// Note that this may include a terminating null character.
1924    pub fn scramble_2_ref(&self) -> Option<&[u8]> {
1925        self.scramble_2.as_ref().map(|x| x.as_bytes())
1926    }
1927
1928    /// Returns concatenated auth plugin nonce.
1929    pub fn nonce(&self) -> Vec<u8> {
1930        let mut out = Vec::from(self.scramble_1_ref());
1931        out.extend_from_slice(self.scramble_2_ref().unwrap_or(&[][..]));
1932
1933        // Trim zero terminator. Fill with zeroes if nonce
1934        // is somehow smaller than 20 bytes.
1935        out.resize(20, 0);
1936        out
1937    }
1938
1939    /// Value of a server capabilities.
1940    pub fn capabilities(&self) -> CapabilityFlags {
1941        self.capabilities_1.0 | self.capabilities_2.0
1942    }
1943
1944    /// Value of MariaDB specific server capabilities
1945    pub fn mariadb_ext_capabilities(&self) -> MariadbCapabilities {
1946        self.mariadb_ext_capabilities.0
1947    }
1948    /// Value of the default_collation field of an initial handshake packet.
1949    pub fn default_collation(&self) -> u8 {
1950        self.default_collation.0
1951    }
1952
1953    /// Value of a status flags.
1954    pub fn status_flags(&self) -> StatusFlags {
1955        self.status_flags.0
1956    }
1957
1958    /// Value of the auth_plugin_name field of an initial handshake packet as a byte slice.
1959    pub fn auth_plugin_name_ref(&self) -> Option<&[u8]> {
1960        self.auth_plugin_name.as_ref().map(|x| x.as_bytes())
1961    }
1962
1963    /// Value of the auth_plugin_name field of an initial handshake packet as a string
1964    /// (lossy converted).
1965    pub fn auth_plugin_name_str(&self) -> Option<Cow<'_, str>> {
1966        self.auth_plugin_name.as_ref().map(|x| x.as_str())
1967    }
1968
1969    /// Auth plugin of a handshake packet
1970    pub fn auth_plugin(&self) -> Option<AuthPlugin<'_>> {
1971        self.auth_plugin_name.as_ref().map(|x| match x.as_bytes() {
1972            [name @ .., 0] => ParseBuf(name).parse_unchecked(()).expect("infallible"),
1973            all => ParseBuf(all).parse_unchecked(()).expect("infallible"),
1974        })
1975    }
1976}
1977
1978define_header!(
1979    ComChangeUserHeader,
1980    InvalidComChangeUserHeader("Invalid COM_CHANGE_USER header"),
1981    0x11
1982);
1983
1984#[derive(Debug, Clone, PartialEq, Eq)]
1985pub struct ComChangeUser<'a> {
1986    __header: ComChangeUserHeader,
1987    user: RawBytes<'a, NullBytes>,
1988    // Only CLIENT_SECURE_CONNECTION capable servers are supported
1989    auth_plugin_data: RawBytes<'a, U8Bytes>,
1990    database: RawBytes<'a, NullBytes>,
1991    more_data: Option<ComChangeUserMoreData<'a>>,
1992}
1993
1994impl<'a> ComChangeUser<'a> {
1995    pub fn new() -> Self {
1996        Self {
1997            __header: ComChangeUserHeader::new(),
1998            user: Default::default(),
1999            auth_plugin_data: Default::default(),
2000            database: Default::default(),
2001            more_data: None,
2002        }
2003    }
2004
2005    pub fn with_user(mut self, user: Option<impl Into<Cow<'a, [u8]>>>) -> Self {
2006        self.user = user.map(RawBytes::new).unwrap_or_default();
2007        self
2008    }
2009
2010    pub fn with_database(mut self, database: Option<impl Into<Cow<'a, [u8]>>>) -> Self {
2011        self.database = database.map(RawBytes::new).unwrap_or_default();
2012        self
2013    }
2014
2015    pub fn with_auth_plugin_data(
2016        mut self,
2017        auth_plugin_data: Option<impl Into<Cow<'a, [u8]>>>,
2018    ) -> Self {
2019        self.auth_plugin_data = auth_plugin_data.map(RawBytes::new).unwrap_or_default();
2020        self
2021    }
2022
2023    pub fn with_more_data(mut self, more_data: Option<ComChangeUserMoreData<'a>>) -> Self {
2024        self.more_data = more_data;
2025        self
2026    }
2027
2028    pub fn into_owned(self) -> ComChangeUser<'static> {
2029        ComChangeUser {
2030            __header: self.__header,
2031            user: self.user.into_owned(),
2032            auth_plugin_data: self.auth_plugin_data.into_owned(),
2033            database: self.database.into_owned(),
2034            more_data: self.more_data.map(|x| x.into_owned()),
2035        }
2036    }
2037}
2038
2039impl Default for ComChangeUser<'_> {
2040    fn default() -> Self {
2041        Self::new()
2042    }
2043}
2044
2045impl<'de> MyDeserialize<'de> for ComChangeUser<'de> {
2046    const SIZE: Option<usize> = None;
2047
2048    type Ctx = CapabilityFlags;
2049
2050    fn deserialize(flags: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2051        Ok(Self {
2052            __header: buf.parse(())?,
2053            user: buf.parse(())?,
2054            auth_plugin_data: buf.parse(())?,
2055            database: buf.parse(())?,
2056            more_data: if !buf.is_empty() {
2057                Some(buf.parse(flags)?)
2058            } else {
2059                None
2060            },
2061        })
2062    }
2063}
2064
2065impl MySerialize for ComChangeUser<'_> {
2066    fn serialize(&self, buf: &mut Vec<u8>) {
2067        self.__header.serialize(&mut *buf);
2068        self.user.serialize(&mut *buf);
2069        self.auth_plugin_data.serialize(&mut *buf);
2070        self.database.serialize(&mut *buf);
2071        if let Some(ref more_data) = self.more_data {
2072            more_data.serialize(&mut *buf);
2073        }
2074    }
2075}
2076
2077#[derive(Debug, Clone, PartialEq, Eq)]
2078pub struct ComChangeUserMoreData<'a> {
2079    character_set: RawInt<LeU16>,
2080    auth_plugin: Option<AuthPlugin<'a>>,
2081    connect_attributes: Option<HashMap<RawBytes<'a, LenEnc>, RawBytes<'a, LenEnc>>>,
2082}
2083
2084impl<'a> ComChangeUserMoreData<'a> {
2085    pub fn new(character_set: u16) -> Self {
2086        Self {
2087            character_set: RawInt::new(character_set),
2088            auth_plugin: None,
2089            connect_attributes: None,
2090        }
2091    }
2092
2093    pub fn with_auth_plugin(mut self, auth_plugin: Option<AuthPlugin<'a>>) -> Self {
2094        self.auth_plugin = auth_plugin;
2095        self
2096    }
2097
2098    pub fn with_connect_attributes(
2099        mut self,
2100        connect_attributes: Option<HashMap<String, String>>,
2101    ) -> Self {
2102        self.connect_attributes = connect_attributes.map(|attrs| {
2103            attrs
2104                .into_iter()
2105                .map(|(k, v)| (RawBytes::new(k.into_bytes()), RawBytes::new(v.into_bytes())))
2106                .collect()
2107        });
2108        self
2109    }
2110
2111    pub fn into_owned(self) -> ComChangeUserMoreData<'static> {
2112        ComChangeUserMoreData {
2113            character_set: self.character_set,
2114            auth_plugin: self.auth_plugin.map(|x| x.into_owned()),
2115            connect_attributes: self.connect_attributes.map(|x| {
2116                x.into_iter()
2117                    .map(|(k, v)| (k.into_owned(), v.into_owned()))
2118                    .collect()
2119            }),
2120        }
2121    }
2122}
2123
2124// Helper that deserializes connect attributes.
2125fn deserialize_connect_attrs<'de>(
2126    buf: &mut ParseBuf<'de>,
2127) -> io::Result<HashMap<RawBytes<'de, LenEnc>, RawBytes<'de, LenEnc>>> {
2128    let data_len = buf.parse::<RawInt<LenEnc>>(())?;
2129    let mut data: ParseBuf<'_> = buf.parse(data_len.0 as usize)?;
2130    let mut attrs = HashMap::new();
2131    while !data.is_empty() {
2132        let key = data.parse::<RawBytes<'_, LenEnc>>(())?;
2133        let value = data.parse::<RawBytes<'_, LenEnc>>(())?;
2134        attrs.insert(key, value);
2135    }
2136    Ok(attrs)
2137}
2138
2139// Helper that serializes connect attributes.
2140fn serialize_connect_attrs<'a>(
2141    connect_attributes: &HashMap<RawBytes<'a, LenEnc>, RawBytes<'a, LenEnc>>,
2142    buf: &mut Vec<u8>,
2143) {
2144    let len = connect_attributes
2145        .iter()
2146        .map(|(k, v)| lenenc_str_len(k.as_bytes()) + lenenc_str_len(v.as_bytes()))
2147        .sum::<u64>();
2148    buf.put_lenenc_int(len);
2149
2150    for (name, value) in connect_attributes {
2151        name.serialize(&mut *buf);
2152        value.serialize(&mut *buf);
2153    }
2154}
2155
2156impl<'de> MyDeserialize<'de> for ComChangeUserMoreData<'de> {
2157    const SIZE: Option<usize> = None;
2158    type Ctx = CapabilityFlags;
2159
2160    fn deserialize(flags: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2161        // always assume CLIENT_PROTOCOL_41
2162        let character_set = buf.parse(())?;
2163        let mut auth_plugin = None;
2164        let mut connect_attributes = None;
2165
2166        if flags.contains(CapabilityFlags::CLIENT_PLUGIN_AUTH) {
2167            // plugin name is null-terminated here
2168            match buf.parse::<RawBytes<'_, NullBytes>>(())?.0 {
2169                Cow::Borrowed(bytes) => {
2170                    let mut auth_plugin_buf = ParseBuf(bytes);
2171                    auth_plugin = Some(auth_plugin_buf.parse(())?);
2172                }
2173                _ => unreachable!(),
2174            }
2175        };
2176
2177        if flags.contains(CapabilityFlags::CLIENT_CONNECT_ATTRS) {
2178            connect_attributes = Some(deserialize_connect_attrs(&mut *buf)?);
2179        };
2180
2181        Ok(Self {
2182            character_set,
2183            auth_plugin,
2184            connect_attributes,
2185        })
2186    }
2187}
2188
2189impl MySerialize for ComChangeUserMoreData<'_> {
2190    fn serialize(&self, buf: &mut Vec<u8>) {
2191        self.character_set.serialize(&mut *buf);
2192        if let Some(ref auth_plugin) = self.auth_plugin {
2193            auth_plugin.serialize(&mut *buf);
2194        }
2195        if let Some(ref connect_attributes) = self.connect_attributes {
2196            serialize_connect_attrs(connect_attributes, buf);
2197        } else {
2198            // We'll always act like CLIENT_CONNECT_ATTRS is set,
2199            // this is to avoid looking into the actual connection flags.
2200            serialize_connect_attrs(&Default::default(), buf);
2201        }
2202    }
2203}
2204
2205/// Actual serialization of this field depends on capability flags values.
2206type ScrambleBuf<'a> =
2207    Either<RawBytes<'a, LenEnc>, Either<RawBytes<'a, U8Bytes>, RawBytes<'a, NullBytes>>>;
2208
2209#[derive(Debug, Clone, PartialEq, Eq)]
2210pub struct HandshakeResponse<'a> {
2211    capabilities: Const<CapabilityFlags, LeU32>,
2212    max_packet_size: RawInt<LeU32>,
2213    collation: RawInt<u8>,
2214    scramble_buf: ScrambleBuf<'a>,
2215    user: RawBytes<'a, NullBytes>,
2216    db_name: Option<RawBytes<'a, NullBytes>>,
2217    auth_plugin: Option<AuthPlugin<'a>>,
2218    connect_attributes: Option<HashMap<RawBytes<'a, LenEnc>, RawBytes<'a, LenEnc>>>,
2219    mariadb_ext_capabilities: Const<MariadbCapabilities, LeU32>,
2220}
2221
2222impl<'a> HandshakeResponse<'a> {
2223    #[allow(clippy::too_many_arguments)]
2224    pub fn new(
2225        scramble_buf: Option<impl Into<Cow<'a, [u8]>>>,
2226        server_version: (u16, u16, u16),
2227        user: Option<impl Into<Cow<'a, [u8]>>>,
2228        db_name: Option<impl Into<Cow<'a, [u8]>>>,
2229        auth_plugin: Option<AuthPlugin<'a>>,
2230        mut capabilities: CapabilityFlags,
2231        connect_attributes: Option<HashMap<String, String>>,
2232        max_packet_size: u32,
2233    ) -> Self {
2234        let scramble_buf =
2235            if capabilities.contains(CapabilityFlags::CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA) {
2236                Either::Left(RawBytes::new(
2237                    scramble_buf.map(Into::into).unwrap_or_default(),
2238                ))
2239            } else if capabilities.contains(CapabilityFlags::CLIENT_SECURE_CONNECTION) {
2240                Either::Right(Either::Left(RawBytes::new(
2241                    scramble_buf.map(Into::into).unwrap_or_default(),
2242                )))
2243            } else {
2244                Either::Right(Either::Right(RawBytes::new(
2245                    scramble_buf.map(Into::into).unwrap_or_default(),
2246                )))
2247            };
2248
2249        if db_name.is_some() {
2250            capabilities.insert(CapabilityFlags::CLIENT_CONNECT_WITH_DB);
2251        } else {
2252            capabilities.remove(CapabilityFlags::CLIENT_CONNECT_WITH_DB);
2253        }
2254
2255        if auth_plugin.is_some() {
2256            capabilities.insert(CapabilityFlags::CLIENT_PLUGIN_AUTH);
2257        } else {
2258            capabilities.remove(CapabilityFlags::CLIENT_PLUGIN_AUTH);
2259        }
2260
2261        if connect_attributes.is_some() {
2262            capabilities.insert(CapabilityFlags::CLIENT_CONNECT_ATTRS);
2263        } else {
2264            capabilities.remove(CapabilityFlags::CLIENT_CONNECT_ATTRS);
2265        }
2266
2267        Self {
2268            scramble_buf,
2269            collation: if server_version >= (5, 5, 3) {
2270                RawInt::new(CollationId::UTF8MB4_GENERAL_CI as u8)
2271            } else {
2272                RawInt::new(CollationId::UTF8MB3_GENERAL_CI as u8)
2273            },
2274            user: user.map(RawBytes::new).unwrap_or_default(),
2275            db_name: db_name.map(RawBytes::new),
2276            auth_plugin,
2277            capabilities: Const::new(capabilities),
2278            connect_attributes: connect_attributes.map(|attrs| {
2279                attrs
2280                    .into_iter()
2281                    .map(|(k, v)| (RawBytes::new(k.into_bytes()), RawBytes::new(v.into_bytes())))
2282                    .collect()
2283            }),
2284            max_packet_size: RawInt::new(max_packet_size),
2285            mariadb_ext_capabilities: Const::new(MariadbCapabilities::empty()),
2286        }
2287    }
2288
2289    pub fn with_mariadb_ext_capabilities(
2290        mut self,
2291        mariadb_ext_capabilities: MariadbCapabilities,
2292    ) -> Self {
2293        self.mariadb_ext_capabilities = Const::new(mariadb_ext_capabilities);
2294        self
2295    }
2296
2297    pub fn capabilities(&self) -> CapabilityFlags {
2298        self.capabilities.0
2299    }
2300
2301    pub fn mariadb_ext_capabilities(&self) -> MariadbCapabilities {
2302        self.mariadb_ext_capabilities.0
2303    }
2304
2305    pub fn collation(&self) -> u8 {
2306        self.collation.0
2307    }
2308
2309    pub fn scramble_buf(&self) -> &[u8] {
2310        match &self.scramble_buf {
2311            Either::Left(x) => x.as_bytes(),
2312            Either::Right(x) => match x {
2313                Either::Left(x) => x.as_bytes(),
2314                Either::Right(x) => x.as_bytes(),
2315            },
2316        }
2317    }
2318
2319    pub fn user(&self) -> &[u8] {
2320        self.user.as_bytes()
2321    }
2322
2323    pub fn db_name(&self) -> Option<&[u8]> {
2324        self.db_name.as_ref().map(|x| x.as_bytes())
2325    }
2326
2327    pub fn auth_plugin(&self) -> Option<&AuthPlugin<'a>> {
2328        self.auth_plugin.as_ref()
2329    }
2330
2331    #[must_use = "entails computation"]
2332    pub fn connect_attributes(&self) -> Option<HashMap<String, String>> {
2333        self.connect_attributes.as_ref().map(|attrs| {
2334            attrs
2335                .iter()
2336                .map(|(k, v)| (k.as_str().into_owned(), v.as_str().into_owned()))
2337                .collect()
2338        })
2339    }
2340}
2341
2342impl<'de> MyDeserialize<'de> for HandshakeResponse<'de> {
2343    const SIZE: Option<usize> = None;
2344    type Ctx = ();
2345
2346    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2347        let mut sbuf: ParseBuf<'_> = buf.parse(4 + 4 + 1 + 23)?;
2348        let client_flags: RawConst<LeU32, CapabilityFlags> = sbuf.parse_unchecked(())?;
2349        let max_packet_size: RawInt<LeU32> = sbuf.parse_unchecked(())?;
2350        let collation = sbuf.parse_unchecked(())?;
2351        sbuf.parse_unchecked::<Skip<19>>(())?;
2352        let mariadb_flags: RawConst<LeU32, MariadbCapabilities> = sbuf.parse_unchecked(())?;
2353
2354        let user = buf.parse(())?;
2355        let scramble_buf =
2356            if client_flags.0 & CapabilityFlags::CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA.bits() > 0 {
2357                Either::Left(buf.parse(())?)
2358            } else if client_flags.0 & CapabilityFlags::CLIENT_SECURE_CONNECTION.bits() > 0 {
2359                Either::Right(Either::Left(buf.parse(())?))
2360            } else {
2361                Either::Right(Either::Right(buf.parse(())?))
2362            };
2363
2364        let mut db_name = None;
2365        if client_flags.0 & CapabilityFlags::CLIENT_CONNECT_WITH_DB.bits() > 0 {
2366            db_name = buf.parse(()).map(Some)?;
2367        }
2368
2369        let mut auth_plugin = None;
2370        if client_flags.0 & CapabilityFlags::CLIENT_PLUGIN_AUTH.bits() > 0 {
2371            let auth_plugin_name = buf.eat_null_str();
2372            auth_plugin = Some(AuthPlugin::from_bytes(auth_plugin_name));
2373        }
2374
2375        let mut connect_attributes = None;
2376        if client_flags.0 & CapabilityFlags::CLIENT_CONNECT_ATTRS.bits() > 0 {
2377            connect_attributes = Some(deserialize_connect_attrs(&mut *buf)?);
2378        }
2379
2380        Ok(Self {
2381            capabilities: Const::new(CapabilityFlags::from_bits_truncate(client_flags.0)),
2382            max_packet_size,
2383            collation,
2384            scramble_buf,
2385            user,
2386            db_name,
2387            auth_plugin,
2388            connect_attributes,
2389            mariadb_ext_capabilities: Const::new(MariadbCapabilities::from_bits_truncate(
2390                mariadb_flags.0,
2391            )),
2392        })
2393    }
2394}
2395
2396impl MySerialize for HandshakeResponse<'_> {
2397    fn serialize(&self, buf: &mut Vec<u8>) {
2398        self.capabilities.serialize(&mut *buf);
2399        self.max_packet_size.serialize(&mut *buf);
2400        self.collation.serialize(&mut *buf);
2401        buf.put_slice(&[0; 19]);
2402        self.mariadb_ext_capabilities.serialize(&mut *buf);
2403        self.user.serialize(&mut *buf);
2404        self.scramble_buf.serialize(&mut *buf);
2405
2406        if let Some(db_name) = &self.db_name {
2407            db_name.serialize(&mut *buf);
2408        }
2409
2410        if let Some(auth_plugin) = &self.auth_plugin {
2411            auth_plugin.serialize(&mut *buf);
2412        }
2413
2414        if let Some(attrs) = &self.connect_attributes {
2415            let len = attrs
2416                .iter()
2417                .map(|(k, v)| lenenc_str_len(k.as_bytes()) + lenenc_str_len(v.as_bytes()))
2418                .sum::<u64>();
2419            buf.put_lenenc_int(len);
2420
2421            for (name, value) in attrs {
2422                name.serialize(&mut *buf);
2423                value.serialize(&mut *buf);
2424            }
2425        }
2426    }
2427}
2428
2429#[derive(Debug, Clone, Eq, PartialEq)]
2430pub struct SslRequest {
2431    capabilities: Const<CapabilityFlags, LeU32>,
2432    max_packet_size: RawInt<LeU32>,
2433    character_set: RawInt<u8>,
2434    __skip: Skip<19>,
2435    mariadb_capabilities: Const<MariadbCapabilities, LeU32>,
2436}
2437
2438impl SslRequest {
2439    pub fn new(capabilities: CapabilityFlags, max_packet_size: u32, character_set: u8) -> Self {
2440        Self {
2441            capabilities: Const::new(capabilities),
2442            max_packet_size: RawInt::new(max_packet_size),
2443            character_set: RawInt::new(character_set),
2444            __skip: Skip,
2445            mariadb_capabilities: Const::new(MariadbCapabilities::empty()),
2446        }
2447    }
2448
2449    /// This only makes sense for MariaDb server (since v10.2?).
2450    ///
2451    /// Capabilities are zeroed by default.
2452    pub fn with_mariadb_capabilities(mut self, mariadb_capabilities: MariadbCapabilities) -> Self {
2453        self.mariadb_capabilities.0 = mariadb_capabilities;
2454        self
2455    }
2456
2457    pub fn capabilities(&self) -> CapabilityFlags {
2458        self.capabilities.0
2459    }
2460
2461    pub fn max_packet_size(&self) -> u32 {
2462        self.max_packet_size.0
2463    }
2464
2465    pub fn character_set(&self) -> u8 {
2466        self.character_set.0
2467    }
2468}
2469
2470impl<'de> MyDeserialize<'de> for SslRequest {
2471    const SIZE: Option<usize> = Some(4 + 4 + 1 + 23);
2472    type Ctx = ();
2473
2474    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2475        let mut buf: ParseBuf<'_> = buf.parse(Self::SIZE.unwrap())?;
2476        let raw_capabilities = buf.parse_unchecked::<RawConst<LeU32, CapabilityFlags>>(())?;
2477
2478        let capabilities = Const::new(CapabilityFlags::from_bits_truncate(raw_capabilities.0));
2479        let max_packet_size = buf.parse_unchecked(())?;
2480        let character_set = buf.parse_unchecked(())?;
2481        let __skip = buf.parse_unchecked(())?;
2482
2483        let raw_mariadb_capabilities =
2484            buf.parse_unchecked::<RawConst<LeU32, MariadbCapabilities>>(())?;
2485
2486        Ok(Self {
2487            capabilities,
2488            max_packet_size,
2489            character_set,
2490            __skip,
2491            mariadb_capabilities: Const::new(MariadbCapabilities::from_bits_truncate(
2492                raw_mariadb_capabilities.0,
2493            )),
2494        })
2495    }
2496}
2497
2498impl MySerialize for SslRequest {
2499    fn serialize(&self, buf: &mut Vec<u8>) {
2500        self.capabilities.serialize(&mut *buf);
2501        self.max_packet_size.serialize(&mut *buf);
2502        self.character_set.serialize(&mut *buf);
2503        self.__skip.serialize(&mut *buf);
2504        self.mariadb_capabilities.serialize(&mut *buf);
2505    }
2506}
2507
2508#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, thiserror::Error)]
2509#[error("Invalid statement packet status")]
2510pub struct InvalidStmtPacketStatus;
2511
2512/// Represents MySql's statement packet.
2513#[derive(Debug, Clone, Eq, PartialEq, Hash)]
2514pub struct StmtPacket {
2515    status: ConstU8<InvalidStmtPacketStatus, 0x00>,
2516    statement_id: RawInt<LeU32>,
2517    num_columns: RawInt<LeU16>,
2518    num_params: RawInt<LeU16>,
2519    __skip: Skip<1>,
2520    warning_count: RawInt<LeU16>,
2521}
2522
2523impl<'de> MyDeserialize<'de> for StmtPacket {
2524    const SIZE: Option<usize> = Some(12);
2525    type Ctx = ();
2526
2527    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2528        let mut buf: ParseBuf<'_> = buf.parse(Self::SIZE.unwrap())?;
2529        Ok(StmtPacket {
2530            status: buf.parse_unchecked(())?,
2531            statement_id: buf.parse_unchecked(())?,
2532            num_columns: buf.parse_unchecked(())?,
2533            num_params: buf.parse_unchecked(())?,
2534            __skip: buf.parse_unchecked(())?,
2535            warning_count: buf.parse_unchecked(())?,
2536        })
2537    }
2538}
2539
2540impl MySerialize for StmtPacket {
2541    fn serialize(&self, buf: &mut Vec<u8>) {
2542        self.status.serialize(&mut *buf);
2543        self.statement_id.serialize(&mut *buf);
2544        self.num_columns.serialize(&mut *buf);
2545        self.num_params.serialize(&mut *buf);
2546        self.__skip.serialize(&mut *buf);
2547        self.warning_count.serialize(&mut *buf);
2548    }
2549}
2550
2551impl StmtPacket {
2552    /// Value of the statement_id field of a statement packet.
2553    pub fn statement_id(&self) -> u32 {
2554        *self.statement_id
2555    }
2556
2557    /// Value of the num_columns field of a statement packet.
2558    pub fn num_columns(&self) -> u16 {
2559        *self.num_columns
2560    }
2561
2562    /// Value of the num_params field of a statement packet.
2563    pub fn num_params(&self) -> u16 {
2564        *self.num_params
2565    }
2566
2567    /// Value of the warning_count field of a statement packet.
2568    pub fn warning_count(&self) -> u16 {
2569        *self.warning_count
2570    }
2571}
2572
2573/// Null-bitmap.
2574///
2575/// <http://dev.mysql.com/doc/internals/en/null-bitmap.html>
2576#[derive(Debug, Clone, Eq, PartialEq)]
2577pub struct NullBitmap<T, U: AsRef<[u8]> = Vec<u8>>(U, PhantomData<T>);
2578
2579impl<'de, T: SerializationSide> MyDeserialize<'de> for NullBitmap<T, Cow<'de, [u8]>> {
2580    const SIZE: Option<usize> = None;
2581    type Ctx = usize;
2582
2583    fn deserialize(num_columns: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
2584        let bitmap_len = Self::bitmap_len(num_columns);
2585        let bytes = buf.checked_eat(bitmap_len).ok_or_else(unexpected_buf_eof)?;
2586        Ok(Self::from_bytes(Cow::Borrowed(bytes)))
2587    }
2588}
2589
2590impl<T: SerializationSide> NullBitmap<T, Vec<u8>> {
2591    /// Creates new null-bitmap for a given number of columns.
2592    pub fn new(num_columns: usize) -> Self {
2593        Self::from_bytes(vec![0; Self::bitmap_len(num_columns)])
2594    }
2595
2596    /// Will read null-bitmap for a given number of columns from `input`.
2597    pub fn read(input: &mut &[u8], num_columns: usize) -> Self {
2598        let bitmap_len = Self::bitmap_len(num_columns);
2599        assert!(input.len() >= bitmap_len);
2600
2601        let bitmap = Self::from_bytes(input[..bitmap_len].to_vec());
2602        *input = &input[bitmap_len..];
2603
2604        bitmap
2605    }
2606}
2607
2608impl<T: SerializationSide, U: AsRef<[u8]>> NullBitmap<T, U> {
2609    pub fn bitmap_len(num_columns: usize) -> usize {
2610        (num_columns + 7 + T::BIT_OFFSET) / 8
2611    }
2612
2613    fn byte_and_bit(&self, column_index: usize) -> (usize, u8) {
2614        let offset = column_index + T::BIT_OFFSET;
2615        let byte = offset / 8;
2616        let bit = 1 << (offset % 8) as u8;
2617
2618        assert!(byte < self.0.as_ref().len());
2619
2620        (byte, bit)
2621    }
2622
2623    /// Creates new null-bitmap from given bytes.
2624    pub fn from_bytes(bytes: U) -> Self {
2625        Self(bytes, PhantomData)
2626    }
2627
2628    /// Returns `true` if given column is `NULL` in this `NullBitmap`.
2629    pub fn is_null(&self, column_index: usize) -> bool {
2630        let (byte, bit) = self.byte_and_bit(column_index);
2631        self.0.as_ref()[byte] & bit > 0
2632    }
2633}
2634
2635impl<T: SerializationSide, U: AsRef<[u8]> + AsMut<[u8]>> NullBitmap<T, U> {
2636    /// Sets flag value for given column.
2637    pub fn set(&mut self, column_index: usize, is_null: bool) {
2638        let (byte, bit) = self.byte_and_bit(column_index);
2639        if is_null {
2640            self.0.as_mut()[byte] |= bit
2641        } else {
2642            self.0.as_mut()[byte] &= !bit
2643        }
2644    }
2645}
2646
2647impl<T, U: AsRef<[u8]>> AsRef<[u8]> for NullBitmap<T, U> {
2648    fn as_ref(&self) -> &[u8] {
2649        self.0.as_ref()
2650    }
2651}
2652
2653#[derive(Debug, Clone, PartialEq)]
2654pub struct ComStmtExecuteRequestBuilder {
2655    pub stmt_id: u32,
2656}
2657
2658impl ComStmtExecuteRequestBuilder {
2659    pub const NULL_BITMAP_OFFSET: usize = 10;
2660
2661    pub fn new(stmt_id: u32) -> Self {
2662        Self { stmt_id }
2663    }
2664}
2665
2666impl ComStmtExecuteRequestBuilder {
2667    pub fn build(self, params: &[Value]) -> (ComStmtExecuteRequest<'_>, bool) {
2668        let bitmap_len = NullBitmap::<ClientSide>::bitmap_len(params.len());
2669
2670        let mut bitmap_bytes = vec![0; bitmap_len];
2671        let mut bitmap = NullBitmap::<ClientSide, _>::from_bytes(&mut bitmap_bytes);
2672        let params = params.iter().collect::<Vec<_>>();
2673
2674        let meta_len = params.len() * 2;
2675
2676        let mut data_len = 0;
2677        for (i, param) in params.iter().enumerate() {
2678            match param.bin_len() as usize {
2679                0 => bitmap.set(i, true),
2680                x => data_len += x,
2681            }
2682        }
2683
2684        let total_len = 10 + bitmap_len + 1 + meta_len + data_len;
2685
2686        let as_long_data = total_len > MAX_PAYLOAD_LEN;
2687
2688        (
2689            ComStmtExecuteRequest {
2690                com_stmt_execute: ConstU8::new(),
2691                stmt_id: RawInt::new(self.stmt_id),
2692                flags: Const::new(CursorType::CURSOR_TYPE_NO_CURSOR),
2693                iteration_count: ConstU32::new(),
2694                params_flags: Const::new(StmtExecuteParamsFlags::NEW_PARAMS_BOUND),
2695                bitmap: RawBytes::new(bitmap_bytes),
2696                params,
2697                as_long_data,
2698            },
2699            as_long_data,
2700        )
2701    }
2702}
2703
2704define_header!(
2705    ComStmtExecuteHeader,
2706    COM_STMT_EXECUTE,
2707    InvalidComStmtExecuteHeader
2708);
2709
2710define_const!(
2711    ConstU32,
2712    IterationCount,
2713    InvalidIterationCount("Invalid iteration count for COM_STMT_EXECUTE"),
2714    1
2715);
2716
2717#[derive(Debug, Clone, PartialEq)]
2718pub struct ComStmtExecuteRequest<'a> {
2719    com_stmt_execute: ComStmtExecuteHeader,
2720    stmt_id: RawInt<LeU32>,
2721    flags: Const<CursorType, u8>,
2722    iteration_count: IterationCount,
2723    // max params / bits per byte = 8192
2724    bitmap: RawBytes<'a, BareBytes<8192>>,
2725    params_flags: Const<StmtExecuteParamsFlags, u8>,
2726    params: Vec<&'a Value>,
2727    as_long_data: bool,
2728}
2729
2730impl<'a> ComStmtExecuteRequest<'a> {
2731    pub fn stmt_id(&self) -> u32 {
2732        self.stmt_id.0
2733    }
2734
2735    pub fn flags(&self) -> CursorType {
2736        self.flags.0
2737    }
2738
2739    pub fn bitmap(&self) -> &[u8] {
2740        self.bitmap.as_bytes()
2741    }
2742
2743    pub fn params_flags(&self) -> StmtExecuteParamsFlags {
2744        self.params_flags.0
2745    }
2746
2747    pub fn params(&self) -> &[&'a Value] {
2748        self.params.as_ref()
2749    }
2750
2751    pub fn as_long_data(&self) -> bool {
2752        self.as_long_data
2753    }
2754}
2755
2756impl MySerialize for ComStmtExecuteRequest<'_> {
2757    fn serialize(&self, buf: &mut Vec<u8>) {
2758        self.com_stmt_execute.serialize(&mut *buf);
2759        self.stmt_id.serialize(&mut *buf);
2760        self.flags.serialize(&mut *buf);
2761        self.iteration_count.serialize(&mut *buf);
2762
2763        if !self.params.is_empty() {
2764            self.bitmap.serialize(&mut *buf);
2765            self.params_flags.serialize(&mut *buf);
2766        }
2767
2768        for param in &self.params {
2769            let (column_type, flags) = match param {
2770                Value::NULL => (ColumnType::MYSQL_TYPE_NULL, StmtExecuteParamFlags::empty()),
2771                Value::Bytes(_) => (
2772                    ColumnType::MYSQL_TYPE_VAR_STRING,
2773                    StmtExecuteParamFlags::empty(),
2774                ),
2775                Value::Int(_) => (
2776                    ColumnType::MYSQL_TYPE_LONGLONG,
2777                    StmtExecuteParamFlags::empty(),
2778                ),
2779                Value::UInt(_) => (
2780                    ColumnType::MYSQL_TYPE_LONGLONG,
2781                    StmtExecuteParamFlags::UNSIGNED,
2782                ),
2783                Value::Float(_) => (ColumnType::MYSQL_TYPE_FLOAT, StmtExecuteParamFlags::empty()),
2784                Value::Double(_) => (
2785                    ColumnType::MYSQL_TYPE_DOUBLE,
2786                    StmtExecuteParamFlags::empty(),
2787                ),
2788                Value::Date(..) => (
2789                    ColumnType::MYSQL_TYPE_DATETIME,
2790                    StmtExecuteParamFlags::empty(),
2791                ),
2792                Value::Time(..) => (ColumnType::MYSQL_TYPE_TIME, StmtExecuteParamFlags::empty()),
2793            };
2794
2795            buf.put_slice(&[column_type as u8, flags.bits()]);
2796        }
2797
2798        for param in &self.params {
2799            match **param {
2800                Value::Int(_)
2801                | Value::UInt(_)
2802                | Value::Float(_)
2803                | Value::Double(_)
2804                | Value::Date(..)
2805                | Value::Time(..) => {
2806                    param.serialize(buf);
2807                }
2808                Value::Bytes(_) if !self.as_long_data => {
2809                    param.serialize(buf);
2810                }
2811                Value::Bytes(_) | Value::NULL => {}
2812            }
2813        }
2814    }
2815}
2816
2817define_header!(
2818    ComStmtSendLongDataHeader,
2819    COM_STMT_SEND_LONG_DATA,
2820    InvalidComStmtSendLongDataHeader
2821);
2822
2823#[derive(Debug, Clone, Eq, PartialEq)]
2824pub struct ComStmtSendLongData<'a> {
2825    __header: ComStmtSendLongDataHeader,
2826    stmt_id: RawInt<LeU32>,
2827    param_index: RawInt<LeU16>,
2828    data: RawBytes<'a, EofBytes>,
2829}
2830
2831impl<'a> ComStmtSendLongData<'a> {
2832    pub fn new(stmt_id: u32, param_index: u16, data: impl Into<Cow<'a, [u8]>>) -> Self {
2833        Self {
2834            __header: ComStmtSendLongDataHeader::new(),
2835            stmt_id: RawInt::new(stmt_id),
2836            param_index: RawInt::new(param_index),
2837            data: RawBytes::new(data),
2838        }
2839    }
2840
2841    pub fn into_owned(self) -> ComStmtSendLongData<'static> {
2842        ComStmtSendLongData {
2843            __header: self.__header,
2844            stmt_id: self.stmt_id,
2845            param_index: self.param_index,
2846            data: self.data.into_owned(),
2847        }
2848    }
2849}
2850
2851impl MySerialize for ComStmtSendLongData<'_> {
2852    fn serialize(&self, buf: &mut Vec<u8>) {
2853        self.__header.serialize(&mut *buf);
2854        self.stmt_id.serialize(&mut *buf);
2855        self.param_index.serialize(&mut *buf);
2856        self.data.serialize(&mut *buf);
2857    }
2858}
2859
2860#[derive(Debug, Clone, Copy, Eq, PartialEq)]
2861pub struct ComStmtClose {
2862    pub stmt_id: u32,
2863}
2864
2865impl ComStmtClose {
2866    pub fn new(stmt_id: u32) -> Self {
2867        Self { stmt_id }
2868    }
2869}
2870
2871impl MySerialize for ComStmtClose {
2872    fn serialize(&self, buf: &mut Vec<u8>) {
2873        buf.put_u8(Command::COM_STMT_CLOSE as u8);
2874        buf.put_u32_le(self.stmt_id);
2875    }
2876}
2877
2878/// Sends array of parameters to the server for the bulk execution of a prepared statement with
2879/// COM_STMT_BULK_EXECUTE command. This command is MariaDB only and may not be used for queries w/out
2880/// parameters and with empty parameter sets.
2881#[derive(Debug, Clone, PartialEq)]
2882pub struct ComStmtBulkExecuteRequestBuilder<'a> {
2883    stmt_id: u32,
2884    with_types: bool,
2885    with_results: bool,
2886    params_set: Vec<Vec<Value>>,
2887    payload_len: usize,
2888    arity: Option<usize>,
2889    named_params: Option<&'a [Vec<u8>]>,
2890    max_payload_len: usize, /* max_allowed_packet(if known) - 4 */
2891}
2892
2893impl ComStmtBulkExecuteRequestBuilder<'_> {
2894    pub fn new(
2895        stmt_id: u32,
2896        max_allowed_packet: usize,
2897    ) -> ComStmtBulkExecuteRequestBuilder<'static> {
2898        ComStmtBulkExecuteRequestBuilder {
2899            stmt_id,
2900            with_types: true,
2901            with_results: false,
2902            params_set: Vec::new(),
2903            payload_len: 0,
2904            arity: None,
2905            named_params: None,
2906            max_payload_len: max_allowed_packet - 4,
2907        }
2908    }
2909
2910    /// Use named parameters in the same order they were given in the SQL statement.
2911    pub fn with_named_params<'a>(
2912        self,
2913        named_params: Option<&'a [Vec<u8>]>,
2914    ) -> ComStmtBulkExecuteRequestBuilder<'a> {
2915        ComStmtBulkExecuteRequestBuilder {
2916            stmt_id: self.stmt_id,
2917            with_types: self.with_types,
2918            with_results: self.with_results,
2919            params_set: self.params_set,
2920            payload_len: self.payload_len,
2921            arity: self.arity,
2922            named_params,
2923            max_payload_len: self.max_payload_len,
2924        }
2925    }
2926
2927    /// Send [`SEND_UNIT_RESULTS`] flag to the server.
2928    ///
2929    /// This feature is available since MariaDb v11.5.1 and requires
2930    /// `MARIADB_CLIENT_BULK_UNIT_RESULTS` capability from the server.
2931    pub fn with_results(self) -> Self {
2932        ComStmtBulkExecuteRequestBuilder {
2933            stmt_id: self.stmt_id,
2934            with_types: self.with_types,
2935            with_results: true,
2936            params_set: self.params_set,
2937            payload_len: self.payload_len,
2938            arity: self.arity,
2939            named_params: self.named_params,
2940            max_payload_len: self.max_payload_len,
2941        }
2942    }
2943
2944    /// See [`ComStmtBulkExecuteRequestBuilder::add_row`].
2945    pub fn add_params(
2946        &mut self,
2947        params: impl Into<Params>,
2948    ) -> Result<Option<Vec<Value>>, BulkExecuteRequestBuilderError> {
2949        self._add_params(params.into())
2950    }
2951
2952    fn _add_params(
2953        &mut self,
2954        params: Params,
2955    ) -> Result<Option<Vec<Value>>, BulkExecuteRequestBuilderError> {
2956        self.add_row(params.into_values(self.named_params)?)
2957    }
2958
2959    /// Adds a new row of parameters to the bulk execute request.
2960    ///
2961    /// Returns row back if adding it would exceed the max allowed packet size — practically
2962    /// this means that you should call [`ComStmtBulkExecuteRequestBuilder::build`] to consume
2963    /// all the rows added so far and then continue adding more rows to the next bulk request
2964    /// starting from this returned row.
2965    ///
2966    /// # Error
2967    ///
2968    /// This function will emit an error if params' arity differs from previous
2969    /// rows added to this builder or if row is larger than the max payload length
2970    /// (max allowed packet - 4)
2971    pub fn add_row(
2972        &mut self,
2973        params: Vec<Value>,
2974    ) -> Result<Option<Vec<Value>>, BulkExecuteRequestBuilderError> {
2975        let arity = self.arity.get_or_insert(params.len());
2976        if params.len() != *arity {
2977            return Err(BulkExecuteRequestError::MixedArity.into());
2978        }
2979
2980        if self.with_types && self.payload_len == 0 {
2981            self.payload_len = params.len() * 2;
2982        }
2983
2984        let mut data_len = 0;
2985
2986        for p in &params {
2987            // bin_len() includes length encoding bytes
2988            match p.bin_len() as usize {
2989                0 => data_len += 1,     // NULLs take 1 byte for the indicator
2990                x => data_len += x + 1, // non-NULLs take their length + 1 byte for the indicator
2991            }
2992        }
2993
2994        // 7 = 1(command id) + 4 (stmt_id) + 2 (flags). If it's 1st row - we take it to return error
2995        // later(when the packet is sent). In this way we can avoid eternal loops of trying to add this row.
2996        if 7 + self.payload_len + data_len > self.max_payload_len {
2997            if self.params_set.is_empty() {
2998                return Err(BulkExecuteRequestError::RowTooLarge(params).into());
2999            }
3000            return Ok(Some(params));
3001        }
3002
3003        self.params_set.push(params);
3004        self.payload_len += data_len;
3005
3006        Ok(None)
3007    }
3008
3009    pub fn has_rows(&self) -> bool {
3010        !self.params_set.is_empty()
3011    }
3012
3013    /// Builds `COM_STMT_BULK_EXECUTE` consuming rows added so far.
3014    ///
3015    /// After the call you can continue using [`ComStmtBulkExecuteRequestBuilder::add_row`]
3016    /// to build next bulk request for this statement.
3017    ///
3018    /// # Error
3019    ///
3020    /// This will error if no rows was added to the builder.
3021    pub fn build(&mut self) -> Result<ComStmtBulkExecuteRequest<'static>, BulkExecuteRequestError> {
3022        let mut bulk_flags = StmtBulkExecuteFlags::empty();
3023
3024        if self.with_results {
3025            bulk_flags.insert(StmtBulkExecuteFlags::SEND_UNIT_RESULTS);
3026        }
3027
3028        if self.with_types {
3029            bulk_flags.insert(StmtBulkExecuteFlags::SEND_TYPES_TO_SERVER);
3030        }
3031
3032        self.with_types = false;
3033        self.payload_len = 0;
3034        ComStmtBulkExecuteRequest::new(self.stmt_id, bulk_flags, mem::take(&mut self.params_set))
3035    }
3036
3037    /// See [`ComStmtBulkExecuteRequestBuilder::build_iter`].
3038    pub fn build_params_iter(
3039        &mut self,
3040        input: impl IntoIterator<Item = impl Into<Params>>,
3041    ) -> impl Iterator<Item = Result<ComStmtBulkExecuteRequest<'static>, BulkExecuteRequestBuilderError>>
3042    {
3043        let mut done = false;
3044
3045        macro_rules! transpose {
3046            ($e:expr) => {
3047                match $e {
3048                    Ok(x) => x,
3049                    Err(e) => {
3050                        done = true;
3051                        return Some(Err(e.into()));
3052                    }
3053                }
3054            };
3055        }
3056
3057        let mut input = input.into_iter().map(Into::into);
3058        let mut stack = None;
3059        std::iter::from_fn(move || {
3060            if done {
3061                return None;
3062            }
3063
3064            let params_iter = stack
3065                .take()
3066                .map(Params::Positional)
3067                .into_iter()
3068                .chain(input.by_ref());
3069
3070            for params in params_iter {
3071                if let Some(params) = transpose!(self.add_params(params)) {
3072                    stack = Some(params);
3073                    return Some(Ok(transpose!(self.build())));
3074                }
3075            }
3076
3077            if self.has_rows() {
3078                return Some(Ok(transpose!(self.build())));
3079            }
3080
3081            done = true;
3082            None
3083        })
3084    }
3085
3086    /// It's a convenient wrapper over [`ComStmtBulkExecuteRequestBuilder::add_row`]
3087    /// and [`ComStmtBulkExecuteRequestBuilder::build`] that converts a rows iterator
3088    /// to a bulk requests iterator.
3089    ///
3090    /// # Error
3091    ///
3092    /// This won't error if the input iterator contains no rows — it'll just emit no bulk requests
3093    /// but the iterator will emit an error if it encounters a row with different arity.
3094    pub fn build_iter(
3095        &mut self,
3096        input: impl IntoIterator<Item = Vec<Value>>,
3097    ) -> impl Iterator<Item = Result<ComStmtBulkExecuteRequest<'static>, BulkExecuteRequestBuilderError>>
3098    {
3099        self.build_params_iter(input.into_iter().map(Params::Positional))
3100    }
3101}
3102
3103define_header!(
3104    ComStmtBulkExecuteHeader,
3105    COM_STMT_BULK_EXECUTE,
3106    InvalidComStmtBulkExecuteHeader
3107);
3108
3109#[derive(Debug, Clone, Copy, PartialEq, Eq)]
3110pub struct StmtBulkExecuteParamType {
3111    r#type: Const<ColumnType, u8>,
3112    flags: Const<StmtExecuteParamFlags, u8>,
3113}
3114
3115impl StmtBulkExecuteParamType {
3116    pub fn new(r#type: ColumnType, flags: StmtExecuteParamFlags) -> Self {
3117        Self {
3118            r#type: Const::new(r#type),
3119            flags: Const::new(flags),
3120        }
3121    }
3122
3123    pub fn from_value(value: &Value) -> Self {
3124        let (r#type, flags) = match value {
3125            Value::NULL => (ColumnType::MYSQL_TYPE_NULL, StmtExecuteParamFlags::empty()),
3126            Value::Bytes(_) => (
3127                ColumnType::MYSQL_TYPE_VAR_STRING,
3128                StmtExecuteParamFlags::empty(),
3129            ),
3130            Value::Int(_) => (
3131                ColumnType::MYSQL_TYPE_LONGLONG,
3132                StmtExecuteParamFlags::empty(),
3133            ),
3134            Value::UInt(_) => (
3135                ColumnType::MYSQL_TYPE_LONGLONG,
3136                StmtExecuteParamFlags::UNSIGNED,
3137            ),
3138            Value::Float(_) => (ColumnType::MYSQL_TYPE_FLOAT, StmtExecuteParamFlags::empty()),
3139            Value::Double(_) => (
3140                ColumnType::MYSQL_TYPE_DOUBLE,
3141                StmtExecuteParamFlags::empty(),
3142            ),
3143            Value::Date(..) => (
3144                ColumnType::MYSQL_TYPE_DATETIME,
3145                StmtExecuteParamFlags::empty(),
3146            ),
3147            Value::Time(..) => (ColumnType::MYSQL_TYPE_TIME, StmtExecuteParamFlags::empty()),
3148        };
3149
3150        Self {
3151            r#type: Const::new(r#type),
3152            flags: Const::new(flags),
3153        }
3154    }
3155
3156    pub fn column_type(&self) -> ColumnType {
3157        self.r#type.0
3158    }
3159
3160    pub fn flags(&self) -> StmtExecuteParamFlags {
3161        self.flags.0
3162    }
3163}
3164
3165impl<'de> MyDeserialize<'de> for StmtBulkExecuteParamType {
3166    const SIZE: Option<usize> = Some(3);
3167    type Ctx = ();
3168
3169    fn deserialize(_ctx: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3170        Ok(Self {
3171            r#type: buf.parse(())?,
3172            flags: buf.parse(())?,
3173        })
3174    }
3175}
3176
3177impl MySerialize for StmtBulkExecuteParamType {
3178    fn serialize(&self, buf: &mut Vec<u8>) {
3179        self.r#type.serialize(&mut *buf);
3180        self.flags.serialize(&mut *buf);
3181    }
3182}
3183
3184#[derive(Debug, Clone, PartialEq)]
3185pub struct StmtBulkExecuteParamValues {
3186    values: Vec<StmtBulkExecuteParamValue>,
3187}
3188
3189impl StmtBulkExecuteParamValues {
3190    pub fn new(values: impl IntoIterator<Item = Value>) -> Self {
3191        let values = values
3192            .into_iter()
3193            .map(StmtBulkExecuteParamValue::new)
3194            .collect();
3195        Self { values }
3196    }
3197}
3198
3199impl AsRef<[StmtBulkExecuteParamValue]> for StmtBulkExecuteParamValues {
3200    fn as_ref(&self) -> &[StmtBulkExecuteParamValue] {
3201        &self.values
3202    }
3203}
3204
3205impl<'de> MyDeserialize<'de> for StmtBulkExecuteParamValues {
3206    const SIZE: Option<usize> = None;
3207    type Ctx = Vec<(ColumnType, ColumnFlags)>;
3208
3209    fn deserialize(params: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3210        let mut values = Vec::with_capacity(params.len());
3211        for param in params {
3212            values.push(StmtBulkExecuteParamValue::deserialize(param, &mut *buf)?);
3213        }
3214        Ok(Self { values })
3215    }
3216}
3217
3218impl MySerialize for StmtBulkExecuteParamValues {
3219    fn serialize(&self, buf: &mut Vec<u8>) {
3220        for value in &self.values {
3221            value.serialize(&mut *buf);
3222        }
3223    }
3224}
3225
3226#[derive(Debug, Clone, PartialEq)]
3227pub struct StmtBulkExecuteParamValue {
3228    indicator: Const<MariadbBulkIndicator, u8>,
3229    value: Value,
3230}
3231
3232impl StmtBulkExecuteParamValue {
3233    pub fn new(value: Value) -> Self {
3234        let indicator = if matches!(value, Value::NULL) {
3235            MariadbBulkIndicator::BULK_INDICATOR_NULL
3236        } else {
3237            MariadbBulkIndicator::BULK_INDICATOR_NONE
3238        };
3239        Self {
3240            indicator: Const::new(indicator),
3241            value,
3242        }
3243    }
3244
3245    pub fn indicator(&self) -> MariadbBulkIndicator {
3246        self.indicator.0
3247    }
3248
3249    pub fn value(&self) -> &Value {
3250        &self.value
3251    }
3252}
3253
3254impl<'de> MyDeserialize<'de> for StmtBulkExecuteParamValue {
3255    const SIZE: Option<usize> = None;
3256    type Ctx = (ColumnType, ColumnFlags);
3257
3258    fn deserialize(ctx: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3259        let indicator = buf.parse::<Const<MariadbBulkIndicator, u8>>(())?;
3260        let value = if *indicator == MariadbBulkIndicator::BULK_INDICATOR_NULL {
3261            Value::NULL
3262        } else {
3263            ValueDeserializer::<BinValue>::deserialize(ctx, buf)?.0
3264        };
3265
3266        Ok(Self { indicator, value })
3267    }
3268}
3269
3270impl MySerialize for StmtBulkExecuteParamValue {
3271    fn serialize(&self, buf: &mut Vec<u8>) {
3272        self.indicator.serialize(&mut *buf);
3273        self.value.serialize(&mut *buf);
3274    }
3275}
3276
3277#[derive(Debug, Clone, PartialEq, thiserror::Error)]
3278pub enum BulkExecuteRequestError {
3279    #[error("No parameters values given for the bulk operation")]
3280    NoParams,
3281    #[error("Mixed statement arity")]
3282    MixedArity,
3283    #[error("Got row bigger than max payload length")]
3284    RowTooLarge(Vec<Value>),
3285}
3286
3287#[derive(Debug, Clone, PartialEq, thiserror::Error)]
3288pub enum BulkExecuteRequestBuilderError {
3289    #[error(transparent)]
3290    Request(#[from] BulkExecuteRequestError),
3291    #[error(transparent)]
3292    Params(#[from] ParamsError),
3293}
3294
3295#[derive(Debug, Clone, PartialEq)]
3296pub struct ComStmtBulkExecuteRequest<'a> {
3297    header: ComStmtBulkExecuteHeader,
3298    stmt_id: RawInt<LeU32>,
3299    bulk_flags: Const<StmtBulkExecuteFlags, LeU16>,
3300    types: Seq<'a, StmtBulkExecuteParamType, Unknown>,
3301    values: StmtBulkExecuteParamValues,
3302}
3303
3304impl<'a> ComStmtBulkExecuteRequest<'a> {
3305    pub fn new(
3306        stmt_id: u32,
3307        bulk_flags: StmtBulkExecuteFlags,
3308        values: Vec<Vec<Value>>,
3309    ) -> Result<Self, BulkExecuteRequestError> {
3310        let first = values.first().ok_or(BulkExecuteRequestError::NoParams)?;
3311        let arity = first.len();
3312
3313        let mut types = if bulk_flags.contains(StmtBulkExecuteFlags::SEND_TYPES_TO_SERVER) {
3314            first
3315                .iter()
3316                .map(StmtBulkExecuteParamType::from_value)
3317                .collect::<Vec<_>>()
3318        } else {
3319            Vec::default()
3320        };
3321
3322        for values in &values {
3323            if values.len() != arity {
3324                return Err(BulkExecuteRequestError::MixedArity);
3325            }
3326            let row_types = values
3327                .iter()
3328                .map(StmtBulkExecuteParamType::from_value)
3329                .collect::<Vec<_>>();
3330
3331            // The point here is to find proper type for every param, i.e the type
3332            // that covers param values in all rows.
3333            // E.g. ColumnType::MYSQL_TYPE_NULL is a proper type only if all the
3334            // values for the given param are NULLs
3335            for (left, right) in types.iter_mut().zip(&row_types) {
3336                if left != right {
3337                    if left.column_type() == ColumnType::MYSQL_TYPE_NULL {
3338                        *left = *right;
3339                    } else if left.column_type() == right.column_type() {
3340                        *left = StmtBulkExecuteParamType::new(
3341                            right.column_type(),
3342                            // if flag is required by a single param value,
3343                            // then it must be given for the whole batch
3344                            left.flags().union(right.flags()),
3345                        )
3346                    } else {
3347                        // TODO: Values of different types are given for the same parameter
3348                        //       within the batch. Not sure if server will always error here.
3349                        // The error:
3350                        //   ERROR 1210 (HY000): Incorrect arguments to mysqld_stmt_bulk_execute
3351                    }
3352                }
3353            }
3354        }
3355
3356        Ok(Self {
3357            header: ConstU8::new(),
3358            stmt_id: RawInt::new(stmt_id),
3359            bulk_flags: Const::new(bulk_flags),
3360            types: Seq::new(types),
3361            values: StmtBulkExecuteParamValues::new(values.into_iter().flatten()),
3362        })
3363    }
3364
3365    pub fn stmt_id(&self) -> u32 {
3366        self.stmt_id.0
3367    }
3368
3369    pub fn bulk_flags(&self) -> StmtBulkExecuteFlags {
3370        self.bulk_flags.0
3371    }
3372
3373    pub fn types(&self) -> &[StmtBulkExecuteParamType] {
3374        &self.types
3375    }
3376
3377    pub fn values(&self) -> &StmtBulkExecuteParamValues {
3378        &self.values
3379    }
3380}
3381
3382impl<'de> MyDeserialize<'de> for ComStmtBulkExecuteRequest<'de> {
3383    const SIZE: Option<usize> = None;
3384    type Ctx = Vec<(ColumnType, ColumnFlags)>;
3385
3386    fn deserialize(mut params: Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3387        let header = buf.parse(())?;
3388        let stmt_id = buf.parse(())?;
3389        let bulk_flags: Const<StmtBulkExecuteFlags, LeU16> = buf.parse(())?;
3390        let types = if bulk_flags.contains(StmtBulkExecuteFlags::SEND_TYPES_TO_SERVER) {
3391            let types: Seq<'_, StmtBulkExecuteParamType, Unknown> = buf.parse(params.len())?;
3392            for (s_type, r#type) in types.iter().zip(params.iter_mut()) {
3393                r#type.0 = *s_type.r#type;
3394                r#type.1.set(
3395                    ColumnFlags::UNSIGNED_FLAG,
3396                    s_type.flags.contains(StmtExecuteParamFlags::UNSIGNED),
3397                );
3398            }
3399            types
3400        } else {
3401            Seq::empty()
3402        };
3403        let values = buf.parse(params)?;
3404
3405        Ok(Self {
3406            header,
3407            stmt_id,
3408            bulk_flags,
3409            types,
3410            values,
3411        })
3412    }
3413}
3414
3415impl<'a> MySerialize for ComStmtBulkExecuteRequest<'a> {
3416    fn serialize(&self, buf: &mut Vec<u8>) {
3417        self.header.serialize(&mut *buf);
3418        self.stmt_id.serialize(&mut *buf);
3419        self.bulk_flags.serialize(&mut *buf);
3420        if self
3421            .bulk_flags
3422            .contains(StmtBulkExecuteFlags::SEND_TYPES_TO_SERVER)
3423        {
3424            self.types.serialize(&mut *buf);
3425        }
3426        self.values.serialize(&mut *buf);
3427    }
3428}
3429
3430define_header!(
3431    ComRegisterSlaveHeader,
3432    COM_REGISTER_SLAVE,
3433    InvalidComRegisterSlaveHeader
3434);
3435
3436/// Registers a slave at the master. Should be sent before requesting a binlog events
3437/// with `COM_BINLOG_DUMP`.
3438#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3439pub struct ComRegisterSlave<'a> {
3440    header: ComRegisterSlaveHeader,
3441    /// The slaves server-id.
3442    server_id: RawInt<LeU32>,
3443    /// The host name or IP address of the slave to be reported to the master during slave
3444    /// registration. Usually empty.
3445    hostname: RawBytes<'a, U8Bytes>,
3446    /// The account user name of the slave to be reported to the master during slave registration.
3447    /// Usually empty.
3448    ///
3449    /// # Note
3450    ///
3451    /// Serialization will truncate this value if length is greater than 255 bytes.
3452    user: RawBytes<'a, U8Bytes>,
3453    /// The account password of the slave to be reported to the master during slave registration.
3454    /// Usually empty.
3455    ///
3456    /// # Note
3457    ///
3458    /// Serialization will truncate this value if length is greater than 255 bytes.
3459    password: RawBytes<'a, U8Bytes>,
3460    /// The TCP/IP port number for connecting to the slave, to be reported to the master during
3461    /// slave registration. Usually empty.
3462    ///
3463    /// # Note
3464    ///
3465    /// Serialization will truncate this value if length is greater than 255 bytes.
3466    port: RawInt<LeU16>,
3467    /// Ignored.
3468    replication_rank: RawInt<LeU32>,
3469    /// Usually 0. Appears as "master id" in `SHOW SLAVE HOSTS` on the master. Unknown what else
3470    /// it impacts.
3471    master_id: RawInt<LeU32>,
3472}
3473
3474impl<'a> ComRegisterSlave<'a> {
3475    /// Creates new `ComRegisterSlave` with the given server identifier. Other fields will be empty.
3476    pub fn new(server_id: u32) -> Self {
3477        Self {
3478            header: Default::default(),
3479            server_id: RawInt::new(server_id),
3480            hostname: Default::default(),
3481            user: Default::default(),
3482            password: Default::default(),
3483            port: Default::default(),
3484            replication_rank: Default::default(),
3485            master_id: Default::default(),
3486        }
3487    }
3488
3489    /// Sets the `hostname` field of the packet (maximum length is 255 bytes).
3490    pub fn with_hostname(mut self, hostname: impl Into<Cow<'a, [u8]>>) -> Self {
3491        self.hostname = RawBytes::new(hostname);
3492        self
3493    }
3494
3495    /// Sets the `user` field of the packet (maximum length is 255 bytes).
3496    pub fn with_user(mut self, user: impl Into<Cow<'a, [u8]>>) -> Self {
3497        self.user = RawBytes::new(user);
3498        self
3499    }
3500
3501    /// Sets the `password` field of the packet (maximum length is 255 bytes).
3502    pub fn with_password(mut self, password: impl Into<Cow<'a, [u8]>>) -> Self {
3503        self.password = RawBytes::new(password);
3504        self
3505    }
3506
3507    /// Sets the `port` field of the packet.
3508    pub fn with_port(mut self, port: u16) -> Self {
3509        self.port = RawInt::new(port);
3510        self
3511    }
3512
3513    /// Sets the `replication_rank` field of the packet.
3514    pub fn with_replication_rank(mut self, replication_rank: u32) -> Self {
3515        self.replication_rank = RawInt::new(replication_rank);
3516        self
3517    }
3518
3519    /// Sets the `master_id` field of the packet.
3520    pub fn with_master_id(mut self, master_id: u32) -> Self {
3521        self.master_id = RawInt::new(master_id);
3522        self
3523    }
3524
3525    /// Returns the `server_id` field of the packet.
3526    pub fn server_id(&self) -> u32 {
3527        self.server_id.0
3528    }
3529
3530    /// Returns the raw `hostname` field value.
3531    pub fn hostname_raw(&self) -> &[u8] {
3532        self.hostname.as_bytes()
3533    }
3534
3535    /// Returns the `hostname` field as a UTF-8 string (lossy converted).
3536    pub fn hostname(&'a self) -> Cow<'a, str> {
3537        self.hostname.as_str()
3538    }
3539
3540    /// Returns the raw `user` field value.
3541    pub fn user_raw(&self) -> &[u8] {
3542        self.user.as_bytes()
3543    }
3544
3545    /// Returns the `user` field as a UTF-8 string (lossy converted).
3546    pub fn user(&'a self) -> Cow<'a, str> {
3547        self.user.as_str()
3548    }
3549
3550    /// Returns the raw `password` field value.
3551    pub fn password_raw(&self) -> &[u8] {
3552        self.password.as_bytes()
3553    }
3554
3555    /// Returns the `password` field as a UTF-8 string (lossy converted).
3556    pub fn password(&'a self) -> Cow<'a, str> {
3557        self.password.as_str()
3558    }
3559
3560    /// Returns the `port` field of the packet.
3561    pub fn port(&self) -> u16 {
3562        self.port.0
3563    }
3564
3565    /// Returns the `replication_rank` field of the packet.
3566    pub fn replication_rank(&self) -> u32 {
3567        self.replication_rank.0
3568    }
3569
3570    /// Returns the `master_id` field of the packet.
3571    pub fn master_id(&self) -> u32 {
3572        self.master_id.0
3573    }
3574}
3575
3576impl MySerialize for ComRegisterSlave<'_> {
3577    fn serialize(&self, buf: &mut Vec<u8>) {
3578        self.header.serialize(&mut *buf);
3579        self.server_id.serialize(&mut *buf);
3580        self.hostname.serialize(&mut *buf);
3581        self.user.serialize(&mut *buf);
3582        self.password.serialize(&mut *buf);
3583        self.port.serialize(&mut *buf);
3584        self.replication_rank.serialize(&mut *buf);
3585        self.master_id.serialize(&mut *buf);
3586    }
3587}
3588
3589impl<'de> MyDeserialize<'de> for ComRegisterSlave<'de> {
3590    const SIZE: Option<usize> = None;
3591    type Ctx = ();
3592
3593    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3594        let mut sbuf: ParseBuf<'_> = buf.parse(5)?;
3595        let header = sbuf.parse_unchecked(())?;
3596        let server_id = sbuf.parse_unchecked(())?;
3597
3598        let hostname = buf.parse(())?;
3599        let user = buf.parse(())?;
3600        let password = buf.parse(())?;
3601
3602        let mut sbuf: ParseBuf<'_> = buf.parse(10)?;
3603        let port = sbuf.parse_unchecked(())?;
3604        let replication_rank = sbuf.parse_unchecked(())?;
3605        let master_id = sbuf.parse_unchecked(())?;
3606
3607        Ok(Self {
3608            header,
3609            server_id,
3610            hostname,
3611            user,
3612            password,
3613            port,
3614            replication_rank,
3615            master_id,
3616        })
3617    }
3618}
3619
3620define_header!(
3621    ComTableDumpHeader,
3622    COM_TABLE_DUMP,
3623    InvalidComTableDumpHeader
3624);
3625
3626/// COM_TABLE_DUMP command.
3627#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3628pub struct ComTableDump<'a> {
3629    header: ComTableDumpHeader,
3630    /// Database name.
3631    ///
3632    /// # Note
3633    ///
3634    /// Serialization will truncate this value if length is greater than 255 bytes.
3635    database: RawBytes<'a, U8Bytes>,
3636    /// Table name.
3637    ///
3638    /// # Note
3639    ///
3640    /// Serialization will truncate this value if length is greater than 255 bytes.
3641    table: RawBytes<'a, U8Bytes>,
3642}
3643
3644impl<'a> ComTableDump<'a> {
3645    /// Creates new instance.
3646    pub fn new(database: impl Into<Cow<'a, [u8]>>, table: impl Into<Cow<'a, [u8]>>) -> Self {
3647        Self {
3648            header: Default::default(),
3649            database: RawBytes::new(database),
3650            table: RawBytes::new(table),
3651        }
3652    }
3653
3654    /// Returns the raw `database` field value.
3655    pub fn database_raw(&self) -> &[u8] {
3656        self.database.as_bytes()
3657    }
3658
3659    /// Returns the `database` field value as a UTF-8 string (lossy converted).
3660    pub fn database(&self) -> Cow<'_, str> {
3661        self.database.as_str()
3662    }
3663
3664    /// Returns the raw `table` field value.
3665    pub fn table_raw(&self) -> &[u8] {
3666        self.table.as_bytes()
3667    }
3668
3669    /// Returns the `table` field value as a UTF-8 string (lossy converted).
3670    pub fn table(&self) -> Cow<'_, str> {
3671        self.table.as_str()
3672    }
3673}
3674
3675impl MySerialize for ComTableDump<'_> {
3676    fn serialize(&self, buf: &mut Vec<u8>) {
3677        self.header.serialize(&mut *buf);
3678        self.database.serialize(&mut *buf);
3679        self.table.serialize(&mut *buf);
3680    }
3681}
3682
3683impl<'de> MyDeserialize<'de> for ComTableDump<'de> {
3684    const SIZE: Option<usize> = None;
3685    type Ctx = ();
3686
3687    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3688        Ok(Self {
3689            header: buf.parse(())?,
3690            database: buf.parse(())?,
3691            table: buf.parse(())?,
3692        })
3693    }
3694}
3695
3696my_bitflags! {
3697    BinlogDumpFlags,
3698    #[error("Unknown flags in the raw value of BinlogDumpFlags (raw={0:b})")]
3699    UnknownBinlogDumpFlags,
3700    u16,
3701
3702    /// Empty flags of a `LoadEvent`.
3703    #[derive(PartialEq, Eq, Hash, Debug, Clone, Copy)]
3704    pub struct BinlogDumpFlags: u16 {
3705        /// If there is no more event to send a EOF_Packet instead of blocking the connection
3706        const BINLOG_DUMP_NON_BLOCK = 0x01;
3707        const BINLOG_THROUGH_POSITION = 0x02;
3708        const BINLOG_THROUGH_GTID = 0x04;
3709    }
3710}
3711
3712define_header!(
3713    ComBinlogDumpHeader,
3714    COM_BINLOG_DUMP,
3715    InvalidComBinlogDumpHeader
3716);
3717
3718/// Command to request a binlog-stream from the master starting a given position.
3719#[derive(Clone, Debug, Eq, PartialEq, Hash)]
3720pub struct ComBinlogDump<'a> {
3721    header: ComBinlogDumpHeader,
3722    /// Position in the binlog-file to start the stream with (`0` by default).
3723    pos: RawInt<LeU32>,
3724    /// Command flags (empty by default).
3725    ///
3726    /// Only `BINLOG_DUMP_NON_BLOCK` is supported for this command.
3727    flags: Const<BinlogDumpFlags, LeU16>,
3728    /// Server id of this slave.
3729    server_id: RawInt<LeU32>,
3730    /// Filename of the binlog on the master.
3731    ///
3732    /// If the binlog-filename is empty, the server will send the binlog-stream of the first known
3733    /// binlog.
3734    filename: RawBytes<'a, EofBytes>,
3735}
3736
3737impl<'a> ComBinlogDump<'a> {
3738    /// Creates new instance with default values for `pos` and `flags`.
3739    pub fn new(server_id: u32) -> Self {
3740        Self {
3741            header: Default::default(),
3742            pos: Default::default(),
3743            flags: Default::default(),
3744            server_id: RawInt::new(server_id),
3745            filename: Default::default(),
3746        }
3747    }
3748
3749    /// Defines position for this instance.
3750    pub fn with_pos(mut self, pos: u32) -> Self {
3751        self.pos = RawInt::new(pos);
3752        self
3753    }
3754
3755    /// Defines flags for this instance.
3756    pub fn with_flags(mut self, flags: BinlogDumpFlags) -> Self {
3757        self.flags = Const::new(flags);
3758        self
3759    }
3760
3761    /// Defines filename for this instance.
3762    pub fn with_filename(mut self, filename: impl Into<Cow<'a, [u8]>>) -> Self {
3763        self.filename = RawBytes::new(filename);
3764        self
3765    }
3766
3767    /// Returns parsed `pos` field with unknown bits truncated.
3768    pub fn pos(&self) -> u32 {
3769        *self.pos
3770    }
3771
3772    /// Returns parsed `flags` field with unknown bits truncated.
3773    pub fn flags(&self) -> BinlogDumpFlags {
3774        *self.flags
3775    }
3776
3777    /// Returns parsed `server_id` field with unknown bits truncated.
3778    pub fn server_id(&self) -> u32 {
3779        *self.server_id
3780    }
3781
3782    /// Returns the raw `filename` field value.
3783    pub fn filename_raw(&self) -> &[u8] {
3784        self.filename.as_bytes()
3785    }
3786
3787    /// Returns the `filename` field value as a UTF-8 string (lossy converted).
3788    pub fn filename(&self) -> Cow<'_, str> {
3789        self.filename.as_str()
3790    }
3791}
3792
3793impl MySerialize for ComBinlogDump<'_> {
3794    fn serialize(&self, buf: &mut Vec<u8>) {
3795        self.header.serialize(&mut *buf);
3796        self.pos.serialize(&mut *buf);
3797        self.flags.serialize(&mut *buf);
3798        self.server_id.serialize(&mut *buf);
3799        self.filename.serialize(&mut *buf);
3800    }
3801}
3802
3803impl<'de> MyDeserialize<'de> for ComBinlogDump<'de> {
3804    const SIZE: Option<usize> = None;
3805    type Ctx = ();
3806
3807    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3808        let mut sbuf: ParseBuf<'_> = buf.parse(11)?;
3809        Ok(Self {
3810            header: sbuf.parse_unchecked(())?,
3811            pos: sbuf.parse_unchecked(())?,
3812            flags: sbuf.parse_unchecked(())?,
3813            server_id: sbuf.parse_unchecked(())?,
3814            filename: buf.parse(())?,
3815        })
3816    }
3817}
3818
3819/// GnoInterval. Stored within [`Sid`]
3820#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
3821pub struct GnoInterval {
3822    start: RawInt<LeU64>,
3823    end: RawInt<LeU64>,
3824}
3825
3826impl GnoInterval {
3827    /// Creates a new interval.
3828    pub fn new(start: u64, end: u64) -> Self {
3829        Self {
3830            start: RawInt::new(start),
3831            end: RawInt::new(end),
3832        }
3833    }
3834
3835    /// Returns the start of the interval (inclusive).
3836    pub fn start(&self) -> u64 {
3837        self.start.0
3838    }
3839
3840    /// Returns the end of the interval (exclusive).
3841    pub fn end(&self) -> u64 {
3842        self.end.0
3843    }
3844    /// Checks if the [start, end) interval is valid and creates it.
3845    pub fn check_and_new(start: u64, end: u64) -> io::Result<Self> {
3846        if start >= end {
3847            return Err(io::Error::new(
3848                io::ErrorKind::InvalidData,
3849                format!("start({}) >= end({}) in GnoInterval", start, end),
3850            ));
3851        }
3852        if start == 0 || end == 0 {
3853            return Err(io::Error::new(
3854                io::ErrorKind::InvalidData,
3855                "Gno can't be zero",
3856            ));
3857        }
3858        Ok(Self::new(start, end))
3859    }
3860}
3861
3862impl MySerialize for GnoInterval {
3863    fn serialize(&self, buf: &mut Vec<u8>) {
3864        self.start.serialize(&mut *buf);
3865        self.end.serialize(&mut *buf);
3866    }
3867}
3868
3869impl<'de> MyDeserialize<'de> for GnoInterval {
3870    const SIZE: Option<usize> = Some(16);
3871    type Ctx = ();
3872
3873    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
3874        Ok(Self {
3875            start: buf.parse_unchecked(())?,
3876            end: buf.parse_unchecked(())?,
3877        })
3878    }
3879}
3880
3881/// Length of a Uuid in `COM_BINLOG_DUMP_GTID` command packet.
3882pub const UUID_LEN: usize = 16;
3883
3884/// SID is a part of the `COM_BINLOG_DUMP_GTID` command. It's a GtidSet whose
3885/// A GTID tag (MySQL 8.4+).
3886///
3887/// Tags are used to group transactions in tagged GTIDs.
3888/// The format is: `[a-z_][a-z0-9_]{0,31}` (1-32 characters).
3889///
3890/// Tagged GTIDs have the format: `UUID:tag:transaction_id`
3891#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
3892pub struct Tag<'a>(Cow<'a, str>);
3893
3894impl<'a> Tag<'a> {
3895    /// Maximum length of a tag in characters.
3896    pub const MAX_LEN: usize = 32;
3897
3898    /// Creates a new tag, validating the format.
3899    ///
3900    /// Tags must:
3901    /// - Start with a lowercase letter or underscore
3902    /// - Contain only lowercase letters, digits, or underscores
3903    /// - Be 1-32 characters long
3904    pub fn new(s: impl Into<Cow<'a, str>>) -> Result<Self, InvalidTag> {
3905        let s = s.into();
3906        Self::validate(&s)?;
3907        Ok(Self(s))
3908    }
3909
3910    /// Validates a tag string.
3911    fn validate(s: &str) -> Result<(), InvalidTag> {
3912        if s.is_empty() {
3913            return Err(InvalidTag::Empty);
3914        }
3915        if s.len() > Self::MAX_LEN {
3916            return Err(InvalidTag::TooLong(s.len()));
3917        }
3918
3919        let mut chars = s.chars();
3920
3921        // First character must be lowercase letter or underscore
3922        match chars.next() {
3923            Some(c) if c.is_ascii_lowercase() || c == '_' => {}
3924            Some(c) => return Err(InvalidTag::InvalidFirstChar(c)),
3925            None => return Err(InvalidTag::Empty),
3926        }
3927
3928        // Remaining characters must be lowercase letters, digits, or underscores
3929        for c in chars {
3930            if !c.is_ascii_lowercase() && !c.is_ascii_digit() && c != '_' {
3931                return Err(InvalidTag::InvalidChar(c));
3932            }
3933        }
3934
3935        Ok(())
3936    }
3937
3938    /// Returns the tag as a string slice.
3939    pub fn as_str(&self) -> &str {
3940        &self.0
3941    }
3942
3943    /// Converts this tag into an owned version.
3944    pub fn into_owned(self) -> Tag<'static> {
3945        Tag(Cow::Owned(self.0.into_owned()))
3946    }
3947}
3948
3949impl std::fmt::Display for Tag<'_> {
3950    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
3951        write!(f, "{}", self.0)
3952    }
3953}
3954
3955impl std::ops::Deref for Tag<'_> {
3956    type Target = str;
3957
3958    fn deref(&self) -> &Self::Target {
3959        &self.0
3960    }
3961}
3962
3963impl<'a> TryFrom<&'a str> for Tag<'a> {
3964    type Error = InvalidTag;
3965
3966    fn try_from(s: &'a str) -> Result<Self, Self::Error> {
3967        Self::new(s)
3968    }
3969}
3970
3971impl TryFrom<String> for Tag<'static> {
3972    type Error = InvalidTag;
3973
3974    fn try_from(s: String) -> Result<Self, Self::Error> {
3975        Self::new(s)
3976    }
3977}
3978
3979/// Error for invalid GTID tag format.
3980#[derive(Debug, Clone, PartialEq, Eq, Hash, thiserror::Error)]
3981pub enum InvalidTag {
3982    #[error("tag is empty")]
3983    Empty,
3984    #[error("tag is too long ({len} > {max})", len = .0, max = Tag::MAX_LEN)]
3985    TooLong(usize),
3986    #[error("invalid first character '{0}' (must be lowercase letter or underscore)")]
3987    InvalidFirstChar(char),
3988    #[error("invalid character '{0}' (must be lowercase letter, digit, or underscore)")]
3989    InvalidChar(char),
3990}
3991
3992/// has only one Uuid.
3993///
3994/// # Tagged GTIDs (MySQL 8.4+)
3995///
3996/// Supports tagged GTIDs with the format `UUID:tag:intervals`.
3997/// When a tag is present, transactions are grouped by the tag component.
3998#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3999pub struct Sid<'a> {
4000    uuid: [u8; UUID_LEN],
4001    /// Optional tag for tagged GTIDs (MySQL 8.4+).
4002    ///
4003    /// When present in any Sid within a [`ComBinlogDumpGtid`] packet, the
4004    /// entire sid block is encoded in the tagged (TSID) wire format. In that
4005    /// format every entry carries a varlen-encoded tag length prefix (0 for
4006    /// untagged entries).
4007    tag: Option<Tag<'a>>,
4008    intervals: Seq<'a, GnoInterval, LeU64>,
4009}
4010
4011impl Sid<'_> {
4012    /// Creates a new instance.
4013    pub fn new(uuid: [u8; UUID_LEN]) -> Self {
4014        Self {
4015            uuid,
4016            tag: None,
4017            intervals: Default::default(),
4018        }
4019    }
4020
4021    /// Returns the `uuid` field value.
4022    pub fn uuid(&self) -> [u8; UUID_LEN] {
4023        self.uuid
4024    }
4025
4026    /// Returns the `tag` field value (MySQL 8.4+).
4027    pub fn tag(&self) -> Option<&Tag<'_>> {
4028        self.tag.as_ref()
4029    }
4030
4031    /// Returns the `intervals` field value.
4032    pub fn intervals(&self) -> &[GnoInterval] {
4033        &self.intervals[..]
4034    }
4035
4036    /// Appends an GnoInterval to this block.
4037    pub fn with_interval(mut self, interval: GnoInterval) -> Self {
4038        let mut intervals = self.intervals.0.into_owned();
4039        intervals.push(interval);
4040        self.intervals = Seq::new(intervals);
4041        self
4042    }
4043
4044    /// Sets the `intervals` value for this block.
4045    pub fn with_intervals(mut self, intervals: Vec<GnoInterval>) -> Self {
4046        self.intervals = Seq::new(intervals);
4047        self
4048    }
4049
4050    fn len(&self) -> u64 {
4051        use saturating::Saturating as S;
4052        let mut len = S(UUID_LEN as u64); // SID
4053        len += S(8); // n_intervals
4054        len += S((self.intervals.len() * 16) as u64);
4055        len.0
4056    }
4057
4058    /// Returns the serialized length in the tagged (TSID) wire format.
4059    ///
4060    /// Layout: UUID (16) + varlen(tag_len) + tag_bytes + n_intervals (8) + intervals (16 each)
4061    fn tagged_len(&self) -> u64 {
4062        use saturating::Saturating as S;
4063        let tag_str = self.tag.as_ref().map(|t| t.as_str()).unwrap_or("");
4064        let mut len = S(UUID_LEN as u64);
4065        len += S(varlen_uint_size(tag_str.len() as u64) as u64);
4066        len += S(tag_str.len() as u64);
4067        len += S(8); // n_intervals
4068        len += S((self.intervals.len() * 16) as u64);
4069        len.0
4070    }
4071
4072    /// Serializes this Sid in the tagged (TSID) wire format.
4073    ///
4074    /// Layout: UUID (16) + varlen(tag_len) + tag_bytes + intervals (Seq<GnoInterval, LeU64>)
4075    fn serialize_tagged(&self, buf: &mut Vec<u8>) {
4076        self.uuid.serialize(&mut *buf);
4077        let tag_str = self.tag.as_ref().map(|t| t.as_str()).unwrap_or("");
4078        write_varlen_uint(buf, tag_str.len() as u64);
4079        if !tag_str.is_empty() {
4080            buf.extend_from_slice(tag_str.as_bytes());
4081        }
4082        self.intervals.serialize(buf);
4083    }
4084}
4085
4086impl<'de> Sid<'de> {
4087    /// Deserializes a Sid from the tagged (TSID) wire format.
4088    fn deserialize_tagged(buf: &mut ParseBuf<'de>) -> io::Result<Self> {
4089        let uuid: [u8; UUID_LEN] = buf.parse(())?;
4090        let tag_len = read_varlen_uint(buf)? as usize;
4091        let tag = if tag_len > 0 {
4092            if buf.len() < tag_len {
4093                return Err(io::Error::new(
4094                    io::ErrorKind::UnexpectedEof,
4095                    "unexpected end of buffer reading Sid tag",
4096                ));
4097            }
4098            let tag_bytes = &buf.0[..tag_len];
4099            let tag_str = std::str::from_utf8(tag_bytes).map_err(|e| {
4100                io::Error::new(
4101                    io::ErrorKind::InvalidData,
4102                    format!("invalid UTF-8 in Sid tag: {}", e),
4103                )
4104            })?;
4105            buf.0 = &buf.0[tag_len..];
4106            Some(
4107                Tag::new(tag_str.to_owned())
4108                    .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?,
4109            )
4110        } else {
4111            None
4112        };
4113        let intervals: Seq<'de, GnoInterval, LeU64> = buf.parse(())?;
4114        Ok(Sid {
4115            uuid,
4116            tag,
4117            intervals,
4118        })
4119    }
4120}
4121
4122impl<'a> Sid<'a> {
4123    /// Sets the `tag` value for this SID (MySQL 8.4+).
4124    ///
4125    /// When any Sid in a [`ComBinlogDumpGtid`] packet carries a tag, the
4126    /// entire sid block switches to the tagged (TSID) wire format.
4127    pub fn with_tag(mut self, tag: Tag<'a>) -> Self {
4128        self.tag = Some(tag);
4129        self
4130    }
4131}
4132
4133impl MySerialize for Sid<'_> {
4134    fn serialize(&self, buf: &mut Vec<u8>) {
4135        self.uuid.serialize(&mut *buf);
4136        self.intervals.serialize(buf);
4137    }
4138}
4139
4140impl<'de> MyDeserialize<'de> for Sid<'de> {
4141    const SIZE: Option<usize> = None;
4142    type Ctx = ();
4143
4144    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
4145        Ok(Self {
4146            uuid: buf.parse(())?,
4147            // Tags are not present in the binary wire protocol format.
4148            // Tagged GTIDs are conveyed via GTID_TAGGED_LOG_EVENT, not via Sid on the wire.
4149            tag: None,
4150            intervals: buf.parse(())?,
4151        })
4152    }
4153}
4154
4155impl fmt::Display for Sid<'_> {
4156    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
4157        let uuid = Uuid::from_bytes(self.uuid);
4158        write!(f, "{}", uuid.as_hyphenated())?;
4159
4160        if let Some(tag) = &self.tag {
4161            write!(f, ":{tag}")?;
4162        }
4163
4164        for interval in self.intervals.iter() {
4165            let start = *interval.start;
4166            let end = *interval.end;
4167            if end == start + 1 {
4168                write!(f, ":{start}")?;
4169            } else {
4170                write!(f, ":{}-{}", start, end - 1)?;
4171            }
4172        }
4173
4174        Ok(())
4175    }
4176}
4177
4178impl Sid<'_> {
4179    fn wrap_err(msg: String) -> io::Error {
4180        io::Error::new(io::ErrorKind::InvalidInput, msg)
4181    }
4182
4183    fn parse_interval_num(to_parse: &str, full: &str) -> Result<u64, io::Error> {
4184        let n: u64 = to_parse.parse().map_err(|e| {
4185            Sid::wrap_err(format!(
4186                "invalid GnoInterval format: {}, error: {}",
4187                full, e
4188            ))
4189        })?;
4190        Ok(n)
4191    }
4192
4193    /// Parses the tag and intervals portion of a SID string.
4194    ///
4195    /// Returns (tag, intervals_str) where tag is Some if a tag is present.
4196    /// Tagged format: `tag:intervals` where tag matches `[a-z_][a-z0-9_]*`
4197    /// Non-tagged format: `intervals` (starts with a digit)
4198    fn parse_tag_and_intervals<'a>(
4199        rest: &'a str,
4200        full: &str,
4201    ) -> Result<(Option<Tag<'static>>, &'a str), io::Error> {
4202        // Check if the first component looks like a tag
4203        // Tags start with [a-z_], intervals start with digits
4204        let first_char = rest.chars().next().ok_or_else(|| {
4205            Sid::wrap_err(format!("invalid sid format (empty after UUID): {}", full))
4206        })?;
4207
4208        if first_char.is_ascii_digit() {
4209            // No tag, rest is all intervals
4210            return Ok((None, rest));
4211        }
4212
4213        // Might be a tag - find the first colon to separate tag from intervals
4214        if let Some((potential_tag, intervals)) = rest.split_once(':') {
4215            // Validate the tag format
4216            match Tag::new(potential_tag.to_owned()) {
4217                Ok(tag) => Ok((Some(tag), intervals)),
4218                Err(e) => Err(Sid::wrap_err(format!(
4219                    "invalid tag format in sid {}: {}",
4220                    full, e
4221                ))),
4222            }
4223        } else {
4224            // No colon after potential tag - this might be an error
4225            // or it could be a single-component tag with no intervals (unlikely but possible)
4226            Err(Sid::wrap_err(format!(
4227                "invalid sid format (no intervals after tag?): {}",
4228                full
4229            )))
4230        }
4231    }
4232}
4233
4234impl FromStr for Sid<'_> {
4235    type Err = io::Error;
4236
4237    fn from_str(s: &str) -> Result<Self, Self::Err> {
4238        let (uuid, rest) = s
4239            .split_once(':')
4240            .ok_or_else(|| Sid::wrap_err(format!("invalid sid format: {}", s)))?;
4241        let uuid = Uuid::parse_str(uuid)
4242            .map_err(|e| Sid::wrap_err(format!("invalid uuid format: {}, error: {}", s, e)))?;
4243
4244        // Check if the first component after UUID is a tag or an interval
4245        // Tags start with [a-z_] and contain only [a-z0-9_]
4246        // Intervals are numeric (possibly with a dash for ranges)
4247        let (tag, intervals_str) = Self::parse_tag_and_intervals(rest, s)?;
4248
4249        let intervals: Vec<GnoInterval> = intervals_str
4250            .split(':')
4251            .map(|interval: &str| {
4252                let numbers = interval.split('-').collect::<Vec<_>>();
4253                if numbers.len() != 1 && numbers.len() != 2 {
4254                    return Err(Sid::wrap_err(format!("invalid GnoInterval format: {}", s)));
4255                }
4256                if numbers.len() == 1 {
4257                    let start = Sid::parse_interval_num(numbers[0], s)?;
4258                    let interval = GnoInterval::check_and_new(start, start + 1)?;
4259                    Ok(interval)
4260                } else {
4261                    let start = Sid::parse_interval_num(numbers[0], s)?;
4262                    let end = Sid::parse_interval_num(numbers[1], s)?;
4263                    let interval = GnoInterval::check_and_new(start, end + 1)?;
4264                    Ok(interval)
4265                }
4266            })
4267            .collect::<Result<Vec<_>, _>>()?;
4268        Ok(Self {
4269            uuid: *uuid.as_bytes(),
4270            tag,
4271            intervals: Seq::new(intervals),
4272        })
4273    }
4274}
4275
4276define_header!(
4277    ComBinlogDumpGtidHeader,
4278    COM_BINLOG_DUMP_GTID,
4279    InvalidComBinlogDumpGtidHeader
4280);
4281
4282/// Command to request a binlog-stream from the master starting a given position.
4283#[derive(Debug, Clone, Eq, PartialEq, Hash)]
4284pub struct ComBinlogDumpGtid<'a> {
4285    header: ComBinlogDumpGtidHeader,
4286    /// Command flags (empty by default).
4287    flags: Const<BinlogDumpFlags, LeU16>,
4288    /// Server id of this slave.
4289    server_id: RawInt<LeU32>,
4290    /// Filename of the binlog on the master.
4291    ///
4292    /// If the binlog-filename is empty, the server will send the binlog-stream of the first known
4293    /// binlog.
4294    ///
4295    /// # Note
4296    ///
4297    /// Serialization will truncate this value if length is greater than 2^32 - 1 bytes.
4298    filename: RawBytes<'a, U32Bytes>,
4299    /// Position in the binlog-file to start the stream with (`0` by default).
4300    pos: RawInt<LeU64>,
4301    /// SID block.
4302    sid_block: Seq<'a, Sid<'a>, LeU64>,
4303}
4304
4305impl<'a> ComBinlogDumpGtid<'a> {
4306    /// Creates new instance with default values for `pos`, `data` and `flags` fields.
4307    pub fn new(server_id: u32) -> Self {
4308        Self {
4309            header: Default::default(),
4310            pos: Default::default(),
4311            flags: Default::default(),
4312            server_id: RawInt::new(server_id),
4313            filename: Default::default(),
4314            sid_block: Default::default(),
4315        }
4316    }
4317
4318    /// Returns the `server_id` field value.
4319    pub fn server_id(&self) -> u32 {
4320        self.server_id.0
4321    }
4322
4323    /// Returns the `flags` field value.
4324    pub fn flags(&self) -> BinlogDumpFlags {
4325        self.flags.0
4326    }
4327
4328    /// Returns the `filename` field value.
4329    pub fn filename_raw(&self) -> &[u8] {
4330        self.filename.as_bytes()
4331    }
4332
4333    /// Returns the `filename` field value as a UTF-8 string (lossy converted).
4334    pub fn filename(&self) -> Cow<'_, str> {
4335        self.filename.as_str()
4336    }
4337
4338    /// Returns the `pos` field value.
4339    pub fn pos(&self) -> u64 {
4340        self.pos.0
4341    }
4342
4343    /// Returns the sequence of sids in this packet.
4344    pub fn sids(&self) -> &[Sid<'a>] {
4345        &self.sid_block
4346    }
4347
4348    /// Defines filename for this instance.
4349    pub fn with_filename(self, filename: impl Into<Cow<'a, [u8]>>) -> Self {
4350        Self {
4351            header: self.header,
4352            flags: self.flags,
4353            server_id: self.server_id,
4354            filename: RawBytes::new(filename),
4355            pos: self.pos,
4356            sid_block: self.sid_block,
4357        }
4358    }
4359
4360    /// Sets the `server_id` field value.
4361    pub fn with_server_id(mut self, server_id: u32) -> Self {
4362        self.server_id.0 = server_id;
4363        self
4364    }
4365
4366    /// Sets the `flags` field value.
4367    pub fn with_flags(mut self, mut flags: BinlogDumpFlags) -> Self {
4368        if self.sid_block.is_empty() {
4369            flags.remove(BinlogDumpFlags::BINLOG_THROUGH_GTID);
4370        } else {
4371            flags.insert(BinlogDumpFlags::BINLOG_THROUGH_GTID);
4372        }
4373        self.flags.0 = flags;
4374        self
4375    }
4376
4377    /// Sets the `pos` field value.
4378    pub fn with_pos(mut self, pos: u64) -> Self {
4379        self.pos.0 = pos;
4380        self
4381    }
4382
4383    /// Sets the `sid_block` field value.
4384    pub fn with_sid(mut self, sid: Sid<'a>) -> Self {
4385        self.flags.0.insert(BinlogDumpFlags::BINLOG_THROUGH_GTID);
4386        self.sid_block.push(sid);
4387        self
4388    }
4389
4390    /// Sets the `sid_block` field value.
4391    pub fn with_sids(mut self, sids: impl Into<Cow<'a, [Sid<'a>]>>) -> Self {
4392        self.sid_block = Seq::new(sids);
4393        if self.sid_block.is_empty() {
4394            self.flags.0.remove(BinlogDumpFlags::BINLOG_THROUGH_GTID);
4395        } else {
4396            self.flags.0.insert(BinlogDumpFlags::BINLOG_THROUGH_GTID);
4397        }
4398        self
4399    }
4400
4401    /// Returns `true` if any Sid in the block has a tag, requiring the tagged
4402    /// (TSID) wire format.
4403    fn is_tagged_format(&self) -> bool {
4404        self.sid_block.iter().any(|sid| sid.tag().is_some())
4405    }
4406
4407    fn sid_block_len(&self) -> u32 {
4408        use saturating::Saturating as S;
4409        let mut len = S(8_u32); // n_sids header (8 bytes including format indicator)
4410        let tagged = self.is_tagged_format();
4411        for sid in self.sid_block.iter() {
4412            if tagged {
4413                len += S(sid.tagged_len() as u32);
4414            } else {
4415                len += S(sid.len() as u32);
4416            }
4417        }
4418        len.0
4419    }
4420}
4421
4422impl MySerialize for ComBinlogDumpGtid<'_> {
4423    fn serialize(&self, buf: &mut Vec<u8>) {
4424        self.header.serialize(&mut *buf);
4425        self.flags.serialize(&mut *buf);
4426        self.server_id.serialize(&mut *buf);
4427        self.filename.serialize(&mut *buf);
4428        self.pos.serialize(&mut *buf);
4429        buf.put_u32_le(self.sid_block_len());
4430
4431        let tagged = self.is_tagged_format();
4432        let format_byte: u8 = if tagged { 0x01 } else { 0x00 };
4433        let n_sids = self.sid_block.len() as u64;
4434        // MySQL encodes the n_sids header differently for tagged vs untagged:
4435        //   Untagged: bits [55:0] = n_sids, bits [63:56] = format (0x00)
4436        //   Tagged:   bits [55:8] = n_sids, bits [63:56] = format (0x01), bits [7:0] = reserved
4437        let n_sids_header = if tagged {
4438            (n_sids << 8) | ((format_byte as u64) << 56)
4439        } else {
4440            n_sids | ((format_byte as u64) << 56)
4441        };
4442        buf.put_u64_le(n_sids_header);
4443
4444        for sid in self.sid_block.iter() {
4445            if tagged {
4446                sid.serialize_tagged(&mut *buf);
4447            } else {
4448                sid.serialize(&mut *buf);
4449            }
4450        }
4451    }
4452}
4453
4454impl<'de> MyDeserialize<'de> for ComBinlogDumpGtid<'de> {
4455    const SIZE: Option<usize> = None;
4456    type Ctx = ();
4457
4458    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
4459        let mut sbuf: ParseBuf<'_> = buf.parse(7)?;
4460        let header = sbuf.parse_unchecked(())?;
4461        let flags: Const<BinlogDumpFlags, LeU16> = sbuf.parse_unchecked(())?;
4462        let server_id = sbuf.parse_unchecked(())?;
4463
4464        let filename = buf.parse(())?;
4465        let pos = buf.parse(())?;
4466
4467        // `flags` should contain `BINLOG_THROUGH_GTID` flag if sid_block isn't empty
4468        let sid_data_len: RawInt<LeU32> = buf.parse(())?;
4469        let mut sid_buf: ParseBuf<'de> = buf.parse(sid_data_len.0 as usize)?;
4470
4471        // Read the 8-byte n_sids header which includes the format indicator
4472        let raw_n_sids: RawInt<LeU64> = sid_buf.parse(())?;
4473        let format_byte = (raw_n_sids.0 >> 56) as u8;
4474        // MySQL encodes n_sids differently for tagged vs untagged:
4475        //   Untagged: n_sids in bits [55:0]
4476        //   Tagged:   n_sids in bits [55:8]
4477        let n_sids = if format_byte == 0x01 {
4478            (raw_n_sids.0 & 0x00FF_FFFF_FFFF_FF00) >> 8
4479        } else {
4480            raw_n_sids.0 & 0x00FF_FFFF_FFFF_FFFF
4481        };
4482
4483        let mut sids = Vec::with_capacity(n_sids as usize);
4484        match format_byte {
4485            0x01 => {
4486                // Tagged (TSID) format
4487                for _ in 0..n_sids {
4488                    sids.push(Sid::deserialize_tagged(&mut sid_buf)?);
4489                }
4490            }
4491            _ => {
4492                // Untagged (legacy) format — format_byte 0x00 or any unknown value
4493                for _ in 0..n_sids {
4494                    sids.push(sid_buf.parse::<Sid<'de>>(())?);
4495                }
4496            }
4497        }
4498
4499        Ok(Self {
4500            header,
4501            flags,
4502            server_id,
4503            filename,
4504            pos,
4505            sid_block: Seq::new(sids),
4506        })
4507    }
4508}
4509
4510define_header!(
4511    SemiSyncAckPacketPacketHeader,
4512    InvalidSemiSyncAckPacketPacketHeader("Invalid semi-sync ack packet header"),
4513    0xEF
4514);
4515
4516/// Each Semi Sync Binlog Event with the `SEMI_SYNC_ACK_REQ` flag set the slave has to acknowledge
4517/// with Semi-Sync ACK packet.
4518pub struct SemiSyncAckPacket<'a> {
4519    header: SemiSyncAckPacketPacketHeader,
4520    position: RawInt<LeU64>,
4521    filename: RawBytes<'a, EofBytes>,
4522}
4523
4524impl<'a> SemiSyncAckPacket<'a> {
4525    pub fn new(position: u64, filename: impl Into<Cow<'a, [u8]>>) -> Self {
4526        Self {
4527            header: Default::default(),
4528            position: RawInt::new(position),
4529            filename: RawBytes::new(filename),
4530        }
4531    }
4532
4533    /// Sets the `position` field value.
4534    pub fn with_position(mut self, position: u64) -> Self {
4535        self.position.0 = position;
4536        self
4537    }
4538
4539    /// Sets the `filename` field value.
4540    pub fn with_filename(mut self, filename: impl Into<Cow<'a, [u8]>>) -> Self {
4541        self.filename = RawBytes::new(filename);
4542        self
4543    }
4544
4545    /// Returns the `position` field value.
4546    pub fn position(&self) -> u64 {
4547        self.position.0
4548    }
4549
4550    /// Returns the raw `filename` field value.
4551    pub fn filename_raw(&self) -> &[u8] {
4552        self.filename.as_bytes()
4553    }
4554
4555    /// Returns the `filename` field value as a string (lossy converted).
4556    pub fn filename(&self) -> Cow<'_, str> {
4557        self.filename.as_str()
4558    }
4559}
4560
4561impl MySerialize for SemiSyncAckPacket<'_> {
4562    fn serialize(&self, buf: &mut Vec<u8>) {
4563        self.header.serialize(&mut *buf);
4564        self.position.serialize(&mut *buf);
4565        self.filename.serialize(&mut *buf);
4566    }
4567}
4568
4569impl<'de> MyDeserialize<'de> for SemiSyncAckPacket<'de> {
4570    const SIZE: Option<usize> = None;
4571    type Ctx = ();
4572
4573    fn deserialize((): Self::Ctx, buf: &mut ParseBuf<'de>) -> io::Result<Self> {
4574        let mut sbuf: ParseBuf<'_> = buf.parse(9)?;
4575        Ok(Self {
4576            header: sbuf.parse_unchecked(())?,
4577            position: sbuf.parse_unchecked(())?,
4578            filename: buf.parse(())?,
4579        })
4580    }
4581}
4582
4583#[cfg(test)]
4584mod test {
4585    use super::*;
4586    use crate::{
4587        constants::{CapabilityFlags, ColumnFlags, ColumnType, StatusFlags},
4588        proto::{MyDeserialize, MySerialize},
4589    };
4590
4591    proptest::proptest! {
4592        #[test]
4593        fn com_table_dump_roundtrip(database: Vec<u8>, table: Vec<u8>) {
4594            let cmd = ComTableDump::new(database, table);
4595
4596            let mut output = Vec::new();
4597            cmd.serialize(&mut output);
4598
4599            assert_eq!(cmd, ComTableDump::deserialize((), &mut ParseBuf(&output[..]))?);
4600        }
4601
4602        #[test]
4603        fn com_binlog_dump_roundtrip(
4604            server_id: u32,
4605            filename: Vec<u8>,
4606            pos: u32,
4607            flags: u16,
4608        ) {
4609            let cmd = ComBinlogDump::new(server_id)
4610                .with_filename(filename)
4611                .with_pos(pos)
4612                .with_flags(crate::packets::BinlogDumpFlags::from_bits_truncate(flags));
4613
4614            let mut output = Vec::new();
4615            cmd.serialize(&mut output);
4616
4617            assert_eq!(cmd, ComBinlogDump::deserialize((), &mut ParseBuf(&output[..]))?);
4618        }
4619
4620        #[test]
4621        fn com_register_slave_roundtrip(
4622            server_id: u32,
4623            hostname in r"\w{0,256}",
4624            user in r"\w{0,256}",
4625            password in r"\w{0,256}",
4626            port: u16,
4627            replication_rank: u32,
4628            master_id: u32,
4629        ) {
4630            let cmd = ComRegisterSlave::new(server_id)
4631                .with_hostname(hostname.as_bytes())
4632                .with_user(user.as_bytes())
4633                .with_password(password.as_bytes())
4634                .with_port(port)
4635                .with_replication_rank(replication_rank)
4636                .with_master_id(master_id);
4637
4638            let mut output = Vec::new();
4639            cmd.serialize(&mut output);
4640            let parsed = ComRegisterSlave::deserialize((), &mut ParseBuf(&output[..]))?;
4641
4642            if hostname.len() > 255 || user.len() > 255 || password.len() > 255 {
4643                assert_ne!(cmd, parsed);
4644            } else {
4645                assert_eq!(cmd, parsed);
4646            }
4647        }
4648
4649        #[test]
4650        fn com_binlog_dump_gtid_roundtrip(
4651            flags: u16,
4652            server_id: u32,
4653            filename: Vec<u8>,
4654            pos: u64,
4655            n_sid_blocks in 0_u64..1024,
4656        ) {
4657            let mut cmd = ComBinlogDumpGtid::new(server_id)
4658                .with_filename(filename)
4659                .with_pos(pos)
4660                .with_flags(crate::packets::BinlogDumpFlags::from_bits_truncate(flags));
4661
4662            let mut sids = Vec::new();
4663            for i in 0..n_sid_blocks {
4664                let mut block = Sid::new([i as u8; 16]);
4665                for j in 0..i {
4666                    block = block.with_interval(GnoInterval::new(i, j));
4667                }
4668                sids.push(block);
4669            }
4670
4671            cmd = cmd.with_sids(sids);
4672
4673            let mut output = Vec::new();
4674            cmd.serialize(&mut output);
4675
4676            assert_eq!(cmd, ComBinlogDumpGtid::deserialize((), &mut ParseBuf(&output[..]))?);
4677        }
4678
4679        #[test]
4680        fn com_binlog_dump_gtid_tagged_roundtrip(
4681            flags: u16,
4682            server_id: u32,
4683            filename: Vec<u8>,
4684            pos: u64,
4685            n_sid_blocks in 0_u64..128,
4686        ) {
4687            let mut cmd = ComBinlogDumpGtid::new(server_id)
4688                .with_filename(filename)
4689                .with_pos(pos)
4690                .with_flags(crate::packets::BinlogDumpFlags::from_bits_truncate(flags));
4691
4692            let mut sids = Vec::new();
4693            for i in 0..n_sid_blocks {
4694                let mut block = Sid::new([i as u8; 16]);
4695                // Tag every other Sid
4696                if i % 2 == 0 {
4697                    block = block.with_tag(Tag::new(format!("tag_{}", i)).unwrap());
4698                }
4699                for j in 0..i {
4700                    block = block.with_interval(GnoInterval::new(i, j));
4701                }
4702                sids.push(block);
4703            }
4704
4705            cmd = cmd.with_sids(sids);
4706
4707            let mut output = Vec::new();
4708            cmd.serialize(&mut output);
4709
4710            assert_eq!(cmd, ComBinlogDumpGtid::deserialize((), &mut ParseBuf(&output[..]))?);
4711        }
4712    }
4713
4714    #[test]
4715    fn sid_display_roundtrip() {
4716        // Untagged SID
4717        let input = "3e11fa47-71ca-11e1-9e33-c80aa9429562:1-5:10-20";
4718        let sid: Sid = input.parse().unwrap();
4719        assert_eq!(sid.to_string(), input);
4720
4721        // Tagged SID (MySQL 8.4+)
4722        let input_tagged = "3e11fa47-71ca-11e1-9e33-c80aa9429562:domain_a:1-5:10-20";
4723        let sid_tagged: Sid = input_tagged.parse().unwrap();
4724        assert!(sid_tagged.tag().is_some());
4725        assert_eq!(sid_tagged.to_string(), input_tagged);
4726
4727        // Single interval
4728        let input_single = "3e11fa47-71ca-11e1-9e33-c80aa9429562:42";
4729        let sid_single: Sid = input_single.parse().unwrap();
4730        assert_eq!(sid_single.to_string(), input_single);
4731
4732        // Tagged single interval
4733        let input_tagged_single = "3e11fa47-71ca-11e1-9e33-c80aa9429562:app:7";
4734        let sid_tagged_single: Sid = input_tagged_single.parse().unwrap();
4735        assert_eq!(sid_tagged_single.to_string(), input_tagged_single);
4736    }
4737
4738    #[test]
4739    fn should_parse_local_infile_packet() {
4740        const LIP: &[u8] = b"\xfbfile_name";
4741
4742        let lip = LocalInfilePacket::deserialize((), &mut ParseBuf(LIP)).unwrap();
4743        assert_eq!(lip.file_name_str(), "file_name");
4744    }
4745
4746    #[test]
4747    fn should_parse_stmt_packet() {
4748        const SP: &[u8] = b"\x00\x01\x00\x00\x00\x01\x00\x02\x00\x00\x00\x00";
4749        const SP_2: &[u8] = b"\x00\x01\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00";
4750
4751        let sp = StmtPacket::deserialize((), &mut ParseBuf(SP)).unwrap();
4752        assert_eq!(sp.statement_id(), 0x01);
4753        assert_eq!(sp.num_columns(), 0x01);
4754        assert_eq!(sp.num_params(), 0x02);
4755        assert_eq!(sp.warning_count(), 0x00);
4756
4757        let sp = StmtPacket::deserialize((), &mut ParseBuf(SP_2)).unwrap();
4758        assert_eq!(sp.statement_id(), 0x01);
4759        assert_eq!(sp.num_columns(), 0x00);
4760        assert_eq!(sp.num_params(), 0x00);
4761        assert_eq!(sp.warning_count(), 0x00);
4762    }
4763
4764    #[test]
4765    fn should_parse_handshake_packet() {
4766        const HSP: &[u8] = b"\x0a5.5.5-10.0.17-MariaDB-log\x00\x0b\x00\
4767                             \x00\x00\x64\x76\x48\x40\x49\x2d\x43\x4a\x00\xff\xf7\x08\x02\x00\
4768                             \x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x2a\x34\x64\
4769                             \x7c\x63\x5a\x77\x6b\x34\x5e\x5d\x3a\x00";
4770
4771        const HSP_2: &[u8] = b"\x0a\x35\x2e\x36\x2e\x34\x2d\x6d\x37\x2d\x6c\x6f\
4772                               \x67\x00\x56\x0a\x00\x00\x52\x42\x33\x76\x7a\x26\x47\x72\x00\xff\
4773                               \xff\x08\x02\x00\x0f\xc0\x15\x00\x00\x00\x00\x00\x00\x00\x00\x00\
4774                               \x00\x2b\x79\x44\x26\x2f\x5a\x5a\x33\x30\x35\x5a\x47\x00\x6d\x79\
4775                               \x73\x71\x6c\x5f\x6e\x61\x74\x69\x76\x65\x5f\x70\x61\x73\x73\x77\
4776                               \x6f\x72\x64\x00";
4777
4778        const HSP_3: &[u8] = b"\x0a\x35\x2e\x36\x2e\x34\x2d\x6d\x37\x2d\x6c\x6f\
4779                                \x67\x00\x56\x0a\x00\x00\x52\x42\x33\x76\x7a\x26\x47\x72\x00\xff\
4780                                \xff\x08\x02\x00\x0f\xc0\x15\x00\x00\x00\x00\x00\x00\x00\x00\x00\
4781                                \x00\x2b\x79\x44\x26\x2f\x5a\x5a\x33\x30\x35\x5a\x47\x00\x6d\x79\
4782                                \x73\x71\x6c\x5f\x6e\x61\x74\x69\x76\x65\x5f\x70\x61\x73\x73\x77\
4783                                \x6f\x72\x64\x00";
4784
4785        let hsp = HandshakePacket::deserialize((), &mut ParseBuf(HSP)).unwrap();
4786        assert_eq!(hsp.protocol_version(), 0x0a);
4787        assert_eq!(hsp.server_version_str(), "5.5.5-10.0.17-MariaDB-log");
4788        assert_eq!(hsp.server_version_parsed(), Some((5, 5, 5)));
4789        assert_eq!(hsp.maria_db_server_version_parsed(), Some((10, 0, 17)));
4790        assert_eq!(hsp.connection_id(), 0x0b);
4791        assert_eq!(hsp.scramble_1_ref(), b"dvH@I-CJ");
4792        assert_eq!(
4793            hsp.capabilities(),
4794            CapabilityFlags::from_bits_truncate(0xf7ff)
4795        );
4796        assert_eq!(hsp.default_collation(), 0x08);
4797        assert_eq!(hsp.status_flags(), StatusFlags::from_bits_truncate(0x0002));
4798        assert_eq!(hsp.scramble_2_ref(), Some(&b"*4d|cZwk4^]:\x00"[..]));
4799        assert_eq!(hsp.auth_plugin_name_ref(), None);
4800
4801        let mut output = Vec::new();
4802        hsp.serialize(&mut output);
4803        assert_eq!(&output, HSP);
4804
4805        let hsp = HandshakePacket::deserialize((), &mut ParseBuf(HSP_2)).unwrap();
4806        assert_eq!(hsp.protocol_version(), 0x0a);
4807        assert_eq!(hsp.server_version_str(), "5.6.4-m7-log");
4808        assert_eq!(hsp.server_version_parsed(), Some((5, 6, 4)));
4809        assert_eq!(hsp.maria_db_server_version_parsed(), None);
4810        assert_eq!(hsp.connection_id(), 0x0a56);
4811        assert_eq!(hsp.scramble_1_ref(), b"RB3vz&Gr");
4812        assert_eq!(
4813            hsp.capabilities(),
4814            CapabilityFlags::from_bits_truncate(0xc00fffff)
4815        );
4816        assert_eq!(hsp.default_collation(), 0x08);
4817        assert_eq!(hsp.status_flags(), StatusFlags::from_bits_truncate(0x0002));
4818        assert_eq!(hsp.scramble_2_ref(), Some(&b"+yD&/ZZ305ZG\0"[..]));
4819        assert_eq!(
4820            hsp.auth_plugin_name_ref(),
4821            Some(&b"mysql_native_password"[..])
4822        );
4823
4824        let mut output = Vec::new();
4825        hsp.serialize(&mut output);
4826        assert_eq!(&output, HSP_2);
4827
4828        let hsp = HandshakePacket::deserialize((), &mut ParseBuf(HSP_3)).unwrap();
4829        assert_eq!(hsp.protocol_version(), 0x0a);
4830        assert_eq!(hsp.server_version_str(), "5.6.4-m7-log");
4831        assert_eq!(hsp.server_version_parsed(), Some((5, 6, 4)));
4832        assert_eq!(hsp.maria_db_server_version_parsed(), None);
4833        assert_eq!(hsp.connection_id(), 0x0a56);
4834        assert_eq!(hsp.scramble_1_ref(), b"RB3vz&Gr");
4835        assert_eq!(
4836            hsp.capabilities(),
4837            CapabilityFlags::from_bits_truncate(0xc00fffff)
4838        );
4839        assert_eq!(hsp.default_collation(), 0x08);
4840        assert_eq!(hsp.status_flags(), StatusFlags::from_bits_truncate(0x0002));
4841        assert_eq!(hsp.scramble_2_ref(), Some(&b"+yD&/ZZ305ZG\0"[..]));
4842        assert_eq!(
4843            hsp.auth_plugin_name_ref(),
4844            Some(&b"mysql_native_password"[..])
4845        );
4846
4847        let mut output = Vec::new();
4848        hsp.serialize(&mut output);
4849        assert_eq!(&output, HSP_3);
4850    }
4851
4852    #[test]
4853    fn should_parse_err_packet() {
4854        const ERR_PACKET: &[u8] = b"\xff\x48\x04\x23\x48\x59\x30\x30\x30\x4e\x6f\x20\x74\x61\x62\
4855        \x6c\x65\x73\x20\x75\x73\x65\x64";
4856        const ERR_PACKET_NO_STATE: &[u8] = b"\xff\x10\x04\x54\x6f\x6f\x20\x6d\x61\x6e\x79\x20\x63\
4857        \x6f\x6e\x6e\x65\x63\x74\x69\x6f\x6e\x73";
4858        const PROGRESS_PACKET: &[u8] = b"\xff\xff\xff\x01\x01\x0a\xcc\x5b\x00\x0astage name";
4859
4860        let err_packet = ErrPacket::deserialize(
4861            CapabilityFlags::CLIENT_PROTOCOL_41,
4862            &mut ParseBuf(ERR_PACKET),
4863        )
4864        .unwrap();
4865        let err_packet = err_packet.server_error();
4866        assert_eq!(err_packet.error_code(), 1096);
4867        assert_eq!(err_packet.sql_state_ref().unwrap().as_str(), "HY000");
4868        assert_eq!(err_packet.message_str(), "No tables used");
4869
4870        let err_packet =
4871            ErrPacket::deserialize(CapabilityFlags::empty(), &mut ParseBuf(ERR_PACKET_NO_STATE))
4872                .unwrap();
4873        let server_error = err_packet.server_error();
4874        assert_eq!(server_error.error_code(), 1040);
4875        assert_eq!(server_error.sql_state_ref(), None);
4876        assert_eq!(server_error.message_str(), "Too many connections");
4877
4878        let err_packet = ErrPacket::deserialize(
4879            CapabilityFlags::CLIENT_PROGRESS_OBSOLETE,
4880            &mut ParseBuf(PROGRESS_PACKET),
4881        )
4882        .unwrap();
4883        assert!(err_packet.is_progress_report());
4884        let progress_report = err_packet.progress_report();
4885        assert_eq!(progress_report.stage(), 1);
4886        assert_eq!(progress_report.max_stage(), 10);
4887        assert_eq!(progress_report.progress(), 23500);
4888        assert_eq!(progress_report.stage_info_str(), "stage name");
4889    }
4890
4891    #[test]
4892    fn should_parse_column_packet() {
4893        const COLUMN_PACKET: &[u8] = b"\x03def\x06schema\x05table\x09org_table\x04name\
4894              \x08org_name\x0c\x21\x00\x0F\x00\x00\x00\x00\x01\x00\x08\x00\x00";
4895        let column = Column::deserialize((), &mut ParseBuf(COLUMN_PACKET)).unwrap();
4896        assert_eq!(column.schema_str(), "schema");
4897        assert_eq!(column.table_str(), "table");
4898        assert_eq!(column.org_table_str(), "org_table");
4899        assert_eq!(column.name_str(), "name");
4900        assert_eq!(column.org_name_str(), "org_name");
4901        assert_eq!(
4902            column.character_set(),
4903            CollationId::UTF8MB3_GENERAL_CI as u16
4904        );
4905        assert_eq!(column.column_length(), 15);
4906        assert_eq!(column.column_type(), ColumnType::MYSQL_TYPE_DECIMAL);
4907        assert_eq!(column.flags(), ColumnFlags::NOT_NULL_FLAG);
4908        assert_eq!(column.decimals(), 8);
4909    }
4910
4911    #[test]
4912    fn should_parse_auth_switch_request() {
4913        const PAYLOAD: &[u8] = b"\xfe\x6d\x79\x73\x71\x6c\x5f\x6e\x61\x74\x69\x76\x65\x5f\x70\x61\
4914                                 \x73\x73\x77\x6f\x72\x64\x00\x7a\x51\x67\x34\x69\x36\x6f\x4e\x79\
4915                                 \x36\x3d\x72\x48\x4e\x2f\x3e\x2d\x62\x29\x41\x00";
4916        let packet = AuthSwitchRequest::deserialize((), &mut ParseBuf(PAYLOAD)).unwrap();
4917        assert_eq!(packet.auth_plugin().as_bytes(), b"mysql_native_password",);
4918        assert_eq!(packet.plugin_data(), b"zQg4i6oNy6=rHN/>-b)A",)
4919    }
4920
4921    #[test]
4922    fn should_parse_auth_more_data() {
4923        const PAYLOAD: &[u8] = b"\x01\x04";
4924        let packet = AuthMoreData::deserialize((), &mut ParseBuf(PAYLOAD)).unwrap();
4925        assert_eq!(packet.data(), b"\x04",);
4926    }
4927
4928    #[test]
4929    fn should_parse_ok_packet() {
4930        const PLAIN_OK: &[u8] = b"\x00\x01\x00\x02\x00\x00\x00";
4931        const RESULT_SET_TERMINATOR: &[u8] = &[
4932            0xfe, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x42, 0x52, 0x65, 0x61, 0x64, 0x20, 0x31,
4933            0x20, 0x72, 0x6f, 0x77, 0x73, 0x2c, 0x20, 0x31, 0x2e, 0x30, 0x30, 0x20, 0x42, 0x20,
4934            0x69, 0x6e, 0x20, 0x30, 0x2e, 0x30, 0x30, 0x32, 0x20, 0x73, 0x65, 0x63, 0x2e, 0x2c,
4935            0x20, 0x36, 0x31, 0x31, 0x2e, 0x33, 0x34, 0x20, 0x72, 0x6f, 0x77, 0x73, 0x2f, 0x73,
4936            0x65, 0x63, 0x2e, 0x2c, 0x20, 0x36, 0x31, 0x31, 0x2e, 0x33, 0x34, 0x20, 0x42, 0x2f,
4937            0x73, 0x65, 0x63, 0x2e,
4938        ];
4939        const SESSION_STATE_SYS_VAR_OK: &[u8] =
4940            b"\x00\x00\x00\x02\x40\x00\x00\x00\x11\x00\x0f\x0a\x61\
4941              \x75\x74\x6f\x63\x6f\x6d\x6d\x69\x74\x03\x4f\x46\x46";
4942        const SESSION_STATE_SCHEMA_OK: &[u8] =
4943            b"\x00\x00\x00\x02\x40\x00\x00\x00\x07\x01\x05\x04\x74\x65\x73\x74";
4944        const SESSION_STATE_TRACK_OK: &[u8] =
4945            b"\x00\x00\x00\x02\x40\x00\x00\x00\x04\x02\x02\x01\x31";
4946        const EOF: &[u8] = b"\xfe\x00\x00\x02\x00";
4947
4948        // packet starting with 0x00 is not an ok packet if it terminates a result set
4949        OkPacketDeserializer::<ResultSetTerminator>::deserialize(
4950            CapabilityFlags::empty(),
4951            &mut ParseBuf(PLAIN_OK),
4952        )
4953        .unwrap_err();
4954
4955        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<CommonOkPacket>::deserialize(
4956            CapabilityFlags::empty(),
4957            &mut ParseBuf(PLAIN_OK),
4958        )
4959        .unwrap()
4960        .into();
4961        assert_eq!(ok_packet.affected_rows(), 1);
4962        assert_eq!(ok_packet.last_insert_id(), None);
4963        assert_eq!(
4964            ok_packet.status_flags(),
4965            StatusFlags::SERVER_STATUS_AUTOCOMMIT
4966        );
4967        assert_eq!(ok_packet.warnings(), 0);
4968        assert_eq!(ok_packet.info_ref(), None);
4969        assert_eq!(ok_packet.session_state_info_ref(), None);
4970
4971        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<CommonOkPacket>::deserialize(
4972            CapabilityFlags::CLIENT_SESSION_TRACK,
4973            &mut ParseBuf(PLAIN_OK),
4974        )
4975        .unwrap()
4976        .into();
4977        assert_eq!(ok_packet.affected_rows(), 1);
4978        assert_eq!(ok_packet.last_insert_id(), None);
4979        assert_eq!(
4980            ok_packet.status_flags(),
4981            StatusFlags::SERVER_STATUS_AUTOCOMMIT
4982        );
4983        assert_eq!(ok_packet.warnings(), 0);
4984        assert_eq!(ok_packet.info_ref(), None);
4985        assert_eq!(ok_packet.session_state_info_ref(), None);
4986
4987        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<ResultSetTerminator>::deserialize(
4988            CapabilityFlags::CLIENT_SESSION_TRACK,
4989            &mut ParseBuf(RESULT_SET_TERMINATOR),
4990        )
4991        .unwrap()
4992        .into();
4993        assert_eq!(ok_packet.affected_rows(), 0);
4994        assert_eq!(ok_packet.last_insert_id(), None);
4995        assert_eq!(ok_packet.status_flags(), StatusFlags::empty());
4996        assert_eq!(ok_packet.warnings(), 0);
4997        assert_eq!(
4998            ok_packet.info_str(),
4999            Some(Cow::Borrowed(
5000                "Read 1 rows, 1.00 B in 0.002 sec., 611.34 rows/sec., 611.34 B/sec."
5001            ))
5002        );
5003        assert_eq!(ok_packet.session_state_info_ref(), None);
5004
5005        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<CommonOkPacket>::deserialize(
5006            CapabilityFlags::CLIENT_SESSION_TRACK,
5007            &mut ParseBuf(SESSION_STATE_SYS_VAR_OK),
5008        )
5009        .unwrap()
5010        .into();
5011        assert_eq!(ok_packet.affected_rows(), 0);
5012        assert_eq!(ok_packet.last_insert_id(), None);
5013        assert_eq!(
5014            ok_packet.status_flags(),
5015            StatusFlags::SERVER_STATUS_AUTOCOMMIT | StatusFlags::SERVER_SESSION_STATE_CHANGED
5016        );
5017        assert_eq!(ok_packet.warnings(), 0);
5018        assert_eq!(ok_packet.info_ref(), None);
5019        let session_state_info = ok_packet.session_state_info().unwrap().pop().unwrap();
5020
5021        match session_state_info.decode().unwrap() {
5022            SessionStateChange::SystemVariables(mut vals) => {
5023                let val = vals.pop().unwrap();
5024                assert_eq!(val.name_bytes(), b"autocommit");
5025                assert_eq!(val.value_bytes(), b"OFF");
5026                assert!(vals.is_empty());
5027            }
5028            _ => panic!(),
5029        }
5030
5031        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<CommonOkPacket>::deserialize(
5032            CapabilityFlags::CLIENT_SESSION_TRACK,
5033            &mut ParseBuf(SESSION_STATE_SCHEMA_OK),
5034        )
5035        .unwrap()
5036        .into();
5037        assert_eq!(ok_packet.affected_rows(), 0);
5038        assert_eq!(ok_packet.last_insert_id(), None);
5039        assert_eq!(
5040            ok_packet.status_flags(),
5041            StatusFlags::SERVER_STATUS_AUTOCOMMIT | StatusFlags::SERVER_SESSION_STATE_CHANGED
5042        );
5043        assert_eq!(ok_packet.warnings(), 0);
5044        assert_eq!(ok_packet.info_ref(), None);
5045        let session_state_info = ok_packet.session_state_info().unwrap().pop().unwrap();
5046        match session_state_info.decode().unwrap() {
5047            SessionStateChange::Schema(schema) => assert_eq!(schema.as_bytes(), b"test"),
5048            _ => panic!(),
5049        }
5050
5051        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<CommonOkPacket>::deserialize(
5052            CapabilityFlags::CLIENT_SESSION_TRACK,
5053            &mut ParseBuf(SESSION_STATE_TRACK_OK),
5054        )
5055        .unwrap()
5056        .into();
5057        assert_eq!(ok_packet.affected_rows(), 0);
5058        assert_eq!(ok_packet.last_insert_id(), None);
5059        assert_eq!(
5060            ok_packet.status_flags(),
5061            StatusFlags::SERVER_STATUS_AUTOCOMMIT | StatusFlags::SERVER_SESSION_STATE_CHANGED
5062        );
5063        assert_eq!(ok_packet.warnings(), 0);
5064        assert_eq!(ok_packet.info_ref(), None);
5065        let session_state_info = ok_packet.session_state_info().unwrap().pop().unwrap();
5066        assert_eq!(
5067            session_state_info.decode().unwrap(),
5068            SessionStateChange::IsTracked(true),
5069        );
5070
5071        let ok_packet: OkPacket<'_> = OkPacketDeserializer::<OldEofPacket>::deserialize(
5072            CapabilityFlags::CLIENT_SESSION_TRACK,
5073            &mut ParseBuf(EOF),
5074        )
5075        .unwrap()
5076        .into();
5077        assert_eq!(ok_packet.affected_rows(), 0);
5078        assert_eq!(ok_packet.last_insert_id(), None);
5079        assert_eq!(
5080            ok_packet.status_flags(),
5081            StatusFlags::SERVER_STATUS_AUTOCOMMIT
5082        );
5083        assert_eq!(ok_packet.warnings(), 0);
5084        assert_eq!(ok_packet.info_ref(), None);
5085        assert_eq!(ok_packet.session_state_info_ref(), None);
5086    }
5087
5088    #[test]
5089    fn should_build_handshake_response() {
5090        let flags_without_db_name = CapabilityFlags::from_bits_truncate(0x81aea205);
5091        let response = HandshakeResponse::new(
5092            Some(&[][..]),
5093            (5u16, 5, 5),
5094            Some(&b"root"[..]),
5095            None::<&'static [u8]>,
5096            Some(AuthPlugin::MysqlNativePassword),
5097            flags_without_db_name,
5098            None,
5099            1_u32.to_be(),
5100        );
5101        let mut actual = Vec::new();
5102        response.serialize(&mut actual);
5103
5104        let expected: Vec<u8> = [
5105            0x05, 0xa2, 0xae, 0x81, // client capabilities
5106            0x00, 0x00, 0x00, 0x01, // max packet
5107            0x2d, // charset
5108            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5109            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // reserved
5110            0x72, 0x6f, 0x6f, 0x74, 0x00, // username=root
5111            0x00, // blank scramble
5112            0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5113            0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, // mysql_native_password
5114        ]
5115        .to_vec();
5116
5117        assert_eq!(expected, actual);
5118
5119        let flags_with_db_name = flags_without_db_name | CapabilityFlags::CLIENT_CONNECT_WITH_DB;
5120        let response = HandshakeResponse::new(
5121            Some(&[][..]),
5122            (5u16, 5, 5),
5123            Some(&b"root"[..]),
5124            Some(&b"mydb"[..]),
5125            Some(AuthPlugin::MysqlNativePassword),
5126            flags_with_db_name,
5127            None,
5128            1_u32.to_be(),
5129        );
5130        let mut actual = Vec::new();
5131        response.serialize(&mut actual);
5132
5133        let expected: Vec<u8> = [
5134            0x0d, 0xa2, 0xae, 0x81, // client capabilities
5135            0x00, 0x00, 0x00, 0x01, // max packet
5136            0x2d, // charset
5137            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5138            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // reserved
5139            0x72, 0x6f, 0x6f, 0x74, 0x00, // username=root
5140            0x00, // blank scramble
5141            0x6d, 0x79, 0x64, 0x62, 0x00, // db name
5142            0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5143            0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, // mysql_native_password
5144        ]
5145        .to_vec();
5146
5147        assert_eq!(expected, actual);
5148
5149        let response = HandshakeResponse::new(
5150            Some(&[][..]),
5151            (5u16, 5, 5),
5152            Some(&b"root"[..]),
5153            Some(&b"mydb"[..]),
5154            Some(AuthPlugin::MysqlNativePassword),
5155            flags_without_db_name,
5156            None,
5157            1_u32.to_be(),
5158        );
5159        let mut actual = Vec::new();
5160        response.serialize(&mut actual);
5161        assert_eq!(expected, actual);
5162
5163        let response = HandshakeResponse::new(
5164            Some(&[][..]),
5165            (5u16, 5, 5),
5166            Some(&b"root"[..]),
5167            Some(&[][..]),
5168            Some(AuthPlugin::MysqlNativePassword),
5169            flags_with_db_name,
5170            None,
5171            1_u32.to_be(),
5172        );
5173        let mut actual = Vec::new();
5174        response.serialize(&mut actual);
5175
5176        let expected: Vec<u8> = [
5177            0x0d, 0xa2, 0xae, 0x81, // client capabilities
5178            0x00, 0x00, 0x00, 0x01, // max packet
5179            0x2d, // charset
5180            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5181            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // reserved
5182            0x72, 0x6f, 0x6f, 0x74, 0x00, // username=root
5183            0x00, // blank db_name
5184            0x00, // blank scramble
5185            0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5186            0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, // mysql_native_password
5187        ]
5188        .to_vec();
5189        assert_eq!(expected, actual);
5190    }
5191
5192    #[test]
5193    fn should_parse_handshake_packet_with_mariadb_ext_capabilities() {
5194        const HSP: &[u8] = b"\x0a5.5.5-11.4.7-MariaDB-log\x00\x0b\x00\
5195                             \x00\x00\x64\x76\x48\x40\x49\x2d\x43\x4a\x00\xff\xf7\x08\x02\x00\
5196                             \x00\x00\x00\x00\x00\x00\x00\x00\x00\x14\x00\x00\x00\x2a\x34\x64\
5197                             \x7c\x63\x5a\x77\x6b\x34\x5e\x5d\x3a\x00";
5198
5199        let hsp = HandshakePacket::deserialize((), &mut ParseBuf(HSP)).unwrap();
5200        assert_eq!(hsp.protocol_version(), 0x0a);
5201        assert_eq!(hsp.server_version_str(), "5.5.5-11.4.7-MariaDB-log");
5202        assert_eq!(hsp.server_version_parsed(), Some((5, 5, 5)));
5203        assert_eq!(hsp.maria_db_server_version_parsed(), Some((11, 4, 7)));
5204        assert_eq!(hsp.connection_id(), 0x0b);
5205        assert_eq!(hsp.scramble_1_ref(), b"dvH@I-CJ");
5206        assert_eq!(
5207            hsp.capabilities(),
5208            CapabilityFlags::from_bits_truncate(0xf7ff)
5209        );
5210        assert_eq!(hsp.default_collation(), 0x08);
5211        assert_eq!(hsp.status_flags(), StatusFlags::from_bits_truncate(0x0002));
5212        assert_eq!(hsp.scramble_2_ref(), Some(&b"*4d|cZwk4^]:\x00"[..]));
5213        assert_eq!(hsp.auth_plugin_name_ref(), None);
5214        assert_eq!(
5215            hsp.mariadb_ext_capabilities(),
5216            MariadbCapabilities::MARIADB_CLIENT_CACHE_METADATA
5217                | MariadbCapabilities::MARIADB_CLIENT_STMT_BULK_OPERATIONS
5218        );
5219        let mut output = Vec::new();
5220        hsp.serialize(&mut output);
5221        assert_eq!(&output, HSP);
5222    }
5223
5224    #[test]
5225    fn should_build_handshake_response_with_mariadb_capabilities() {
5226        let flags_without_db_name = CapabilityFlags::from_bits_truncate(0x81aea205);
5227        let response = HandshakeResponse::new(
5228            Some(&[][..]),
5229            (5u16, 5, 5),
5230            Some(&b"root"[..]),
5231            None::<&'static [u8]>,
5232            Some(AuthPlugin::MysqlNativePassword),
5233            flags_without_db_name,
5234            None,
5235            1_u32.to_be(),
5236        )
5237        .with_mariadb_ext_capabilities(
5238            MariadbCapabilities::MARIADB_CLIENT_CACHE_METADATA
5239                | MariadbCapabilities::MARIADB_CLIENT_STMT_BULK_OPERATIONS,
5240        );
5241        let mut actual = Vec::new();
5242        response.serialize(&mut actual);
5243
5244        let expected: Vec<u8> = [
5245            0x05, 0xa2, 0xae, 0x81, // client capabilities
5246            0x00, 0x00, 0x00, 0x01, // max packet
5247            0x2d, // charset
5248            0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5249            0x00, 0x00, 0x00, 0x00, 0x00, // reserved
5250            0x14, 0x00, 0x00, 0x00, // mariadb capabilities
5251            0x72, 0x6f, 0x6f, 0x74, 0x00, // username=root
5252            0x00, // blank scramble
5253            0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5254            0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, // mysql_native_password
5255        ]
5256        .to_vec();
5257
5258        assert_eq!(expected, actual);
5259    }
5260
5261    #[test]
5262    fn parse_str_to_sid() {
5263        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:23";
5264        let sid = input.parse::<Sid<'_>>().unwrap();
5265        let expected_sid = Uuid::parse_str("3E11FA47-71CA-11E1-9E33-C80AA9429562").unwrap();
5266        assert_eq!(sid.uuid, *expected_sid.as_bytes());
5267        assert_eq!(sid.intervals.len(), 1);
5268        assert_eq!(sid.intervals[0].start.0, 23);
5269        assert_eq!(sid.intervals[0].end.0, 24);
5270
5271        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:1-5:10-15";
5272        let sid = input.parse::<Sid<'_>>().unwrap();
5273        assert_eq!(sid.uuid, *expected_sid.as_bytes());
5274        assert_eq!(sid.intervals.len(), 2);
5275        assert_eq!(sid.intervals[0].start.0, 1);
5276        assert_eq!(sid.intervals[0].end.0, 6);
5277        assert_eq!(sid.intervals[1].start.0, 10);
5278        assert_eq!(sid.intervals[1].end.0, 16);
5279
5280        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562";
5281        let e = input.parse::<Sid<'_>>().unwrap_err();
5282        assert_eq!(
5283            e.to_string(),
5284            "invalid sid format: 3E11FA47-71CA-11E1-9E33-C80AA9429562".to_string()
5285        );
5286
5287        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:1-5:10-15:20-";
5288        let e = input.parse::<Sid<'_>>().unwrap_err();
5289        assert_eq!(e.to_string(), "invalid GnoInterval format: 3E11FA47-71CA-11E1-9E33-C80AA9429562:1-5:10-15:20-, error: cannot parse integer from empty string".to_string());
5290
5291        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:1-5:1aaa";
5292        let e = input.parse::<Sid<'_>>().unwrap_err();
5293        assert_eq!(e.to_string(), "invalid GnoInterval format: 3E11FA47-71CA-11E1-9E33-C80AA9429562:1-5:1aaa, error: invalid digit found in string".to_string());
5294
5295        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:0-3";
5296        let e = input.parse::<Sid<'_>>().unwrap_err();
5297        assert_eq!(e.to_string(), "Gno can't be zero".to_string());
5298
5299        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:4-3";
5300        let e = input.parse::<Sid<'_>>().unwrap_err();
5301        assert_eq!(
5302            e.to_string(),
5303            "start(4) >= end(4) in GnoInterval".to_string()
5304        );
5305
5306        // Tagged GTID tests (MySQL 8.4+)
5307        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:domain_1:23";
5308        let sid = input.parse::<Sid<'_>>().unwrap();
5309        let expected_sid = Uuid::parse_str("3E11FA47-71CA-11E1-9E33-C80AA9429562").unwrap();
5310        assert_eq!(sid.uuid, *expected_sid.as_bytes());
5311        assert_eq!(sid.tag.as_ref().map(|t| t.as_str()), Some("domain_1"));
5312        assert_eq!(sid.intervals.len(), 1);
5313        assert_eq!(sid.intervals[0].start.0, 23);
5314        assert_eq!(sid.intervals[0].end.0, 24);
5315
5316        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:_private_tag:1-5:10-15";
5317        let sid = input.parse::<Sid<'_>>().unwrap();
5318        assert_eq!(sid.uuid, *expected_sid.as_bytes());
5319        assert_eq!(sid.tag.as_ref().map(|t| t.as_str()), Some("_private_tag"));
5320        assert_eq!(sid.intervals.len(), 2);
5321        assert_eq!(sid.intervals[0].start.0, 1);
5322        assert_eq!(sid.intervals[0].end.0, 6);
5323        assert_eq!(sid.intervals[1].start.0, 10);
5324        assert_eq!(sid.intervals[1].end.0, 16);
5325
5326        // Invalid tag format (uppercase)
5327        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:InvalidTag:23";
5328        let e = input.parse::<Sid<'_>>().unwrap_err();
5329        assert!(e.to_string().contains("invalid first character"));
5330
5331        // Invalid tag format (starts with digit)
5332        let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:1tag:23";
5333        // This should parse as intervals starting with "1tag" which is invalid
5334        let e = input.parse::<Sid<'_>>().unwrap_err();
5335        assert!(e.to_string().contains("invalid"));
5336    }
5337
5338    #[test]
5339    fn should_parse_rsa_public_key_response_packet() {
5340        const PUBLIC_RSA_KEY_RESPONSE: &[u8] = b"\x01test";
5341
5342        let rsa_public_key_response =
5343            PublicKeyResponse::deserialize((), &mut ParseBuf(PUBLIC_RSA_KEY_RESPONSE));
5344
5345        assert!(rsa_public_key_response.is_ok());
5346        assert_eq!(rsa_public_key_response.unwrap().rsa_key(), "test");
5347    }
5348
5349    #[test]
5350    fn should_build_rsa_public_key_response_packet() {
5351        let rsa_public_key_response = PublicKeyResponse::new("test".as_bytes());
5352
5353        let mut actual = Vec::new();
5354        rsa_public_key_response.serialize(&mut actual);
5355
5356        let expected = b"\x01test".to_vec();
5357
5358        assert_eq!(expected, actual);
5359    }
5360
5361    #[test]
5362    fn valid_tags() {
5363        assert!(Tag::new("a").is_ok());
5364        assert!(Tag::new("_").is_ok());
5365        assert!(Tag::new("domain_1").is_ok());
5366        assert!(Tag::new("_private").is_ok());
5367        assert!(Tag::new("tag123").is_ok());
5368        assert!(Tag::new("a_b_c_1_2_3").is_ok());
5369        // Max length (32 chars)
5370        assert!(Tag::new("abcdefghijklmnopqrstuvwxyz123456").is_ok());
5371    }
5372
5373    #[test]
5374    fn invalid_tags() {
5375        // Empty
5376        assert_eq!(Tag::new(""), Err(InvalidTag::Empty));
5377
5378        // Too long (33 chars)
5379        assert_eq!(
5380            Tag::new("abcdefghijklmnopqrstuvwxyz1234567"),
5381            Err(InvalidTag::TooLong(33))
5382        );
5383
5384        // Invalid first character
5385        assert_eq!(Tag::new("1tag"), Err(InvalidTag::InvalidFirstChar('1')));
5386        assert_eq!(Tag::new("Tag"), Err(InvalidTag::InvalidFirstChar('T')));
5387        assert_eq!(Tag::new("-tag"), Err(InvalidTag::InvalidFirstChar('-')));
5388
5389        // Invalid characters
5390        assert_eq!(Tag::new("tag-name"), Err(InvalidTag::InvalidChar('-')));
5391        assert_eq!(Tag::new("TAG"), Err(InvalidTag::InvalidFirstChar('T')));
5392        assert_eq!(Tag::new("tAg"), Err(InvalidTag::InvalidChar('A')));
5393        assert_eq!(Tag::new("tag.name"), Err(InvalidTag::InvalidChar('.')));
5394    }
5395
5396    #[test]
5397    fn sid_tagged_wire_roundtrip() {
5398        let tag = Tag::new("domain_a".to_owned()).unwrap();
5399        let sid = Sid::new(
5400            *uuid::Uuid::parse_str("3E11FA47-71CA-11E1-9E33-C80AA9429562")
5401                .unwrap()
5402                .as_bytes(),
5403        )
5404        .with_tag(tag)
5405        .with_interval(GnoInterval::new(1, 10))
5406        .with_interval(GnoInterval::new(20, 30));
5407
5408        let mut buf = Vec::new();
5409        sid.serialize_tagged(&mut buf);
5410
5411        let mut parse_buf = ParseBuf(&buf);
5412        let decoded = Sid::deserialize_tagged(&mut parse_buf).unwrap();
5413
5414        assert!(
5415            parse_buf.is_empty(),
5416            "leftover bytes after deserialize_tagged"
5417        );
5418        assert_eq!(sid, decoded);
5419        assert_eq!(decoded.tag().unwrap().as_str(), "domain_a");
5420    }
5421
5422    #[test]
5423    fn sid_tagged_zero_tag_roundtrip() {
5424        let sid = Sid::new(
5425            *uuid::Uuid::parse_str("3E11FA47-71CA-11E1-9E33-C80AA9429562")
5426                .unwrap()
5427                .as_bytes(),
5428        )
5429        .with_interval(GnoInterval::new(1, 5));
5430
5431        let mut buf = Vec::new();
5432        sid.serialize_tagged(&mut buf);
5433
5434        let mut parse_buf = ParseBuf(&buf);
5435        let decoded = Sid::deserialize_tagged(&mut parse_buf).unwrap();
5436
5437        assert!(
5438            parse_buf.is_empty(),
5439            "leftover bytes after deserialize_tagged"
5440        );
5441        assert_eq!(sid, decoded);
5442        assert!(decoded.tag().is_none());
5443    }
5444
5445    #[test]
5446    fn com_binlog_dump_gtid_tagged_format_byte() {
5447        let uuid1 = *uuid::Uuid::parse_str("3E11FA47-71CA-11E1-9E33-C80AA9429562")
5448            .unwrap()
5449            .as_bytes();
5450        let uuid2 = *uuid::Uuid::parse_str("A0B1C2D3-E4F5-6789-ABCD-EF0123456789")
5451            .unwrap()
5452            .as_bytes();
5453
5454        // Mix of tagged and untagged Sids
5455        let sid_tagged = Sid::new(uuid1)
5456            .with_tag(Tag::new("app".to_owned()).unwrap())
5457            .with_interval(GnoInterval::new(1, 100));
5458        let sid_untagged = Sid::new(uuid2).with_interval(GnoInterval::new(1, 50));
5459
5460        let cmd = ComBinlogDumpGtid::new(42)
5461            .with_pos(4)
5462            .with_sid(sid_tagged)
5463            .with_sid(sid_untagged);
5464
5465        let mut output = Vec::new();
5466        cmd.serialize(&mut output);
5467
5468        // Verify the format byte is 0x01 in the raw serialized output.
5469        // The n_sids header starts after: header(1) + flags(2) + server_id(4) +
5470        // filename_len(4) + filename(0) + pos(8) + sid_data_len(4) = 23
5471        let n_sids_offset = 1 + 2 + 4 + 4 + 8 + 4;
5472        let n_sids_bytes = &output[n_sids_offset..n_sids_offset + 8];
5473        let n_sids_raw = u64::from_le_bytes(n_sids_bytes.try_into().unwrap());
5474        let format_byte = (n_sids_raw >> 56) as u8;
5475        assert_eq!(format_byte, 0x01, "format byte should be 0x01 for tagged");
5476
5477        let decoded = ComBinlogDumpGtid::deserialize((), &mut ParseBuf(&output)).unwrap();
5478        assert_eq!(cmd, decoded);
5479
5480        // Verify the tags survived the roundtrip
5481        assert_eq!(decoded.sids()[0].tag().unwrap().as_str(), "app");
5482        assert!(decoded.sids()[1].tag().is_none());
5483    }
5484}