1use 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#[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 pub fn schema_ref(&self) -> &[u8] {
145 self.schema.as_bytes()
146 }
147
148 pub fn schema_str(&self) -> Cow<'_, str> {
150 String::from_utf8_lossy(self.schema_ref())
151 }
152
153 pub fn table_ref(&self) -> &[u8] {
155 self.table.as_bytes()
156 }
157
158 pub fn table_str(&self) -> Cow<'_, str> {
160 String::from_utf8_lossy(self.table_ref())
161 }
162
163 pub fn org_table_ref(&self) -> &[u8] {
167 self.org_table.as_bytes()
168 }
169
170 pub fn org_table_str(&self) -> Cow<'_, str> {
172 String::from_utf8_lossy(self.org_table_ref())
173 }
174
175 pub fn name_ref(&self) -> &[u8] {
177 self.name.as_bytes()
178 }
179
180 pub fn name_str(&self) -> Cow<'_, str> {
182 String::from_utf8_lossy(self.name_ref())
183 }
184
185 pub fn org_name_ref(&self) -> &[u8] {
189 self.org_name.as_bytes()
190 }
191
192 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#[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 }
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 pub fn column_length(&self) -> u32 {
339 *self.column_length
340 }
341
342 pub fn column_type(&self) -> ColumnType {
344 *self.column_type
345 }
346
347 pub fn character_set(&self) -> u16 {
349 *self.character_set
350 }
351
352 pub fn flags(&self) -> ColumnFlags {
354 *self.flags
355 }
356
357 pub fn decimals(&self) -> u8 {
365 *self.decimals
366 }
367
368 #[inline(always)]
370 pub fn schema_ref(&self) -> &[u8] {
371 self.meta.schema_ref()
372 }
373
374 #[inline(always)]
376 pub fn schema_str(&self) -> Cow<'_, str> {
377 self.meta.schema_str()
378 }
379
380 #[inline(always)]
382 pub fn table_ref(&self) -> &[u8] {
383 self.meta.table_ref()
384 }
385
386 #[inline(always)]
388 pub fn table_str(&self) -> Cow<'_, str> {
389 self.meta.table_str()
390 }
391
392 #[inline(always)]
396 pub fn org_table_ref(&self) -> &[u8] {
397 self.meta.org_table_ref()
398 }
399
400 #[inline(always)]
402 pub fn org_table_str(&self) -> Cow<'_, str> {
403 self.meta.org_table_str()
404 }
405
406 #[inline(always)]
408 pub fn name_ref(&self) -> &[u8] {
409 self.meta.name_ref()
410 }
411
412 #[inline(always)]
414 pub fn name_str(&self) -> Cow<'_, str> {
415 self.meta.name_str()
416 }
417
418 #[inline(always)]
422 pub fn org_name_ref(&self) -> &[u8] {
423 self.meta.org_name_ref()
424 }
425
426 #[inline(always)]
428 pub fn org_name_str(&self) -> Cow<'_, str> {
429 self.meta.org_name_str()
430 }
431}
432
433#[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 pub fn data_ref(&self) -> &[u8] {
455 self.data.as_bytes()
456 }
457
458 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#[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
494pub 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#[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 buf.parse::<RawInt<LenEnc>>(())?;
522 buf.parse::<RawInt<LenEnc>>(())?;
523
524 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 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#[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 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#[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#[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 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 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#[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 pub fn affected_rows(&self) -> u64 {
704 self.affected_rows
705 }
706
707 pub fn last_insert_id(&self) -> Option<u64> {
709 self.last_insert_id
710 }
711
712 pub fn status_flags(&self) -> StatusFlags {
714 self.status_flags
715 }
716
717 pub fn warnings(&self) -> u16 {
719 self.warnings
720 }
721
722 pub fn info_ref(&self) -> Option<&[u8]> {
724 self.info.as_ref().map(|x| x.as_bytes())
725 }
726
727 pub fn info_str(&self) -> Option<Cow<'_, str>> {
729 self.info.as_ref().map(|x| x.as_str())
730 }
731
732 pub fn session_state_info_ref(&self) -> Option<&[u8]> {
734 self.session_state_info.as_ref().map(|x| x.as_bytes())
735 }
736
737 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#[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 pub fn stage(&self) -> u8 {
816 *self.stage
817 }
818
819 pub fn max_stage(&self) -> u8 {
820 *self.max_stage
821 }
822
823 pub fn progress(&self) -> u32 {
825 *self.progress
826 }
827
828 pub fn stage_info_ref(&self) -> &[u8] {
830 self.stage_info.as_bytes()
831 }
832
833 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); 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#[derive(Debug, Clone, PartialEq)]
899pub enum ErrPacket<'a> {
900 Error(ServerError<'a>),
901 Progress(ProgressReport<'a>),
902}
903
904impl ErrPacket<'_> {
905 pub fn is_error(&self) -> bool {
907 matches!(self, ErrPacket::Error { .. })
908 }
909
910 pub fn is_progress_report(&self) -> bool {
912 !self.is_error()
913 }
914
915 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 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#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
986pub struct SqlState {
987 __state_marker: SqlStateMarker,
988 state: [u8; 5],
989}
990
991impl SqlState {
992 pub fn new(state: [u8; 5]) -> Self {
994 Self {
995 __state_marker: SqlStateMarker::new(),
996 state,
997 }
998 }
999
1000 pub fn as_bytes(&self) -> [u8; 5] {
1002 self.state
1003 }
1004
1005 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#[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 pub fn error_code(&self) -> u16 {
1051 *self.code
1052 }
1053
1054 pub fn sql_state_ref(&self) -> Option<&SqlState> {
1056 self.state.as_ref()
1057 }
1058
1059 pub fn message_ref(&self) -> &[u8] {
1061 self.message.as_bytes()
1062 }
1063
1064 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 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#[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 pub fn file_name_ref(&self) -> &[u8] {
1150 self.file_name.as_bytes()
1151 }
1152
1153 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 Old([u8; 8]),
1196 Native([u8; 20]),
1198 Sha2([u8; 32]),
1200 Clear(Cow<'a, [u8]>),
1202 Ed25519([u8; 64]),
1207 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#[derive(Debug, Clone, Eq, PartialEq, Hash)]
1262pub enum AuthPlugin<'a> {
1263 MysqlOldPassword,
1265 MysqlClearPassword,
1267 MysqlNativePassword,
1269 CachingSha2Password,
1271 Ed25519,
1276 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 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 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#[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#[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 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#[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#[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
1656pub 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 if *algo != b'P' || *factor > 20 {
1667 return None;
1668 }
1669
1670 Some((1024 << factor, rest.try_into().expect("infallible")))
1671}
1672
1673#[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 capabilities_1: Const<CapabilityFlags, LeU32LowerHalf>,
1683 default_collation: RawInt<u8>,
1684 status_flags: Const<StatusFlags, LeU16>,
1685 capabilities_2: Const<CapabilityFlags, LeU32UpperHalf>,
1687 auth_plugin_data_len: RawInt<u8>,
1688 __reserved: Skip<6>,
1689 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 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 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 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 if let Some(scramble_2) = &self.scramble_2 {
1784 scramble_2.serialize(&mut *buf);
1785 }
1786
1787 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 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 pub fn protocol_version(&self) -> u8 {
1867 self.protocol_version.0
1868 }
1869
1870 pub fn server_version_ref(&self) -> &[u8] {
1872 self.server_version.as_bytes()
1873 }
1874
1875 pub fn server_version_str(&self) -> Cow<'_, str> {
1878 self.server_version.as_str()
1879 }
1880
1881 pub fn server_version_parsed(&self) -> Option<(u16, u16, u16)> {
1885 VERSION_RE
1886 .captures(self.server_version_ref())
1887 .map(|captures| {
1888 (
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 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 (
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 pub fn connection_id(&self) -> u32 {
1913 self.connection_id.0
1914 }
1915
1916 pub fn scramble_1_ref(&self) -> &[u8] {
1918 self.scramble_1.as_ref()
1919 }
1920
1921 pub fn scramble_2_ref(&self) -> Option<&[u8]> {
1925 self.scramble_2.as_ref().map(|x| x.as_bytes())
1926 }
1927
1928 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 out.resize(20, 0);
1936 out
1937 }
1938
1939 pub fn capabilities(&self) -> CapabilityFlags {
1941 self.capabilities_1.0 | self.capabilities_2.0
1942 }
1943
1944 pub fn mariadb_ext_capabilities(&self) -> MariadbCapabilities {
1946 self.mariadb_ext_capabilities.0
1947 }
1948 pub fn default_collation(&self) -> u8 {
1950 self.default_collation.0
1951 }
1952
1953 pub fn status_flags(&self) -> StatusFlags {
1955 self.status_flags.0
1956 }
1957
1958 pub fn auth_plugin_name_ref(&self) -> Option<&[u8]> {
1960 self.auth_plugin_name.as_ref().map(|x| x.as_bytes())
1961 }
1962
1963 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 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 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
2124fn 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
2139fn 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 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 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 serialize_connect_attrs(&Default::default(), buf);
2201 }
2202 }
2203}
2204
2205type 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 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#[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 pub fn statement_id(&self) -> u32 {
2554 *self.statement_id
2555 }
2556
2557 pub fn num_columns(&self) -> u16 {
2559 *self.num_columns
2560 }
2561
2562 pub fn num_params(&self) -> u16 {
2564 *self.num_params
2565 }
2566
2567 pub fn warning_count(&self) -> u16 {
2569 *self.warning_count
2570 }
2571}
2572
2573#[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 pub fn new(num_columns: usize) -> Self {
2593 Self::from_bytes(vec![0; Self::bitmap_len(num_columns)])
2594 }
2595
2596 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 pub fn from_bytes(bytes: U) -> Self {
2625 Self(bytes, PhantomData)
2626 }
2627
2628 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 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 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#[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, }
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 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 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 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 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 ¶ms {
2987 match p.bin_len() as usize {
2989 0 => data_len += 1, x => data_len += x + 1, }
2992 }
2993
2994 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 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 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 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 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 left.flags().union(right.flags()),
3345 )
3346 } else {
3347 }
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#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3439pub struct ComRegisterSlave<'a> {
3440 header: ComRegisterSlaveHeader,
3441 server_id: RawInt<LeU32>,
3443 hostname: RawBytes<'a, U8Bytes>,
3446 user: RawBytes<'a, U8Bytes>,
3453 password: RawBytes<'a, U8Bytes>,
3460 port: RawInt<LeU16>,
3467 replication_rank: RawInt<LeU32>,
3469 master_id: RawInt<LeU32>,
3472}
3473
3474impl<'a> ComRegisterSlave<'a> {
3475 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 pub fn with_hostname(mut self, hostname: impl Into<Cow<'a, [u8]>>) -> Self {
3491 self.hostname = RawBytes::new(hostname);
3492 self
3493 }
3494
3495 pub fn with_user(mut self, user: impl Into<Cow<'a, [u8]>>) -> Self {
3497 self.user = RawBytes::new(user);
3498 self
3499 }
3500
3501 pub fn with_password(mut self, password: impl Into<Cow<'a, [u8]>>) -> Self {
3503 self.password = RawBytes::new(password);
3504 self
3505 }
3506
3507 pub fn with_port(mut self, port: u16) -> Self {
3509 self.port = RawInt::new(port);
3510 self
3511 }
3512
3513 pub fn with_replication_rank(mut self, replication_rank: u32) -> Self {
3515 self.replication_rank = RawInt::new(replication_rank);
3516 self
3517 }
3518
3519 pub fn with_master_id(mut self, master_id: u32) -> Self {
3521 self.master_id = RawInt::new(master_id);
3522 self
3523 }
3524
3525 pub fn server_id(&self) -> u32 {
3527 self.server_id.0
3528 }
3529
3530 pub fn hostname_raw(&self) -> &[u8] {
3532 self.hostname.as_bytes()
3533 }
3534
3535 pub fn hostname(&'a self) -> Cow<'a, str> {
3537 self.hostname.as_str()
3538 }
3539
3540 pub fn user_raw(&self) -> &[u8] {
3542 self.user.as_bytes()
3543 }
3544
3545 pub fn user(&'a self) -> Cow<'a, str> {
3547 self.user.as_str()
3548 }
3549
3550 pub fn password_raw(&self) -> &[u8] {
3552 self.password.as_bytes()
3553 }
3554
3555 pub fn password(&'a self) -> Cow<'a, str> {
3557 self.password.as_str()
3558 }
3559
3560 pub fn port(&self) -> u16 {
3562 self.port.0
3563 }
3564
3565 pub fn replication_rank(&self) -> u32 {
3567 self.replication_rank.0
3568 }
3569
3570 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#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3628pub struct ComTableDump<'a> {
3629 header: ComTableDumpHeader,
3630 database: RawBytes<'a, U8Bytes>,
3636 table: RawBytes<'a, U8Bytes>,
3642}
3643
3644impl<'a> ComTableDump<'a> {
3645 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 pub fn database_raw(&self) -> &[u8] {
3656 self.database.as_bytes()
3657 }
3658
3659 pub fn database(&self) -> Cow<'_, str> {
3661 self.database.as_str()
3662 }
3663
3664 pub fn table_raw(&self) -> &[u8] {
3666 self.table.as_bytes()
3667 }
3668
3669 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 #[derive(PartialEq, Eq, Hash, Debug, Clone, Copy)]
3704 pub struct BinlogDumpFlags: u16 {
3705 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#[derive(Clone, Debug, Eq, PartialEq, Hash)]
3720pub struct ComBinlogDump<'a> {
3721 header: ComBinlogDumpHeader,
3722 pos: RawInt<LeU32>,
3724 flags: Const<BinlogDumpFlags, LeU16>,
3728 server_id: RawInt<LeU32>,
3730 filename: RawBytes<'a, EofBytes>,
3735}
3736
3737impl<'a> ComBinlogDump<'a> {
3738 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 pub fn with_pos(mut self, pos: u32) -> Self {
3751 self.pos = RawInt::new(pos);
3752 self
3753 }
3754
3755 pub fn with_flags(mut self, flags: BinlogDumpFlags) -> Self {
3757 self.flags = Const::new(flags);
3758 self
3759 }
3760
3761 pub fn with_filename(mut self, filename: impl Into<Cow<'a, [u8]>>) -> Self {
3763 self.filename = RawBytes::new(filename);
3764 self
3765 }
3766
3767 pub fn pos(&self) -> u32 {
3769 *self.pos
3770 }
3771
3772 pub fn flags(&self) -> BinlogDumpFlags {
3774 *self.flags
3775 }
3776
3777 pub fn server_id(&self) -> u32 {
3779 *self.server_id
3780 }
3781
3782 pub fn filename_raw(&self) -> &[u8] {
3784 self.filename.as_bytes()
3785 }
3786
3787 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#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
3821pub struct GnoInterval {
3822 start: RawInt<LeU64>,
3823 end: RawInt<LeU64>,
3824}
3825
3826impl GnoInterval {
3827 pub fn new(start: u64, end: u64) -> Self {
3829 Self {
3830 start: RawInt::new(start),
3831 end: RawInt::new(end),
3832 }
3833 }
3834
3835 pub fn start(&self) -> u64 {
3837 self.start.0
3838 }
3839
3840 pub fn end(&self) -> u64 {
3842 self.end.0
3843 }
3844 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
3881pub const UUID_LEN: usize = 16;
3883
3884#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
3892pub struct Tag<'a>(Cow<'a, str>);
3893
3894impl<'a> Tag<'a> {
3895 pub const MAX_LEN: usize = 32;
3897
3898 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 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 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 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 pub fn as_str(&self) -> &str {
3940 &self.0
3941 }
3942
3943 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#[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#[derive(Debug, Clone, Eq, PartialEq, Hash)]
3999pub struct Sid<'a> {
4000 uuid: [u8; UUID_LEN],
4001 tag: Option<Tag<'a>>,
4008 intervals: Seq<'a, GnoInterval, LeU64>,
4009}
4010
4011impl Sid<'_> {
4012 pub fn new(uuid: [u8; UUID_LEN]) -> Self {
4014 Self {
4015 uuid,
4016 tag: None,
4017 intervals: Default::default(),
4018 }
4019 }
4020
4021 pub fn uuid(&self) -> [u8; UUID_LEN] {
4023 self.uuid
4024 }
4025
4026 pub fn tag(&self) -> Option<&Tag<'_>> {
4028 self.tag.as_ref()
4029 }
4030
4031 pub fn intervals(&self) -> &[GnoInterval] {
4033 &self.intervals[..]
4034 }
4035
4036 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 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); len += S(8); len += S((self.intervals.len() * 16) as u64);
4055 len.0
4056 }
4057
4058 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); len += S((self.intervals.len() * 16) as u64);
4069 len.0
4070 }
4071
4072 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 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 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 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 fn parse_tag_and_intervals<'a>(
4199 rest: &'a str,
4200 full: &str,
4201 ) -> Result<(Option<Tag<'static>>, &'a str), io::Error> {
4202 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 return Ok((None, rest));
4211 }
4212
4213 if let Some((potential_tag, intervals)) = rest.split_once(':') {
4215 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 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 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#[derive(Debug, Clone, Eq, PartialEq, Hash)]
4284pub struct ComBinlogDumpGtid<'a> {
4285 header: ComBinlogDumpGtidHeader,
4286 flags: Const<BinlogDumpFlags, LeU16>,
4288 server_id: RawInt<LeU32>,
4290 filename: RawBytes<'a, U32Bytes>,
4299 pos: RawInt<LeU64>,
4301 sid_block: Seq<'a, Sid<'a>, LeU64>,
4303}
4304
4305impl<'a> ComBinlogDumpGtid<'a> {
4306 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 pub fn server_id(&self) -> u32 {
4320 self.server_id.0
4321 }
4322
4323 pub fn flags(&self) -> BinlogDumpFlags {
4325 self.flags.0
4326 }
4327
4328 pub fn filename_raw(&self) -> &[u8] {
4330 self.filename.as_bytes()
4331 }
4332
4333 pub fn filename(&self) -> Cow<'_, str> {
4335 self.filename.as_str()
4336 }
4337
4338 pub fn pos(&self) -> u64 {
4340 self.pos.0
4341 }
4342
4343 pub fn sids(&self) -> &[Sid<'a>] {
4345 &self.sid_block
4346 }
4347
4348 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 pub fn with_server_id(mut self, server_id: u32) -> Self {
4362 self.server_id.0 = server_id;
4363 self
4364 }
4365
4366 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 pub fn with_pos(mut self, pos: u64) -> Self {
4379 self.pos.0 = pos;
4380 self
4381 }
4382
4383 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 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 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); 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 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 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 let raw_n_sids: RawInt<LeU64> = sid_buf.parse(())?;
4473 let format_byte = (raw_n_sids.0 >> 56) as u8;
4474 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 for _ in 0..n_sids {
4488 sids.push(Sid::deserialize_tagged(&mut sid_buf)?);
4489 }
4490 }
4491 _ => {
4492 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
4516pub 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 pub fn with_position(mut self, position: u64) -> Self {
4535 self.position.0 = position;
4536 self
4537 }
4538
4539 pub fn with_filename(mut self, filename: impl Into<Cow<'a, [u8]>>) -> Self {
4541 self.filename = RawBytes::new(filename);
4542 self
4543 }
4544
4545 pub fn position(&self) -> u64 {
4547 self.position.0
4548 }
4549
4550 pub fn filename_raw(&self) -> &[u8] {
4552 self.filename.as_bytes()
4553 }
4554
4555 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 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 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 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 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 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 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, 0x00, 0x00, 0x00, 0x01, 0x2d, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5109 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x72, 0x6f, 0x6f, 0x74, 0x00, 0x00, 0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5113 0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, ]
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, 0x00, 0x00, 0x00, 0x01, 0x2d, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5138 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x72, 0x6f, 0x6f, 0x74, 0x00, 0x00, 0x6d, 0x79, 0x64, 0x62, 0x00, 0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5143 0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, ]
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, 0x00, 0x00, 0x00, 0x01, 0x2d, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5181 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x72, 0x6f, 0x6f, 0x74, 0x00, 0x00, 0x00, 0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5186 0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, ]
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, 0x00, 0x00, 0x00, 0x01, 0x2d, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
5249 0x00, 0x00, 0x00, 0x00, 0x00, 0x14, 0x00, 0x00, 0x00, 0x72, 0x6f, 0x6f, 0x74, 0x00, 0x00, 0x6d, 0x79, 0x73, 0x71, 0x6c, 0x5f, 0x6e, 0x61, 0x74, 0x69, 0x76, 0x65, 0x5f, 0x70,
5254 0x61, 0x73, 0x73, 0x77, 0x6f, 0x72, 0x64, 0x00, ]
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 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 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 let input = "3E11FA47-71CA-11E1-9E33-C80AA9429562:1tag:23";
5333 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 assert!(Tag::new("abcdefghijklmnopqrstuvwxyz123456").is_ok());
5371 }
5372
5373 #[test]
5374 fn invalid_tags() {
5375 assert_eq!(Tag::new(""), Err(InvalidTag::Empty));
5377
5378 assert_eq!(
5380 Tag::new("abcdefghijklmnopqrstuvwxyz1234567"),
5381 Err(InvalidTag::TooLong(33))
5382 );
5383
5384 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 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 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 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 assert_eq!(decoded.sids()[0].tag().unwrap().as_str(), "app");
5482 assert!(decoded.sids()[1].tag().is_none());
5483 }
5484}