use crate::cx::Cx;
use crate::database::transaction::trace_database_transaction;
use crate::io::{AsyncRead, AsyncWrite, ReadBuf};
use crate::net::TcpStream;
use crate::obligation::graded::{ObligationToken, TransactionKind};
use crate::security::SecretString;
use crate::types::{CancelReason, Outcome};
use std::collections::{BTreeMap, HashMap, VecDeque};
use std::fmt;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Poll;
#[derive(Debug)]
pub enum MySqlError {
Io(io::Error),
Protocol(String),
InvalidPacket(String),
AuthenticationFailed(String),
Server {
code: u16,
sql_state: String,
message: String,
},
Cancelled(CancelReason),
ConnectionClosed,
ColumnNotFound(String),
TypeConversion {
column: String,
expected: &'static str,
actual: String,
},
InvalidUrl(String),
InvalidParameter(String),
TlsRequired,
TransactionFinished,
UnsupportedAuthPlugin(String),
IsolationLevelMismatch {
requested: IsolationLevel,
observed: String,
},
}
impl MySqlError {
#[must_use]
pub fn server_code(&self) -> Option<u16> {
match self {
Self::Server { code, .. } => Some(*code),
_ => None,
}
}
#[must_use]
pub fn sql_state(&self) -> Option<&str> {
match self {
Self::Server { sql_state, .. } => Some(sql_state),
_ => None,
}
}
#[must_use]
pub fn error_code(&self) -> Option<String> {
self.server_code().map(|c| c.to_string())
}
#[must_use]
pub fn is_serialization_failure(&self) -> bool {
self.server_code() == Some(1213)
}
#[must_use]
pub fn is_deadlock(&self) -> bool {
matches!(self.server_code(), Some(1205 | 1213))
}
#[must_use]
pub fn is_unique_violation(&self) -> bool {
self.server_code() == Some(1062)
}
#[must_use]
pub fn is_constraint_violation(&self) -> bool {
matches!(self.server_code(), Some(1062 | 1451 | 1452))
}
#[must_use]
pub fn is_connection_error(&self) -> bool {
matches!(
self,
Self::Io(_) | Self::ConnectionClosed | Self::TlsRequired
) || matches!(self.server_code(), Some(2006 | 2013))
}
#[must_use]
pub fn debug_details(&self) -> String {
match self {
Self::Server {
code,
sql_state,
message,
} => format!("MySQL error [{}] ({}): {}", code, sql_state, message),
_ => self.to_string(), }
}
#[must_use]
pub fn is_transient(&self) -> bool {
if matches!(self, Self::Io(_) | Self::ConnectionClosed) {
return true;
}
matches!(self.server_code(), Some(1205 | 1213 | 2006 | 2013))
}
#[must_use]
pub fn is_retryable(&self) -> bool {
self.is_transient()
}
}
impl fmt::Display for MySqlError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(e) => write!(f, "MySQL I/O error: {e}"),
Self::Protocol(msg) => write!(f, "MySQL protocol error: {msg}"),
Self::InvalidPacket(msg) => write!(f, "Invalid MySQL packet: {msg}"),
Self::AuthenticationFailed(msg) => write!(f, "MySQL authentication failed: {msg}"),
Self::Server {
code,
sql_state: _,
message: _,
} => {
match *code {
1045 => write!(f, "Authentication failed"),
1046 => write!(f, "No database selected"),
1049 => write!(f, "Database does not exist"),
1050 => write!(f, "Table already exists"),
1051 => write!(f, "Table does not exist"),
1054 => write!(f, "Column not found"),
1062 => write!(f, "Duplicate entry"),
1064 => write!(f, "SQL syntax error"),
1146 => write!(f, "Table does not exist"),
1364 => write!(f, "Field missing default value"),
1452 => write!(f, "Foreign key constraint failed"),
_ => write!(f, "Database operation failed"),
}
}
Self::Cancelled(reason) => write!(f, "MySQL operation cancelled: {reason}"),
Self::ConnectionClosed => write!(f, "MySQL connection is closed"),
Self::ColumnNotFound(name) => write!(f, "Column not found: {name}"),
Self::TypeConversion {
column,
expected,
actual,
} => write!(
f,
"Type conversion error for column {column}: expected {expected}, got {actual}"
),
Self::InvalidUrl(msg) => write!(f, "Invalid MySQL URL: {msg}"),
Self::InvalidParameter(msg) => write!(f, "Invalid MySQL parameter: {msg}"),
Self::TlsRequired => write!(f, "TLS required but not available"),
Self::TransactionFinished => write!(f, "Transaction already finished"),
Self::UnsupportedAuthPlugin(plugin) => {
write!(f, "Unsupported authentication plugin: {plugin}")
}
Self::IsolationLevelMismatch {
requested,
observed,
} => write!(
f,
"MySQL isolation level mismatch: requested {requested}, server reported {observed:?} \
— silent downgrade detected, transaction rolled back (br-asupersync-dvgvcu)"
),
}
}
}
impl std::error::Error for MySqlError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Io(e) => Some(e),
_ => None,
}
}
}
impl From<io::Error> for MySqlError {
fn from(err: io::Error) -> Self {
Self::Io(err)
}
}
#[allow(dead_code)]
mod capability {
pub const CLIENT_LONG_PASSWORD: u32 = 1;
pub const CLIENT_FOUND_ROWS: u32 = 2;
pub const CLIENT_LONG_FLAG: u32 = 4;
pub const CLIENT_CONNECT_WITH_DB: u32 = 8;
pub const CLIENT_NO_SCHEMA: u32 = 16;
pub const CLIENT_COMPRESS: u32 = 32;
pub const CLIENT_ODBC: u32 = 64;
pub const CLIENT_LOCAL_FILES: u32 = 128;
pub const CLIENT_IGNORE_SPACE: u32 = 256;
pub const CLIENT_PROTOCOL_41: u32 = 512;
pub const CLIENT_INTERACTIVE: u32 = 1024;
pub const CLIENT_SSL: u32 = 2048;
pub const CLIENT_IGNORE_SIGPIPE: u32 = 4096;
pub const CLIENT_TRANSACTIONS: u32 = 8192;
pub const CLIENT_RESERVED: u32 = 16384;
pub const CLIENT_SECURE_CONNECTION: u32 = 32768;
pub const CLIENT_MULTI_STATEMENTS: u32 = 1 << 16;
pub const CLIENT_MULTI_RESULTS: u32 = 1 << 17;
pub const CLIENT_PS_MULTI_RESULTS: u32 = 1 << 18;
pub const CLIENT_PLUGIN_AUTH: u32 = 1 << 19;
pub const CLIENT_CONNECT_ATTRS: u32 = 1 << 20;
pub const CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA: u32 = 1 << 21;
pub const CLIENT_DEPRECATE_EOF: u32 = 1 << 24;
}
#[allow(dead_code)]
mod command {
pub const COM_QUIT: u8 = 0x01;
pub const COM_INIT_DB: u8 = 0x02;
pub const COM_QUERY: u8 = 0x03;
pub const COM_FIELD_LIST: u8 = 0x04;
pub const COM_PING: u8 = 0x0E;
pub const COM_STMT_PREPARE: u8 = 0x16;
pub const COM_STMT_EXECUTE: u8 = 0x17;
pub const COM_STMT_SEND_LONG_DATA: u8 = 0x18;
pub const COM_STMT_CLOSE: u8 = 0x19;
pub const COM_STMT_RESET: u8 = 0x1A;
}
const MAX_PACKET_SIZE: u32 = 16 * 1024 * 1024 - 1;
const DEFAULT_MAX_RESULT_ROWS: usize = 1_000_000;
pub const DEFAULT_MAX_PREPARED_STATEMENTS: usize = 256;
const MAX_COLUMN_COUNT: u64 = 16_384;
const MAX_REASSEMBLED_PACKET_SIZE: usize = 64 * 1024 * 1024;
const MYSQL_BINARY_CHARSET_ID: u16 = 63;
#[allow(dead_code, missing_docs)]
pub mod column_type {
pub const MYSQL_TYPE_DECIMAL: u8 = 0;
pub const MYSQL_TYPE_TINY: u8 = 1;
pub const MYSQL_TYPE_SHORT: u8 = 2;
pub const MYSQL_TYPE_LONG: u8 = 3;
pub const MYSQL_TYPE_FLOAT: u8 = 4;
pub const MYSQL_TYPE_DOUBLE: u8 = 5;
pub const MYSQL_TYPE_NULL: u8 = 6;
pub const MYSQL_TYPE_TIMESTAMP: u8 = 7;
pub const MYSQL_TYPE_LONGLONG: u8 = 8;
pub const MYSQL_TYPE_INT24: u8 = 9;
pub const MYSQL_TYPE_DATE: u8 = 10;
pub const MYSQL_TYPE_TIME: u8 = 11;
pub const MYSQL_TYPE_DATETIME: u8 = 12;
pub const MYSQL_TYPE_YEAR: u8 = 13;
pub const MYSQL_TYPE_VARCHAR: u8 = 15;
pub const MYSQL_TYPE_BIT: u8 = 16;
pub const MYSQL_TYPE_JSON: u8 = 245;
pub const MYSQL_TYPE_NEWDECIMAL: u8 = 246;
pub const MYSQL_TYPE_ENUM: u8 = 247;
pub const MYSQL_TYPE_SET: u8 = 248;
pub const MYSQL_TYPE_TINY_BLOB: u8 = 249;
pub const MYSQL_TYPE_MEDIUM_BLOB: u8 = 250;
pub const MYSQL_TYPE_LONG_BLOB: u8 = 251;
pub const MYSQL_TYPE_BLOB: u8 = 252;
pub const MYSQL_TYPE_VAR_STRING: u8 = 253;
pub const MYSQL_TYPE_STRING: u8 = 254;
pub const MYSQL_TYPE_GEOMETRY: u8 = 255;
}
#[derive(Debug, Clone)]
pub struct MySqlColumn {
pub catalog: String,
pub schema: String,
pub table: String,
pub org_table: String,
pub name: String,
pub org_name: String,
pub charset: u16,
pub length: u32,
pub column_type: u8,
pub flags: u16,
pub decimals: u8,
}
#[derive(Debug, Clone, PartialEq)]
pub enum MySqlValue {
Null,
Bool(bool),
Tiny(i8),
Short(i16),
Long(i32),
LongLong(i64),
Float(f32),
Double(f64),
Text(String),
Bytes(Vec<u8>),
}
impl MySqlValue {
#[must_use]
pub fn is_null(&self) -> bool {
matches!(self, Self::Null)
}
#[must_use]
pub fn as_bool(&self) -> Option<bool> {
match self {
Self::Bool(v) => Some(*v),
Self::Tiny(v) => Some(*v != 0),
_ => None,
}
}
#[must_use]
pub fn as_i32(&self) -> Option<i32> {
match self {
Self::Long(v) => Some(*v),
Self::LongLong(v) => i32::try_from(*v).ok(),
Self::Short(v) => Some(i32::from(*v)),
Self::Tiny(v) => Some(i32::from(*v)),
_ => None,
}
}
#[must_use]
pub fn as_i64(&self) -> Option<i64> {
match self {
Self::LongLong(v) => Some(*v),
Self::Long(v) => Some(i64::from(*v)),
Self::Short(v) => Some(i64::from(*v)),
Self::Tiny(v) => Some(i64::from(*v)),
_ => None,
}
}
#[must_use]
pub fn as_f64(&self) -> Option<f64> {
match self {
Self::Double(v) => Some(*v),
Self::Float(v) => Some(f64::from(*v)),
_ => None,
}
}
#[must_use]
pub fn as_str(&self) -> Option<&str> {
match self {
Self::Text(v) => Some(v),
_ => None,
}
}
#[must_use]
pub fn as_bytes(&self) -> Option<&[u8]> {
match self {
Self::Bytes(v) => Some(v),
_ => None,
}
}
}
impl fmt::Display for MySqlValue {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Null => write!(f, "NULL"),
Self::Bool(v) => write!(f, "{v}"),
Self::Tiny(v) => write!(f, "{v}"),
Self::Short(v) => write!(f, "{v}"),
Self::Long(v) => write!(f, "{v}"),
Self::LongLong(v) => write!(f, "{v}"),
Self::Float(v) => write!(f, "{v}"),
Self::Double(v) => write!(f, "{v}"),
Self::Text(v) => write!(f, "{v}"),
Self::Bytes(v) => write!(f, "<bytes {} len>", v.len()),
}
}
}
#[derive(Debug, Clone)]
pub struct MySqlRow {
columns: Arc<Vec<MySqlColumn>>,
column_indices: Arc<BTreeMap<String, usize>>,
values: Vec<MySqlValue>,
}
impl MySqlRow {
pub fn get(&self, column: &str) -> Result<&MySqlValue, MySqlError> {
let idx = self
.column_indices
.get(column)
.ok_or_else(|| MySqlError::ColumnNotFound(column.to_string()))?;
self.values
.get(*idx)
.ok_or_else(|| MySqlError::ColumnNotFound(column.to_string()))
}
pub fn get_idx(&self, idx: usize) -> Result<&MySqlValue, MySqlError> {
self.values
.get(idx)
.ok_or_else(|| MySqlError::ColumnNotFound(format!("index {idx}")))
}
pub fn get_i32(&self, column: &str) -> Result<i32, MySqlError> {
let val = self.get(column)?;
val.as_i32().ok_or_else(|| MySqlError::TypeConversion {
column: column.to_string(),
expected: "i32",
actual: format!("{val:?}"),
})
}
pub fn get_i64(&self, column: &str) -> Result<i64, MySqlError> {
let val = self.get(column)?;
val.as_i64().ok_or_else(|| MySqlError::TypeConversion {
column: column.to_string(),
expected: "i64",
actual: format!("{val:?}"),
})
}
pub fn get_str(&self, column: &str) -> Result<&str, MySqlError> {
let val = self.get(column)?;
val.as_str().ok_or_else(|| MySqlError::TypeConversion {
column: column.to_string(),
expected: "string",
actual: format!("{val:?}"),
})
}
pub fn get_bool(&self, column: &str) -> Result<bool, MySqlError> {
let val = self.get(column)?;
val.as_bool().ok_or_else(|| MySqlError::TypeConversion {
column: column.to_string(),
expected: "bool",
actual: format!("{val:?}"),
})
}
#[must_use]
pub fn len(&self) -> usize {
self.values.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.values.is_empty()
}
#[must_use]
pub fn columns(&self) -> &[MySqlColumn] {
&self.columns
}
}
#[must_use]
pub struct MySqlRowStream<'a> {
connection: &'a mut MySqlConnection,
columns: Option<Arc<Vec<MySqlColumn>>>,
column_indices: Option<Arc<BTreeMap<String, usize>>>,
finished: bool,
pending_row_count: u64,
deprecate_eof: bool,
}
impl MySqlRowStream<'_> {
pub async fn next(&mut self, cx: &Cx) -> Outcome<Option<MySqlRow>, MySqlError> {
if self.finished {
return Outcome::Ok(None);
}
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
loop {
let (data, seq) = match self.connection.read_packet().await {
Ok((d, s)) => (d, s),
Err(e) => return Outcome::Err(e),
};
self.connection.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
continue;
}
match data[0] {
0xFF => {
return Outcome::Err(MySqlConnection::parse_error(&data));
}
_ => {
if let (Some(cols), Some(indices)) = (&self.columns, &self.column_indices) {
match MySqlConnection::parse_data_row_or_terminator(
&data,
cols,
self.deprecate_eof,
) {
Ok(Some(values)) => {
self.pending_row_count += 1;
return Outcome::Ok(Some(MySqlRow {
columns: cols.clone(),
column_indices: indices.clone(),
values,
}));
}
Ok(None) => {
self.finished = true;
self.connection.inner.status_flags =
match MySqlConnection::parse_result_set_terminator_status_flags(
&data,
) {
Ok(flags) => flags,
Err(_) => self.connection.inner.status_flags, };
return Outcome::Ok(None);
}
Err(e) => return Outcome::Err(e),
}
} else {
return Outcome::Err(MySqlError::Protocol(
"Streaming query received row data without column metadata".to_string(),
));
}
}
}
}
}
pub fn row_count(&self) -> u64 {
self.pending_row_count
}
}
impl Drop for MySqlRowStream<'_> {
fn drop(&mut self) {
if !self.finished {
self.connection.inner.closed = true;
}
}
}
impl MySqlConnection {
pub async fn query_stream<'a>(
&'a mut self,
cx: &Cx,
sql: &str,
) -> Outcome<MySqlRowStream<'a>, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(self.inner.sequence);
buf.write_byte(command::COM_QUERY);
buf.write_bytes(sql.as_bytes());
let packet = buf.build_packet();
self.inner.closed = true;
match self.write_all(&packet.bytes).await {
Ok(()) => {}
Err(e) => return Outcome::Err(e),
}
self.inner.sequence = packet.next_sequence;
let (first_packet, seq) = match self.read_packet().await {
Ok(p) => p,
Err(e) => return Outcome::Err(e),
};
self.inner.sequence = seq.wrapping_add(1);
if first_packet.is_empty() {
return Outcome::Err(MySqlError::InvalidPacket("Empty response".to_string()));
}
let deprecate_eof = self.inner.capabilities & capability::CLIENT_DEPRECATE_EOF != 0;
match first_packet[0] {
0xFF => {
Outcome::Err(Self::parse_error(&first_packet))
}
0x00 => {
self.inner.closed = false;
Outcome::Ok(MySqlRowStream {
connection: self,
columns: None,
column_indices: None,
finished: true,
pending_row_count: 0,
deprecate_eof,
})
}
_ => {
let mut reader = PacketReader::new(&first_packet);
let column_count_raw = match reader.read_lenenc_int() {
Ok(count) => count,
Err(e) => return Outcome::Err(e),
};
if column_count_raw > MAX_COLUMN_COUNT {
return Outcome::Err(MySqlError::Protocol(format!(
"column count {column_count_raw} exceeds maximum {MAX_COLUMN_COUNT}"
)));
}
let column_count = column_count_raw as usize;
if column_count == 0 {
self.inner.closed = false;
return Outcome::Ok(MySqlRowStream {
connection: self,
columns: None,
column_indices: None,
finished: true,
pending_row_count: 0,
deprecate_eof,
});
}
let (columns, indices) = match self.read_result_set_columns(column_count).await {
Ok((cols, idx)) => (cols, idx),
Err(e) => return Outcome::Err(e),
};
self.inner.closed = false;
Outcome::Ok(MySqlRowStream {
connection: self,
columns: Some(columns),
column_indices: Some(indices),
finished: false,
pending_row_count: 0,
deprecate_eof,
})
}
}
}
}
struct PacketBuffer {
buf: Vec<u8>,
sequence: u8,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct EncodedPacket {
bytes: Vec<u8>,
next_sequence: u8,
}
impl PacketBuffer {
fn new() -> Self {
Self {
buf: Vec::with_capacity(256),
sequence: 0,
}
}
fn set_sequence(&mut self, seq: u8) {
self.sequence = seq;
}
fn write_byte(&mut self, b: u8) {
self.buf.push(b);
}
fn write_bytes(&mut self, data: &[u8]) {
self.buf.extend_from_slice(data);
}
fn write_u16_le(&mut self, v: u16) {
self.buf.extend_from_slice(&v.to_le_bytes());
}
fn write_u32_le(&mut self, v: u32) {
self.buf.extend_from_slice(&v.to_le_bytes());
}
fn write_null_terminated(&mut self, s: &str) {
self.buf.extend_from_slice(s.as_bytes());
self.buf.push(0);
}
fn write_lenenc_int(&mut self, v: u64) {
if v < 251 {
self.buf.push(v as u8);
} else if v < 65536 {
self.buf.push(0xFC);
self.buf.extend_from_slice(&(v as u16).to_le_bytes());
} else if v < 16_777_216 {
self.buf.push(0xFD);
self.buf.push((v & 0xFF) as u8);
self.buf.push(((v >> 8) & 0xFF) as u8);
self.buf.push(((v >> 16) & 0xFF) as u8);
} else {
self.buf.push(0xFE);
self.buf.extend_from_slice(&v.to_le_bytes());
}
}
fn build_packet(&self) -> EncodedPacket {
let mut sequence = self.sequence;
let mut offset = 0usize;
let max_payload = MAX_PACKET_SIZE as usize;
let payload_len = self.buf.len();
let packet_count = if payload_len == 0 {
1
} else {
(payload_len / max_payload).saturating_add(1)
};
let header_size = packet_count.saturating_mul(4);
let mut result = Vec::with_capacity(payload_len.saturating_add(header_size));
loop {
let remaining = payload_len.saturating_sub(offset);
let chunk_len = remaining.min(max_payload);
result.push((chunk_len & 0xFF) as u8);
result.push(((chunk_len >> 8) & 0xFF) as u8);
result.push(((chunk_len >> 16) & 0xFF) as u8);
result.push(sequence);
if chunk_len > 0 {
result.extend_from_slice(&self.buf[offset..offset + chunk_len]);
offset += chunk_len;
}
sequence = sequence.wrapping_add(1);
if chunk_len < max_payload {
break;
}
if offset == payload_len {
result.extend_from_slice(&[0, 0, 0, sequence]);
sequence = sequence.wrapping_add(1);
break;
}
}
EncodedPacket {
bytes: result,
next_sequence: sequence,
}
}
}
struct PacketReader<'a> {
data: &'a [u8],
pos: usize,
}
impl<'a> PacketReader<'a> {
fn new(data: &'a [u8]) -> Self {
Self { data, pos: 0 }
}
fn remaining(&self) -> usize {
self.data.len().saturating_sub(self.pos)
}
fn read_byte(&mut self) -> Result<u8, MySqlError> {
if self.pos >= self.data.len() {
return Err(MySqlError::Protocol("unexpected end of packet".to_string()));
}
let b = self.data[self.pos];
self.pos += 1;
Ok(b)
}
fn read_bytes(&mut self, len: usize) -> Result<&'a [u8], MySqlError> {
if len > self.data.len().saturating_sub(self.pos) {
return Err(MySqlError::Protocol("unexpected end of packet".to_string()));
}
let data = &self.data[self.pos..self.pos + len];
self.pos += len;
Ok(data)
}
fn read_rest(&mut self) -> &'a [u8] {
let data = &self.data[self.pos..];
self.pos = self.data.len();
data
}
fn read_u16_le(&mut self) -> Result<u16, MySqlError> {
let bytes = self.read_bytes(2)?;
Ok(u16::from_le_bytes([bytes[0], bytes[1]]))
}
fn read_u32_le(&mut self) -> Result<u32, MySqlError> {
let bytes = self.read_bytes(4)?;
Ok(u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
}
fn read_u64_le(&mut self) -> Result<u64, MySqlError> {
let bytes = self.read_bytes(8)?;
Ok(u64::from_le_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
]))
}
fn read_null_terminated(&mut self) -> Result<&'a str, MySqlError> {
let start = self.pos;
while self.pos < self.data.len() && self.data[self.pos] != 0 {
self.pos += 1;
}
if self.pos >= self.data.len() {
return Err(MySqlError::Protocol("unterminated string".to_string()));
}
let s = std::str::from_utf8(&self.data[start..self.pos])
.map_err(|e| MySqlError::Protocol(format!("invalid UTF-8: {e}")))?;
self.pos += 1; Ok(s)
}
fn read_lenenc_int(&mut self) -> Result<u64, MySqlError> {
let first = self.read_byte()?;
match first {
0..=250 => Ok(u64::from(first)),
0xFC => Ok(u64::from(self.read_u16_le()?)),
0xFD => {
let bytes = self.read_bytes(3)?;
Ok(u64::from(bytes[0]) | (u64::from(bytes[1]) << 8) | (u64::from(bytes[2]) << 16))
}
0xFE => self.read_u64_le(),
0xFB => Err(MySqlError::Protocol(
"NULL in length-encoded int".to_string(),
)),
_ => Err(MySqlError::Protocol(format!(
"invalid length-encoded int prefix: {first}"
))),
}
}
fn read_lenenc_str(&mut self) -> Result<&'a str, MySqlError> {
let len = usize::try_from(self.read_lenenc_int()?)
.map_err(|_| MySqlError::Protocol("length too large".to_string()))?;
let bytes = self.read_bytes(len)?;
std::str::from_utf8(bytes).map_err(|e| MySqlError::Protocol(format!("invalid UTF-8: {e}")))
}
fn read_lenenc_bytes(&mut self) -> Result<&'a [u8], MySqlError> {
let len = usize::try_from(self.read_lenenc_int()?)
.map_err(|_| MySqlError::Protocol("length too large".to_string()))?;
self.read_bytes(len)
}
}
fn sha1(data: &[u8]) -> [u8; 20] {
use sha1::Digest;
let mut hasher = sha1::Sha1::new();
hasher.update(data);
hasher.finalize().into()
}
fn sha256(data: &[u8]) -> [u8; 32] {
use sha2::Digest;
let mut hasher = sha2::Sha256::new();
hasher.update(data);
hasher.finalize().into()
}
const MIN_AUTH_NONCE_LEN: usize = 20;
const MIN_AUTH_NONCE_DISTINCT_BYTES: usize = 4;
fn validate_auth_nonce(plugin_name: &str, nonce: &[u8]) -> Result<(), MySqlError> {
if nonce.len() < MIN_AUTH_NONCE_LEN {
return Err(MySqlError::Protocol(format!(
"{plugin_name} server nonce too short: {} bytes; need at least {MIN_AUTH_NONCE_LEN}",
nonce.len()
)));
}
let mut seen = [false; 256];
let mut distinct = 0usize;
for &byte in nonce {
let slot = &mut seen[byte as usize];
if !*slot {
*slot = true;
distinct += 1;
}
}
if distinct < MIN_AUTH_NONCE_DISTINCT_BYTES {
return Err(MySqlError::Protocol(format!(
"{plugin_name} server nonce has insufficient entropy: {distinct} distinct byte values"
)));
}
Ok(())
}
struct ZeroizingBytes<T: AsMut<[u8]>>(T);
impl<T: AsMut<[u8]>> ZeroizingBytes<T> {
#[inline]
fn new(inner: T) -> Self {
Self(inner)
}
#[inline]
fn as_slice(&self) -> &[u8]
where
T: AsRef<[u8]>,
{
self.0.as_ref()
}
}
impl<T: AsMut<[u8]>> Drop for ZeroizingBytes<T> {
#[allow(unsafe_code)] fn drop(&mut self) {
let slice: &mut [u8] = self.0.as_mut();
for byte in slice.iter_mut() {
unsafe {
core::ptr::write_volatile(byte, 0);
}
}
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
}
}
fn mysql_native_auth(password: &str, nonce: &[u8]) -> Result<Vec<u8>, MySqlError> {
validate_auth_nonce("mysql_native_password", nonce)?;
if password.is_empty() {
return Ok(Vec::new());
}
let password_hash = ZeroizingBytes::new(sha1(password.as_bytes()));
let double_hash = ZeroizingBytes::new(sha1(password_hash.as_slice()));
let mut combined_bytes = Vec::with_capacity(nonce.len().saturating_add(20));
combined_bytes.extend_from_slice(nonce);
combined_bytes.extend_from_slice(double_hash.as_slice());
let combined = ZeroizingBytes::new(combined_bytes);
let scramble_hash = ZeroizingBytes::new(sha1(combined.as_slice()));
Ok(password_hash
.as_slice()
.iter()
.zip(scramble_hash.as_slice().iter())
.map(|(a, b)| a ^ b)
.collect())
}
fn caching_sha2_auth(password: &str, nonce: &[u8]) -> Result<Vec<u8>, MySqlError> {
validate_auth_nonce("caching_sha2_password", nonce)?;
if password.is_empty() {
return Ok(Vec::new());
}
let password_hash = ZeroizingBytes::new(sha256(password.as_bytes()));
let double_hash = ZeroizingBytes::new(sha256(password_hash.as_slice()));
let mut combined_bytes = Vec::with_capacity(32 + nonce.len());
combined_bytes.extend_from_slice(double_hash.as_slice());
combined_bytes.extend_from_slice(nonce);
let combined = ZeroizingBytes::new(combined_bytes);
let scramble_hash = ZeroizingBytes::new(sha256(combined.as_slice()));
Ok(password_hash
.as_slice()
.iter()
.zip(scramble_hash.as_slice().iter())
.map(|(a, b)| a ^ b)
.collect())
}
#[derive(Clone)]
pub struct MySqlConnectOptions {
pub host: String,
pub port: u16,
pub database: Option<String>,
pub user: String,
pub password: Option<SecretString>,
pub connect_timeout: Option<std::time::Duration>,
pub ssl_mode: SslMode,
pub insecure_legacy_mysql_native_password: bool,
pub insecure_allow_auth_switch_downgrade: bool,
pub requested_charset: Option<String>,
}
impl std::fmt::Debug for MySqlConnectOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MySqlConnectOptions")
.field("host", &self.host)
.field("port", &self.port)
.field("database", &self.database)
.field("user", &self.user)
.field("password", &self.password.as_ref().map(|_| "[REDACTED]"))
.field("connect_timeout", &self.connect_timeout)
.field("ssl_mode", &self.ssl_mode)
.field(
"insecure_legacy_mysql_native_password",
&self.insecure_legacy_mysql_native_password,
)
.field(
"insecure_allow_auth_switch_downgrade",
&self.insecure_allow_auth_switch_downgrade,
)
.field("requested_charset", &self.requested_charset)
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SslMode {
#[default]
Disabled,
Preferred,
Required,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IsolationLevel {
ReadUncommitted,
ReadCommitted,
RepeatableRead,
Serializable,
}
impl IsolationLevel {
#[must_use]
pub const fn as_sql(self) -> &'static str {
match self {
Self::ReadUncommitted => "READ UNCOMMITTED",
Self::ReadCommitted => "READ COMMITTED",
Self::RepeatableRead => "REPEATABLE READ",
Self::Serializable => "SERIALIZABLE",
}
}
#[must_use]
pub fn from_server_string(value: &str) -> Option<Self> {
let normalised: String = value
.trim()
.chars()
.map(|c| {
if c == '-' || c == '_' {
' '
} else {
c.to_ascii_uppercase()
}
})
.collect();
match normalised.as_str() {
"READ UNCOMMITTED" => Some(Self::ReadUncommitted),
"READ COMMITTED" => Some(Self::ReadCommitted),
"REPEATABLE READ" => Some(Self::RepeatableRead),
"SERIALIZABLE" => Some(Self::Serializable),
_ => None,
}
}
}
impl std::fmt::Display for IsolationLevel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_sql())
}
}
fn percent_decode(input: &str) -> String {
let mut out = Vec::with_capacity(input.len());
let bytes = input.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' && i + 2 < bytes.len() {
if let (Some(hi), Some(lo)) = (hex_nibble(bytes[i + 1]), hex_nibble(bytes[i + 2])) {
out.push((hi << 4) | lo);
i += 3;
continue;
}
}
out.push(bytes[i]);
i += 1;
}
String::from_utf8(out).unwrap_or_else(|e| String::from_utf8_lossy(e.as_bytes()).into_owned())
}
fn hex_nibble(b: u8) -> Option<u8> {
match b {
b'0'..=b'9' => Some(b - b'0'),
b'a'..=b'f' => Some(b - b'a' + 10),
b'A'..=b'F' => Some(b - b'A' + 10),
_ => None,
}
}
impl MySqlConnectOptions {
pub fn parse(url: &str) -> Result<Self, MySqlError> {
let url = url
.strip_prefix("mysql://")
.ok_or_else(|| MySqlError::InvalidUrl("URL must start with mysql://".to_string()))?;
let (auth_host, params) = url.split_once('?').unwrap_or((url, ""));
let (auth_host, database) = auth_host
.rsplit_once('/')
.map(|(ah, db)| (ah, Some(percent_decode(db))))
.unwrap_or((auth_host, None));
let (user, password, host_port) = if let Some((auth, host)) = auth_host.rsplit_once('@') {
let (user, password) = auth
.split_once(':')
.map_or((auth, None), |(u, p)| (u, Some(p)));
(percent_decode(user), password.map(percent_decode), host)
} else {
("root".to_string(), None, auth_host)
};
let (host, port) = if host_port.starts_with('[') {
if let Some((bracketed, rest)) = host_port.split_once(']') {
let addr = &bracketed[1..]; let port = if rest.is_empty() {
3306
} else if let Some(port_str) = rest.strip_prefix(':') {
port_str
.parse()
.map_err(|_| MySqlError::InvalidUrl(format!("invalid port: {port_str}")))?
} else {
return Err(MySqlError::InvalidUrl(format!(
"invalid host/port segment: {host_port}"
)));
};
(addr, port)
} else {
return Err(MySqlError::InvalidUrl(
"unclosed IPv6 bracket in host".to_string(),
));
}
} else if host_port.matches(':').count() > 1 {
(host_port, 3306)
} else {
match host_port.rsplit_once(':') {
Some((h, p)) => (
h,
p.parse()
.map_err(|_| MySqlError::InvalidUrl(format!("invalid port: {p}")))?,
),
None => (host_port, 3306),
}
};
if host.is_empty() {
return Err(MySqlError::InvalidUrl("missing host".to_string()));
}
let mut connect_timeout = None;
let mut ssl_mode = SslMode::Disabled;
let mut requested_charset = None;
if !params.is_empty() {
for pair in params.split('&') {
let (raw_key, raw_value) = pair.split_once('=').unwrap_or((pair, ""));
let key = percent_decode(raw_key);
let value = percent_decode(raw_value);
match key.as_str() {
"ssl-mode" | "sslmode" => {
if value.eq_ignore_ascii_case("disabled") {
ssl_mode = SslMode::Disabled;
} else if value.eq_ignore_ascii_case("preferred") {
ssl_mode = SslMode::Preferred;
} else if value.eq_ignore_ascii_case("required") {
ssl_mode = SslMode::Required;
} else {
return Err(MySqlError::InvalidUrl(format!(
"unknown ssl-mode: {value}"
)));
}
}
"connect_timeout" => {
let secs = value.parse::<u64>().map_err(|_| {
MySqlError::InvalidUrl(format!("invalid connect_timeout: {value}"))
})?;
connect_timeout = Some(std::time::Duration::from_secs(secs));
}
"charset" => {
requested_charset = Some(value);
}
_ => {
}
}
}
}
Ok(Self {
host: host.to_string(),
port,
database,
user,
password: password.map(SecretString::from_string),
connect_timeout,
ssl_mode,
insecure_legacy_mysql_native_password: false,
insecure_allow_auth_switch_downgrade: false,
requested_charset,
})
}
}
#[derive(Debug)]
struct Handshake {
server_version: String,
connection_id: u32,
auth_plugin_data: Vec<u8>,
capabilities: u32,
charset: u8,
status_flags: u16,
auth_plugin_name: String,
}
struct OkPacket {
affected_rows: u64,
status_flags: u16,
}
struct MySqlConnectionInner {
stream: TcpStream,
connection_id: u32,
capabilities: u32,
charset: u8,
status_flags: u16,
sequence: u8,
closed: bool,
server_version: String,
needs_rollback: bool,
session_isolation_restore: Option<IsolationLevel>,
max_result_rows: usize,
prepared_statement_epoch: u64,
prepared_cache: MySqlPreparedStatementCache,
query_in_flight: std::sync::atomic::AtomicBool,
statement_timeout_override: Option<std::time::Duration>,
applied_max_execution_time_ms: Option<u64>,
max_execution_time_unsupported: bool,
}
impl Drop for MySqlConnectionInner {
fn drop(&mut self) {
if !self.closed {
let _ = self.stream.shutdown(std::net::Shutdown::Both);
self.closed = true;
}
}
}
pub struct MySqlConnection {
inner: MySqlConnectionInner,
options: Option<MySqlConnectOptions>,
}
#[doc(hidden)]
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct FuzzHandshakeProtocol41 {
pub server_capabilities: u32,
pub client_capabilities: u32,
pub negotiated_capabilities: u32,
pub auth_plugin_name: String,
pub auth_plugin_data_len: usize,
}
impl fmt::Debug for MySqlConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MySqlConnection")
.field("connection_id", &self.inner.connection_id)
.field("server_version", &self.inner.server_version)
.field("closed", &self.inner.closed)
.field("kill_options_present", &self.options.is_some())
.finish()
}
}
impl Drop for MySqlConnection {
fn drop(&mut self) {
let in_flight = self
.inner
.query_in_flight
.load(std::sync::atomic::Ordering::Acquire);
let already_closed = self.inner.closed;
let kill_options = self.options.clone();
let thread_id = self.inner.connection_id;
if in_flight && !already_closed && thread_id != 0 {
if let Some(options) = kill_options {
std::thread::Builder::new()
.name(format!("asupersync-mysql-kill-{thread_id}"))
.spawn(move || {
let Ok(runtime) = crate::runtime::RuntimeBuilder::new()
.worker_threads(1)
.build()
else {
return;
};
let join = runtime.handle().spawn(async move {
let cx = match crate::cx::Cx::current() {
Some(cx) => cx,
None => return,
};
let killer =
match MySqlConnection::connect_with_options(&cx, options.clone())
.await
{
Outcome::Ok(c) => c,
_ => return,
};
let mut killer = killer;
let sql = format!("KILL QUERY {thread_id}");
let _ = killer.execute_unchecked_internal(&cx, &sql).await;
});
runtime.block_on(join);
})
.ok();
}
}
}
}
#[inline]
fn outcome_from_error<T>(err: MySqlError) -> Outcome<T, MySqlError> {
if let MySqlError::Cancelled(reason) = err {
Outcome::Cancelled(reason)
} else {
Outcome::Err(err)
}
}
fn ambient_cancel_reason() -> Option<CancelReason> {
Cx::with_current(|c| {
if c.checkpoint().is_err() {
Some(
c.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
)
} else {
None
}
})
.flatten()
}
fn stream_io_error(err: io::Error) -> MySqlError {
if err.kind() != io::ErrorKind::Interrupted {
return MySqlError::Io(err);
}
match ambient_cancel_reason() {
Some(reason) => MySqlError::Cancelled(reason),
None => MySqlError::Io(err),
}
}
fn eof_or_cancelled() -> MySqlError {
match ambient_cancel_reason() {
Some(reason) => MySqlError::Cancelled(reason),
None => MySqlError::Io(io::Error::new(
io::ErrorKind::UnexpectedEof,
"unexpected end of stream",
)),
}
}
const INJECTION_SUBSTRING_PATTERNS: &[&str] = &[
" or ",
" and ",
" union ",
" drop ",
" delete ",
" insert ",
" update ",
" alter ",
" create ",
" exec ",
" execute ",
" load ",
" into ",
" outfile ",
" dumpfile ",
"--",
"/*",
"*/",
";",
"'",
"\"",
];
const INJECTION_FUNCTION_PATTERNS: &[&str] = &["concat(", "char(", "ascii(", "substring("];
fn sql_injection_pattern(sql_lower: &str) -> Option<&'static str> {
if let Some(pattern) = INJECTION_SUBSTRING_PATTERNS
.iter()
.find(|pattern| sql_lower.contains(**pattern))
{
return Some(pattern);
}
INJECTION_FUNCTION_PATTERNS
.iter()
.find(|pattern| contains_sql_token(sql_lower, pattern))
.copied()
}
fn contains_sql_token(sql_lower: &str, pattern: &str) -> bool {
let bytes = sql_lower.as_bytes();
sql_lower.match_indices(pattern).any(|(idx, _)| {
idx == 0 || {
let before = bytes[idx - 1];
!(before.is_ascii_alphanumeric() || before == b'_')
}
})
}
async fn read_exact_from<R>(stream: &mut R, buf: &mut [u8]) -> Result<(), MySqlError>
where
R: AsyncRead + Unpin,
{
let mut pos = 0;
while pos < buf.len() {
let mut read_buf = ReadBuf::new(&mut buf[pos..]);
std::future::poll_fn(|task_cx| {
if Cx::with_current(|c| c.checkpoint().is_err()).unwrap_or(false) {
return Poll::Ready(Err(io::Error::new(io::ErrorKind::Interrupted, "cancelled")));
}
Pin::new(&mut *stream).poll_read(task_cx, &mut read_buf)
})
.await
.map_err(stream_io_error)?;
let n = read_buf.filled().len();
if n == 0 {
return Err(eof_or_cancelled());
}
pos += n;
}
Ok(())
}
impl MySqlConnection {
pub async fn connect(cx: &Cx, url: &str) -> Outcome<Self, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
let options = match MySqlConnectOptions::parse(url) {
Ok(opts) => opts,
Err(e) => return outcome_from_error(e),
};
Self::connect_with_options(cx, options).await
}
pub async fn connect_with_options(
cx: &Cx,
options: MySqlConnectOptions,
) -> Outcome<Self, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
let addr = format!("{}:{}", options.host, options.port);
let stream = if let Some(timeout) = options.connect_timeout {
match TcpStream::connect_timeout(addr, timeout).await {
Ok(s) => s,
Err(e) => return Outcome::Err(MySqlError::Io(e)),
}
} else {
match TcpStream::connect(addr).await {
Ok(s) => s,
Err(e) => return Outcome::Err(MySqlError::Io(e)),
}
};
let mut conn = Self {
inner: MySqlConnectionInner {
stream,
connection_id: 0,
capabilities: 0,
charset: 0,
status_flags: 0,
sequence: 0,
closed: false,
server_version: String::new(),
needs_rollback: false,
session_isolation_restore: None,
max_result_rows: DEFAULT_MAX_RESULT_ROWS,
prepared_statement_epoch: 0,
prepared_cache: MySqlPreparedStatementCache::new(DEFAULT_MAX_PREPARED_STATEMENTS),
query_in_flight: std::sync::atomic::AtomicBool::new(false),
statement_timeout_override: None,
applied_max_execution_time_ms: None,
max_execution_time_unsupported: false,
},
options: Some(options.clone()),
};
let handshake = match conn.read_handshake().await {
Ok(h) => h,
Err(e) => return outcome_from_error(e),
};
conn.inner.connection_id = handshake.connection_id;
conn.inner.charset = handshake.charset;
conn.inner.status_flags = handshake.status_flags;
conn.inner.server_version = handshake.server_version.clone();
if Self::should_fail_closed_without_tls(options.ssl_mode, handshake.capabilities) {
return Outcome::Err(MySqlError::TlsRequired);
}
if let Err(e) = conn.send_handshake_response(&options, &handshake).await {
return outcome_from_error(e);
}
if let Err(e) = conn.handle_auth_response(&options, &handshake).await {
return outcome_from_error(e);
}
Outcome::Ok(conn)
}
pub async fn cancel_in_flight_query(&self, cx: &Cx) -> Result<(), MySqlError> {
let options = self.options.clone().ok_or_else(|| {
MySqlError::Protocol(
"cancel_in_flight_query: connection has no stored MySqlConnectOptions \
(constructed outside of connect/connect_with_options — typically a \
test fixture); cannot reopen a fresh connection to issue KILL QUERY"
.to_string(),
)
})?;
let thread_id = self.connection_id();
let mut killer = match Self::connect_with_options(cx, options).await {
Outcome::Ok(c) => c,
Outcome::Err(e) => return Err(e),
Outcome::Cancelled(reason) => return Err(MySqlError::Cancelled(reason)),
Outcome::Panicked(_) => {
return Err(MySqlError::Protocol(
"cancel_in_flight_query: kill connection panicked during connect".to_string(),
));
}
};
let sql = format!("KILL QUERY {thread_id}");
match killer.execute_unchecked_internal(cx, &sql).await {
Outcome::Ok(_) => {
Ok(())
}
Outcome::Err(e) => Err(e),
Outcome::Cancelled(reason) => Err(MySqlError::Cancelled(reason)),
Outcome::Panicked(_) => Err(MySqlError::Protocol(
"cancel_in_flight_query: KILL QUERY panicked during execute".to_string(),
)),
}
}
pub fn set_statement_timeout_override(&mut self, timeout: Option<std::time::Duration>) {
self.inner.statement_timeout_override = timeout;
}
#[must_use]
pub fn statement_timeout_override(&self) -> Option<std::time::Duration> {
self.inner.statement_timeout_override
}
fn is_session_control_statement(sql: &str) -> bool {
let verb = sql
.trim_start()
.split(|c: char| c.is_whitespace() || c == ';')
.next()
.unwrap_or("");
[
"BEGIN",
"COMMIT",
"ROLLBACK",
"START",
"SAVEPOINT",
"RELEASE",
"SET",
"SHOW",
"USE",
"KILL",
]
.iter()
.any(|kw| verb.eq_ignore_ascii_case(kw))
}
async fn apply_statement_timeout(&mut self, cx: &Cx) -> Outcome<(), MySqlError> {
if self.inner.max_execution_time_unsupported {
return Outcome::Ok(());
}
let override_timeout = self.inner.statement_timeout_override;
let effective_ms = crate::database::wire_statement_timeout_ms(cx, override_timeout);
if effective_ms == self.inner.applied_max_execution_time_ms {
return Outcome::Ok(());
}
let remaining_ns = crate::database::remaining_budget(cx)
.map_or_else(|| "none".to_string(), |d| d.as_nanos().to_string());
let base_ms = override_timeout.map_or_else(
|| "none".to_string(),
|d| crate::database::statement_timeout_millis(d).to_string(),
);
let sql = match effective_ms {
Some(ms) => {
cx.trace(&format!(
"client.budget_forwarded proto=mysql base_ms={base_ms} \
remaining_ns={remaining_ns} max_execution_time_ms={ms}"
));
format!("SET SESSION max_execution_time = {ms}")
}
None => {
cx.trace(&format!(
"client.budget_forwarded proto=mysql base_ms={base_ms} \
remaining_ns={remaining_ns} max_execution_time_ms=default"
));
"SET SESSION max_execution_time = DEFAULT".to_string()
}
};
self.inner.applied_max_execution_time_ms = None;
match self.execute_unchecked_inner_impl(cx, &sql).await {
Outcome::Ok(_) => {
self.inner.applied_max_execution_time_ms = effective_ms;
Outcome::Ok(())
}
Outcome::Err(MySqlError::Server {
code: 1193,
message,
..
}) => {
self.inner.max_execution_time_unsupported = true;
cx.trace(&format!(
"client.budget_forwarded proto=mysql outcome=unsupported \
err_code=1193 err={message}"
));
Outcome::Ok(())
}
Outcome::Err(err) => Outcome::Err(err),
Outcome::Cancelled(reason) => Outcome::Cancelled(reason),
Outcome::Panicked(payload) => Outcome::Panicked(payload),
}
}
async fn wire_cancel_in_drain(&self, cx: &Cx) {
const MASKED_WIRE_CANCEL_POLLS: u32 = 4096;
let thread_id = self.inner.connection_id;
if thread_id == 0 {
cx.trace(
"client.wire_cancel proto=mysql outcome=skipped reason=no_connection_id \
fallback=connection_close",
);
return;
}
let Some(options) = self.options.clone() else {
cx.trace(
"client.wire_cancel proto=mysql outcome=skipped reason=no_stored_options \
fallback=connection_close",
);
return;
};
let cap = std::time::Duration::from_millis(500);
let mut kill_options = options;
kill_options.connect_timeout = Some(
kill_options
.connect_timeout
.map_or(cap, |timeout| timeout.min(cap)),
);
let kill = async {
let mut killer = match Self::connect_with_options(cx, kill_options).await {
Outcome::Ok(conn) => conn,
Outcome::Err(err) => return Err(("connect", err.to_string())),
Outcome::Cancelled(_) => {
return Err(("connect", "masked poll budget exhausted".to_string()));
}
Outcome::Panicked(_) => return Err(("connect", "panicked".to_string())),
};
let sql = format!("KILL QUERY {thread_id}");
match Box::pin(killer.execute_unchecked_internal(cx, &sql)).await {
Outcome::Ok(_) => Ok(()),
Outcome::Err(err) => Err(("kill_query", err.to_string())),
Outcome::Cancelled(_) => {
Err(("kill_query", "masked poll budget exhausted".to_string()))
}
Outcome::Panicked(_) => Err(("kill_query", "panicked".to_string())),
}
};
match crate::combinator::commit_section(cx, MASKED_WIRE_CANCEL_POLLS, kill).await {
Ok(()) => cx.trace(&format!(
"client.wire_cancel proto=mysql outcome=sent thread_id={thread_id}"
)),
Err((stage, err)) => cx.trace(&format!(
"client.wire_cancel proto=mysql outcome=send_failed stage={stage} \
fallback=connection_close err={err}"
)),
}
}
async fn read_handshake(&mut self) -> Result<Handshake, MySqlError> {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
const MIN_HANDSHAKE_SIZE: usize = 35;
if data.len() < MIN_HANDSHAKE_SIZE {
return Err(MySqlError::InvalidPacket(format!(
"handshake packet too short: {} bytes, minimum required: {}",
data.len(),
MIN_HANDSHAKE_SIZE
)));
}
let mut reader = PacketReader::new(&data);
let protocol_version = reader.read_byte()?;
if protocol_version != 10 {
return Err(MySqlError::Protocol(format!(
"unsupported protocol version: {protocol_version}"
)));
}
let server_version = reader.read_null_terminated()?.to_string();
let connection_id = reader.read_u32_le()?;
let auth_data_1 = reader.read_bytes(8)?;
let _ = reader.read_byte()?;
let cap_lower = reader.read_u16_le()?;
let charset = reader.read_byte()?;
let status_flags = reader.read_u16_le()?;
let cap_upper = reader.read_u16_le()?;
let capabilities = u32::from(cap_lower) | (u32::from(cap_upper) << 16);
let missing_required_caps =
(capability::CLIENT_PROTOCOL_41 | capability::CLIENT_SECURE_CONNECTION) & !capabilities;
if missing_required_caps != 0 {
let mut missing = Vec::new();
if missing_required_caps & capability::CLIENT_PROTOCOL_41 != 0 {
missing.push("CLIENT_PROTOCOL_41");
}
if missing_required_caps & capability::CLIENT_SECURE_CONNECTION != 0 {
missing.push("CLIENT_SECURE_CONNECTION");
}
return Err(MySqlError::Protocol(format!(
"server handshake missing required capabilities: {}",
missing.join(", ")
)));
}
let auth_data_len = reader.read_byte()?;
let _ = reader.read_bytes(10)?;
let mut auth_plugin_data = auth_data_1.to_vec();
if capabilities & capability::CLIENT_SECURE_CONNECTION != 0 {
let part2_len = std::cmp::max(13, auth_data_len.saturating_sub(8)) as usize;
let auth_data_2 = reader.read_bytes(part2_len.min(reader.remaining()))?;
let end = if auth_data_2.last() == Some(&0) {
auth_data_2.len() - 1
} else {
auth_data_2.len()
};
auth_plugin_data.extend_from_slice(&auth_data_2[..end]);
}
let auth_plugin_name =
if capabilities & capability::CLIENT_PLUGIN_AUTH != 0 && reader.remaining() > 0 {
reader.read_null_terminated()?.to_string()
} else {
"mysql_native_password".to_string()
};
Ok(Handshake {
server_version,
connection_id,
auth_plugin_data,
capabilities,
charset,
status_flags,
auth_plugin_name,
})
}
async fn send_handshake_response(
&mut self,
options: &MySqlConnectOptions,
handshake: &Handshake,
) -> Result<(), MySqlError> {
let mut buf = PacketBuffer::new();
buf.set_sequence(self.inner.sequence);
let client_caps = Self::client_handshake_response_capabilities(options.database.is_some());
self.inner.capabilities =
Self::negotiated_capabilities(handshake.capabilities, client_caps);
if let Some(requested) = &options.requested_charset {
Self::validate_charset_compatibility(requested, handshake.charset)?;
}
buf.write_u32_le(client_caps);
buf.write_u32_le(16_777_215); buf.write_byte(handshake.charset); buf.write_bytes(&[0u8; 23]);
buf.write_null_terminated(&options.user);
let password = options
.password
.as_ref()
.map(SecretString::as_str)
.unwrap_or_default();
let auth_response = match handshake.auth_plugin_name.as_str() {
"mysql_native_password" => {
if !options.insecure_legacy_mysql_native_password {
return Err(MySqlError::UnsupportedAuthPlugin(
"mysql_native_password is permanently disabled due to SHA1 cryptographic \
weaknesses that enable offline password cracking from captured network \
exchanges. Use MySQL 5.7+ with caching_sha2_password (default in MySQL 8.0+) \
or configure your MySQL server to require secure authentication plugins."
.to_string(),
));
}
mysql_native_auth(password, &handshake.auth_plugin_data)?
}
"caching_sha2_password" => caching_sha2_auth(password, &handshake.auth_plugin_data)?,
plugin => {
return Err(MySqlError::UnsupportedAuthPlugin(plugin.to_string()));
}
};
buf.write_lenenc_int(auth_response.len() as u64);
buf.write_bytes(&auth_response);
if let Some(ref db) = options.database {
buf.write_null_terminated(db);
}
buf.write_null_terminated(&handshake.auth_plugin_name);
let packet = buf.build_packet();
self.write_all(&packet.bytes).await?;
self.inner.sequence = packet.next_sequence;
Ok(())
}
#[inline]
const fn negotiated_capabilities(server_caps: u32, client_caps: u32) -> u32 {
server_caps & client_caps
}
#[inline]
const fn client_handshake_response_capabilities(connects_with_db: bool) -> u32 {
let mut client_caps = capability::CLIENT_PROTOCOL_41
| capability::CLIENT_SECURE_CONNECTION
| capability::CLIENT_PLUGIN_AUTH
| capability::CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA
| capability::CLIENT_TRANSACTIONS;
if connects_with_db {
client_caps |= capability::CLIENT_CONNECT_WITH_DB;
}
client_caps
}
fn validate_charset_compatibility(
requested: &str,
server_charset_id: u8,
) -> Result<(), MySqlError> {
let server_charset_name = match server_charset_id {
33 => "utf8", 45 => "utf8mb4", 8 => "latin1", _ => "unknown", };
let requested_lowercase = requested.to_lowercase();
let normalized_requested = match requested_lowercase.as_str() {
"utf8mb4" => "utf8mb4",
"utf8" => "utf8", "utf8mb3" => "utf8",
"latin1" => "latin1",
other => other,
};
match (normalized_requested, server_charset_name) {
("utf8mb4", "utf8mb4") | ("utf8", "utf8") | ("latin1", "latin1") => Ok(()),
("utf8mb4", "utf8") => Err(MySqlError::InvalidParameter(format!(
"charset incompatibility: client requested '{}' but server only supports '{}' \
(charset ID {}). utf8mb3 cannot store 4-byte UTF-8 sequences like emojis. \
Use server charset='utf8mb4' or remove charset parameter to accept server default.",
requested, server_charset_name, server_charset_id
))),
(req, srv) if req != srv => Err(MySqlError::InvalidParameter(format!(
"charset mismatch: client requested '{}' but server uses '{}' (charset ID {})",
requested, server_charset_name, server_charset_id
))),
_ => Ok(()),
}
}
#[inline]
const fn should_fail_closed_without_tls(ssl_mode: SslMode, _server_caps: u32) -> bool {
match ssl_mode {
SslMode::Disabled => false,
SslMode::Required => true,
SslMode::Preferred => true,
}
}
async fn handle_auth_response(
&mut self,
options: &MySqlConnectOptions,
handshake: &Handshake,
) -> Result<(), MySqlError> {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
return Err(MySqlError::Protocol("empty auth response".to_string()));
}
match data[0] {
0x00 => {
let ok = Self::parse_ok_packet(&data)?;
self.inner.status_flags = ok.status_flags;
Ok(())
}
0xFF => {
Err(Self::parse_error(&data))
}
0xFE => {
self.handle_auth_switch(&data[1..], options, handshake)
.await
}
0x01 => {
self.handle_caching_sha2_more_data(&data[1..], options, handshake)
.await
}
_ => Err(MySqlError::Protocol(format!(
"unexpected auth response: {:02x}",
data[0]
))),
}
}
async fn handle_auth_switch(
&mut self,
data: &[u8],
options: &MySqlConnectOptions,
handshake: &Handshake,
) -> Result<(), MySqlError> {
let mut reader = PacketReader::new(data);
let plugin_name = reader.read_null_terminated()?;
let auth_data_raw = reader.read_rest();
let auth_data = if auth_data_raw.last() == Some(&0) {
&auth_data_raw[..auth_data_raw.len() - 1]
} else {
auth_data_raw
};
let password = options
.password
.as_ref()
.map(SecretString::as_str)
.unwrap_or_default();
validate_auth_plugin_switch(handshake.auth_plugin_name.as_str(), plugin_name, options)?;
let auth_response = match plugin_name {
"mysql_native_password" => {
if !options.insecure_legacy_mysql_native_password {
return Err(MySqlError::UnsupportedAuthPlugin(
"mysql_native_password permanently blocked due to SHA1 cryptographic weakness. \
SHA1 enables offline password cracking from captured network exchanges. \
Use caching_sha2_password instead."
.to_string(),
));
}
mysql_native_auth(password, auth_data)?
}
"caching_sha2_password" => caching_sha2_auth(password, auth_data)?,
plugin => {
return Err(MySqlError::UnsupportedAuthPlugin(plugin.to_string()));
}
};
let mut buf = PacketBuffer::new();
buf.set_sequence(self.inner.sequence);
buf.write_bytes(&auth_response);
let packet = buf.build_packet();
self.write_all(&packet.bytes).await?;
self.inner.sequence = packet.next_sequence;
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
match data.first() {
Some(0x00) => {
let ok = Self::parse_ok_packet(&data)?;
self.inner.status_flags = ok.status_flags;
Ok(())
}
Some(0xFF) => Err(Self::parse_error(&data)),
Some(0x01) if plugin_name == "caching_sha2_password" => {
self.handle_caching_sha2_final(&data[1..], options).await
}
_ => Err(MySqlError::Protocol(
"unexpected auth switch response".to_string(),
)),
}
}
async fn handle_caching_sha2_more_data(
&mut self,
data: &[u8],
_options: &MySqlConnectOptions,
_handshake: &Handshake,
) -> Result<(), MySqlError> {
if data.first() == Some(&0x03) {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
match data.first() {
Some(0x00) => {
let ok = Self::parse_ok_packet(&data)?;
self.inner.status_flags = ok.status_flags;
Ok(())
}
Some(0xFF) => Err(Self::parse_error(&data)),
_ => Err(MySqlError::Protocol(
"unexpected response after fast auth".to_string(),
)),
}
} else if data.first() == Some(&0x04) {
Err(MySqlError::AuthenticationFailed(
"caching_sha2_password full auth requires secure connection".to_string(),
))
} else {
Err(MySqlError::Protocol(format!(
"unexpected caching_sha2 status: {:?}",
data.first()
)))
}
}
async fn handle_caching_sha2_final(
&mut self,
data: &[u8],
_options: &MySqlConnectOptions,
) -> Result<(), MySqlError> {
match data.first() {
Some(0x03) => {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
match data.first() {
Some(0x00) => {
let ok = Self::parse_ok_packet(&data)?;
self.inner.status_flags = ok.status_flags;
Ok(())
}
Some(0xFF) => Err(Self::parse_error(&data)),
_ => Err(MySqlError::Protocol(
"unexpected response after fast auth".to_string(),
)),
}
}
Some(0x04) => Err(MySqlError::AuthenticationFailed(
"caching_sha2_password full auth requires secure connection".to_string(),
)),
status => Err(MySqlError::Protocol(format!(
"unexpected caching_sha2 final status: {status:?}"
))),
}
}
fn validate_sql_security(&self, sql: &str) -> Result<(), MySqlError> {
let sql_lower = sql.to_lowercase();
const SAFE_STATIC_PATTERNS: &[&str] = &[
"start transaction",
"commit",
"rollback",
"select @@",
"show ",
"describe ",
"explain ",
"kill query ",
"set ",
];
if sql_lower.starts_with("kill query ") {
let id_part = sql_lower.trim_start_matches("kill query ").trim();
if id_part.chars().all(|c| c.is_ascii_digit()) {
return Ok(()); }
}
for pattern in SAFE_STATIC_PATTERNS {
if sql_lower.starts_with(pattern) {
return Ok(());
}
}
if let Some(pattern) = sql_injection_pattern(&sql_lower) {
return Err(MySqlError::InvalidParameter(format!(
"Potential SQL injection detected: query contains '{pattern}'. Use prepared statements for dynamic content.",
)));
}
if sql.chars().any(|c| matches!(c, '{' | '}' | '%')) {
return Err(MySqlError::InvalidParameter(
"Dynamic SQL pattern detected (contains format markers). Use prepared statements."
.to_string(),
));
}
Ok(())
}
pub async fn execute_static_sql(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
self.execute_unchecked_internal(cx, sql).await
}
pub async fn query_static_sql(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
self.query_unchecked_internal(cx, sql).await
}
pub async fn begin_transaction(
&mut self,
cx: &Cx,
) -> Outcome<MySqlTransaction<'_>, MySqlError> {
self.begin(cx).await
}
pub async fn kill_query(&mut self, cx: &Cx, thread_id: u32) -> Outcome<u64, MySqlError> {
let sql = format!("KILL QUERY {}", thread_id);
self.execute_unchecked_internal(cx, &sql).await
}
#[cfg(test)]
pub async fn query_unchecked_test_only(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
self.query_unchecked_inner_impl(cx, sql).await
}
#[deprecated(
note = "use query_static_sql for trusted-literal SQL or the prepared-statement APIs for parameterized queries (br-asupersync-0fxbp6)"
)]
pub async fn query(&mut self, cx: &Cx, sql: &str) -> Outcome<Vec<MySqlRow>, MySqlError> {
self.query_unchecked_internal(cx, sql).await
}
async fn query_unchecked_internal(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
if let Err(injection_error) = self.validate_sql_security(sql) {
return Outcome::Err(injection_error);
}
let was_closed_at_entry = self.inner.closed;
self.inner
.query_in_flight
.store(true, std::sync::atomic::Ordering::Release);
let apply_outcome = if Self::is_session_control_statement(sql) {
Outcome::Ok(())
} else {
self.apply_statement_timeout(cx).await
};
let result = match apply_outcome {
Outcome::Ok(()) => self.query_unchecked_inner_impl(cx, sql).await,
Outcome::Err(err) => Outcome::Err(err),
Outcome::Cancelled(reason) => Outcome::Cancelled(reason),
Outcome::Panicked(payload) => Outcome::Panicked(payload),
};
if matches!(result, Outcome::Cancelled(_))
&& self.inner.closed
&& !was_closed_at_entry
&& !Self::is_session_control_statement(sql)
{
self.wire_cancel_in_drain(cx).await;
}
self.inner
.query_in_flight
.store(false, std::sync::atomic::Ordering::Release);
result
}
async fn query_unchecked_inner_impl(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
if let Err(e) = self.drain_abandoned_transaction().await {
return outcome_from_error(e);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_QUERY);
buf.write_bytes(sql.as_bytes());
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (data, seq) = match self.read_packet().await {
Ok(p) => p,
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
return Outcome::Err(MySqlError::Protocol("empty query response".to_string()));
}
match data[0] {
0x00 => {
match Self::parse_ok_packet(&data) {
Ok(ok) => {
self.inner.status_flags = ok.status_flags;
self.inner.closed = false;
Outcome::Ok(Vec::new())
}
Err(e) => {
self.inner.closed = false;
Outcome::Err(e)
}
}
}
0xFF => {
let err = Self::parse_error(&data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
Outcome::Err(err)
}
0xFB => Outcome::Err(MySqlError::Protocol(
"LOAD DATA LOCAL INFILE request rejected: client local infile is disabled by default"
.to_string(),
)),
_ => {
match self.read_result_set(cx, &data).await {
Ok(rows) => {
self.inner.closed = false;
Outcome::Ok(rows)
}
Err(MySqlError::Cancelled(r)) => Outcome::Cancelled(r),
Err(e) => outcome_from_error(e),
}
}
}
}
async fn read_result_set(
&mut self,
cx: &Cx,
first_packet: &[u8],
) -> Result<Vec<MySqlRow>, MySqlError> {
let mut reader = PacketReader::new(first_packet);
let column_count_raw = reader.read_lenenc_int()?;
if column_count_raw > MAX_COLUMN_COUNT {
return Err(MySqlError::Protocol(format!(
"column count {column_count_raw} exceeds maximum {MAX_COLUMN_COUNT}"
)));
}
let column_count = column_count_raw as usize;
let deprecate_eof = self.inner.capabilities & capability::CLIENT_DEPRECATE_EOF != 0;
let max_rows = self.inner.max_result_rows;
if column_count == 0 {
return Ok(Vec::new());
}
let (columns, indices) = self.read_result_set_columns(column_count).await?;
let mut rows = Vec::new();
loop {
if cx.checkpoint().is_err() {
self.inner.closed = true;
return Err(MySqlError::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| crate::types::CancelReason::user("cancelled")),
));
}
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
continue;
}
match data[0] {
0xFF => {
return Err(Self::parse_error(&data));
}
_ => {
if let Some(values) =
Self::parse_data_row_or_terminator(&data, &columns, deprecate_eof)?
{
self.push_result_row(&mut rows, &columns, &indices, values, max_rows)?;
} else {
self.inner.status_flags =
Self::parse_result_set_terminator_status_flags(&data)?;
break;
}
}
}
}
Ok(rows)
}
async fn read_binary_result_set(
&mut self,
cx: &Cx,
first_packet: &[u8],
) -> Result<Vec<MySqlRow>, MySqlError> {
let mut reader = PacketReader::new(first_packet);
let column_count_raw = reader.read_lenenc_int()?;
if column_count_raw > MAX_COLUMN_COUNT {
return Err(MySqlError::Protocol(format!(
"column count {column_count_raw} exceeds maximum {MAX_COLUMN_COUNT}"
)));
}
let column_count = column_count_raw as usize;
let deprecate_eof = self.inner.capabilities & capability::CLIENT_DEPRECATE_EOF != 0;
let max_rows = self.inner.max_result_rows;
if column_count == 0 {
return Ok(Vec::new());
}
let (columns, indices) = self.read_result_set_columns(column_count).await?;
let mut rows = Vec::new();
loop {
if cx.checkpoint().is_err() {
self.inner.closed = true;
return Err(MySqlError::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| crate::types::CancelReason::user("cancelled")),
));
}
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
continue;
}
match data[0] {
0xFF => return Err(Self::parse_error(&data)),
_ => {
if let Some(values) =
Self::parse_binary_row_or_terminator(&data, &columns, deprecate_eof)?
{
self.push_result_row(&mut rows, &columns, &indices, values, max_rows)?;
} else {
self.inner.status_flags =
Self::parse_result_set_terminator_status_flags(&data)?;
break;
}
}
}
}
Ok(rows)
}
async fn read_result_set_columns(
&mut self,
column_count: usize,
) -> Result<(Arc<Vec<MySqlColumn>>, Arc<BTreeMap<String, usize>>), MySqlError> {
let mut columns = Vec::with_capacity(column_count);
let mut indices = BTreeMap::new();
for i in 0..column_count {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
let column = Self::parse_column_definition(&data)?;
indices.entry(column.name.clone()).or_insert(i);
columns.push(column);
}
if Self::expects_metadata_eof(self.inner.capabilities) {
self.read_metadata_eof("columns").await?;
}
Ok((Arc::new(columns), Arc::new(indices)))
}
async fn read_metadata_eof(&mut self, label: &'static str) -> Result<(), MySqlError> {
let (data, seq) = self.read_packet().await?;
self.inner.sequence = seq.wrapping_add(1);
if !Self::is_eof_packet(&data) {
return Err(MySqlError::Protocol(format!("expected EOF after {label}")));
}
self.inner.status_flags = Self::parse_eof_packet_status_flags(&data)?;
Ok(())
}
fn parse_column_definition(data: &[u8]) -> Result<MySqlColumn, MySqlError> {
let mut reader = PacketReader::new(data);
let catalog = reader.read_lenenc_str()?.to_string();
let schema = reader.read_lenenc_str()?.to_string();
let table = reader.read_lenenc_str()?.to_string();
let org_table = reader.read_lenenc_str()?.to_string();
let name = reader.read_lenenc_str()?.to_string();
let org_name = reader.read_lenenc_str()?.to_string();
let _ = reader.read_lenenc_int()?;
let charset = reader.read_u16_le()?;
let length = reader.read_u32_le()?;
let column_type = reader.read_byte()?;
let flags = reader.read_u16_le()?;
let decimals = reader.read_byte()?;
Ok(MySqlColumn {
catalog,
schema,
table,
org_table,
name,
org_name,
charset,
length,
column_type,
flags,
decimals,
})
}
fn push_result_row(
&mut self,
rows: &mut Vec<MySqlRow>,
columns: &Arc<Vec<MySqlColumn>>,
indices: &Arc<BTreeMap<String, usize>>,
values: Vec<MySqlValue>,
max_rows: usize,
) -> Result<(), MySqlError> {
if rows.len() >= max_rows {
self.inner.closed = true;
return Err(MySqlError::Protocol(format!(
"result set exceeds maximum row limit ({max_rows})"
)));
}
rows.push(MySqlRow {
columns: Arc::clone(columns),
column_indices: Arc::clone(indices),
values,
});
Ok(())
}
fn parse_text_row(data: &[u8], columns: &[MySqlColumn]) -> Result<Vec<MySqlValue>, MySqlError> {
let mut reader = PacketReader::new(data);
let mut values = Vec::with_capacity(columns.len());
for col in columns {
if reader.remaining() > 0 && data[reader.pos] == 0xFB {
reader.pos += 1;
values.push(MySqlValue::Null);
continue;
}
let raw = reader.read_lenenc_bytes()?;
let value = Self::parse_text_value(raw, col)?;
values.push(value);
}
if reader.remaining() != 0 {
return Err(MySqlError::Protocol(format!(
"row packet has {} trailing bytes",
reader.remaining()
)));
}
Ok(values)
}
fn parse_binary_row_or_terminator(
data: &[u8],
columns: &[MySqlColumn],
deprecate_eof: bool,
) -> Result<Option<Vec<MySqlValue>>, MySqlError> {
if Self::is_eof_packet(data) {
return Ok(None);
}
if data.first() == Some(&0x00) {
return Self::parse_binary_row(data, columns).map(Some);
}
if deprecate_eof && data.first() == Some(&0xFE) && Self::is_deprecate_eof_ok_packet(data) {
return Ok(None);
}
Err(MySqlError::Protocol(
"unexpected binary result-set row packet".to_string(),
))
}
fn parse_binary_row(
data: &[u8],
columns: &[MySqlColumn],
) -> Result<Vec<MySqlValue>, MySqlError> {
let mut reader = PacketReader::new(data);
let header = reader.read_byte()?;
if header != 0x00 {
return Err(MySqlError::Protocol(
"binary row must start with 0x00".to_string(),
));
}
let null_bitmap_len = columns.len().saturating_add(7).saturating_add(2) / 8;
let null_bitmap = reader.read_bytes(null_bitmap_len)?;
if (null_bitmap[0] & 0b0000_0011) != 0 {
return Err(MySqlError::Protocol(
"binary row reserved NULL-bitmap bits must be zero".to_string(),
));
}
let mut values = Vec::with_capacity(columns.len());
for (idx, col) in columns.iter().enumerate() {
let bit_idx = idx + 2;
if (null_bitmap[bit_idx / 8] & (1 << (bit_idx % 8))) != 0 {
values.push(MySqlValue::Null);
continue;
}
values.push(Self::parse_binary_value(&mut reader, col)?);
}
if reader.remaining() != 0 {
return Err(MySqlError::Protocol(format!(
"binary row packet has {} trailing bytes",
reader.remaining()
)));
}
Ok(values)
}
fn parse_binary_value(
reader: &mut PacketReader<'_>,
col: &MySqlColumn,
) -> Result<MySqlValue, MySqlError> {
Ok(match col.column_type {
column_type::MYSQL_TYPE_TINY => {
MySqlValue::Tiny(i8::from_le_bytes([reader.read_byte()?]))
}
column_type::MYSQL_TYPE_SHORT | column_type::MYSQL_TYPE_YEAR => {
MySqlValue::Short(i16::from_le_bytes(reader.read_u16_le()?.to_le_bytes()))
}
column_type::MYSQL_TYPE_LONG | column_type::MYSQL_TYPE_INT24 => {
MySqlValue::Long(i32::from_le_bytes(reader.read_u32_le()?.to_le_bytes()))
}
column_type::MYSQL_TYPE_LONGLONG => {
MySqlValue::LongLong(i64::from_le_bytes(reader.read_u64_le()?.to_le_bytes()))
}
column_type::MYSQL_TYPE_FLOAT => {
MySqlValue::Float(f32::from_bits(reader.read_u32_le()?))
}
column_type::MYSQL_TYPE_DOUBLE => {
MySqlValue::Double(f64::from_bits(reader.read_u64_le()?))
}
column_type::MYSQL_TYPE_DATE
| column_type::MYSQL_TYPE_DATETIME
| column_type::MYSQL_TYPE_TIMESTAMP => {
Self::parse_binary_datetime_value(reader, col.column_type)?
}
column_type::MYSQL_TYPE_TIME => Self::parse_binary_time_value(reader)?,
column_type::MYSQL_TYPE_NULL => MySqlValue::Null,
column_type::MYSQL_TYPE_VARCHAR
| column_type::MYSQL_TYPE_VAR_STRING
| column_type::MYSQL_TYPE_STRING
| column_type::MYSQL_TYPE_TINY_BLOB
| column_type::MYSQL_TYPE_MEDIUM_BLOB
| column_type::MYSQL_TYPE_LONG_BLOB
| column_type::MYSQL_TYPE_BLOB => Self::parse_binary_string_value(reader, col)?,
column_type::MYSQL_TYPE_GEOMETRY | column_type::MYSQL_TYPE_BIT => {
MySqlValue::Bytes(reader.read_lenenc_bytes()?.to_vec())
}
_ => {
let raw = reader.read_lenenc_bytes()?;
match std::str::from_utf8(raw) {
Ok(s) => MySqlValue::Text(s.to_string()),
Err(_) => MySqlValue::Bytes(raw.to_vec()),
}
}
})
}
fn parse_binary_string_value(
reader: &mut PacketReader<'_>,
col: &MySqlColumn,
) -> Result<MySqlValue, MySqlError> {
let raw = reader.read_lenenc_bytes()?;
Self::parse_string_or_bytes_value(raw, col)
}
fn parse_binary_datetime_value(
reader: &mut PacketReader<'_>,
column_type: u8,
) -> Result<MySqlValue, MySqlError> {
let len = usize::from(reader.read_byte()?);
let data = reader.read_bytes(len)?;
let mut value_reader = PacketReader::new(data);
if len == 0 {
return Ok(if column_type == column_type::MYSQL_TYPE_DATE {
MySqlValue::Text("0000-00-00".to_string())
} else {
MySqlValue::Text("0000-00-00 00:00:00".to_string())
});
}
if len != 4 && len != 7 && len != 11 {
return Err(MySqlError::Protocol(format!(
"invalid binary datetime length {len}"
)));
}
let year = value_reader.read_u16_le()?;
let month = value_reader.read_byte()?;
let day = value_reader.read_byte()?;
if column_type == column_type::MYSQL_TYPE_DATE || len == 4 {
return Ok(MySqlValue::Text(format!("{year:04}-{month:02}-{day:02}")));
}
let hour = value_reader.read_byte()?;
let minute = value_reader.read_byte()?;
let second = value_reader.read_byte()?;
if len == 7 {
return Ok(MySqlValue::Text(format!(
"{year:04}-{month:02}-{day:02} {hour:02}:{minute:02}:{second:02}"
)));
}
let micros = value_reader.read_u32_le()?;
Ok(MySqlValue::Text(format!(
"{year:04}-{month:02}-{day:02} {hour:02}:{minute:02}:{second:02}.{micros:06}"
)))
}
fn parse_binary_time_value(reader: &mut PacketReader<'_>) -> Result<MySqlValue, MySqlError> {
let len = usize::from(reader.read_byte()?);
let data = reader.read_bytes(len)?;
let mut value_reader = PacketReader::new(data);
if len == 0 {
return Ok(MySqlValue::Text("00:00:00".to_string()));
}
if len != 8 && len != 12 {
return Err(MySqlError::Protocol(format!(
"invalid binary time length {len}"
)));
}
let negative = value_reader.read_byte()? != 0;
let days = value_reader.read_u32_le()?;
let hour = value_reader.read_byte()?;
let minute = value_reader.read_byte()?;
let second = value_reader.read_byte()?;
let sign = if negative { "-" } else { "" };
if len == 8 {
return Ok(MySqlValue::Text(format!(
"{sign}{days} {hour:02}:{minute:02}:{second:02}"
)));
}
let micros = value_reader.read_u32_le()?;
Ok(MySqlValue::Text(format!(
"{sign}{days} {hour:02}:{minute:02}:{second:02}.{micros:06}"
)))
}
#[inline]
fn is_eof_packet(data: &[u8]) -> bool {
data.first() == Some(&0xFE) && data.len() < 9
}
fn parse_ok_packet(data: &[u8]) -> Result<OkPacket, MySqlError> {
if data.first() != Some(&0x00) {
return Err(MySqlError::Protocol("not an OK packet".to_string()));
}
let mut reader = PacketReader::new(&data[1..]);
let affected_rows = reader.read_lenenc_int()?;
let _last_insert_id = reader.read_lenenc_int()?;
let status_flags = reader.read_u16_le()?;
let _warning_count = reader.read_u16_le()?;
Ok(OkPacket {
affected_rows,
status_flags,
})
}
fn parse_eof_packet_status_flags(data: &[u8]) -> Result<u16, MySqlError> {
if !Self::is_eof_packet(data) {
return Err(MySqlError::Protocol("not an EOF packet".to_string()));
}
let mut reader = PacketReader::new(&data[1..]);
let _warning_count = reader.read_u16_le()?;
reader.read_u16_le()
}
fn parse_result_set_terminator_status_flags(data: &[u8]) -> Result<u16, MySqlError> {
if Self::is_eof_packet(data) {
return Self::parse_eof_packet_status_flags(data);
}
match data.first() {
Some(0x00 | 0xFE) => Self::parse_ok_packet_like_status_flags(data),
_ => Err(MySqlError::Protocol(
"not a result-set terminator packet".to_string(),
)),
}
}
fn parse_ok_packet_like_status_flags(data: &[u8]) -> Result<u16, MySqlError> {
match data.first() {
Some(0x00 | 0xFE) => {}
_ => return Err(MySqlError::Protocol("not an OK-like packet".to_string())),
}
let mut reader = PacketReader::new(&data[1..]);
let _affected_rows = reader.read_lenenc_int()?;
let _last_insert_id = reader.read_lenenc_int()?;
reader.read_u16_le()
}
#[inline]
const fn expects_metadata_eof(capabilities: u32) -> bool {
capabilities & capability::CLIENT_DEPRECATE_EOF == 0
}
#[inline]
fn is_result_set_ok_packet(data: &[u8]) -> bool {
if data.first() != Some(&0x00) {
return false;
}
let mut reader = PacketReader::new(&data[1..]);
reader.read_lenenc_int().is_ok()
&& reader.read_lenenc_int().is_ok()
&& reader.read_u16_le().is_ok()
&& reader.read_u16_le().is_ok()
}
#[inline]
fn is_deprecate_eof_ok_packet(data: &[u8]) -> bool {
if data.first() != Some(&0xFE) {
return false;
}
let mut reader = PacketReader::new(&data[1..]);
reader.read_lenenc_int().is_ok()
&& reader.read_lenenc_int().is_ok()
&& reader.read_u16_le().is_ok()
&& reader.read_u16_le().is_ok()
}
fn parse_data_row_or_terminator(
data: &[u8],
columns: &[MySqlColumn],
deprecate_eof: bool,
) -> Result<Option<Vec<MySqlValue>>, MySqlError> {
if Self::is_eof_packet(data) {
return Ok(None);
}
if deprecate_eof && matches!(data.first(), Some(&0x00 | &0xFE)) {
return match Self::parse_text_row(data, columns) {
Ok(values) => Ok(Some(values)),
Err(row_err) => {
if Self::is_result_set_ok_packet(data) || Self::is_deprecate_eof_ok_packet(data)
{
Ok(None)
} else {
Err(row_err)
}
}
};
}
Self::parse_text_row(data, columns).map(Some)
}
fn parse_text_value(data: &[u8], col: &MySqlColumn) -> Result<MySqlValue, MySqlError> {
let text = match Self::parse_string_or_bytes_value(data, col)? {
MySqlValue::Bytes(bytes) => return Ok(MySqlValue::Bytes(bytes)),
MySqlValue::Text(text) => text,
value => {
return Err(MySqlError::Protocol(format!(
"unexpected string parser value: {value:?}"
)));
}
};
let parse_err = |typ: &str| {
MySqlError::Protocol(format!("cannot parse {typ} from text value: {text:?}"))
};
Ok(match col.column_type {
column_type::MYSQL_TYPE_TINY => {
MySqlValue::Tiny(text.parse().map_err(|_| parse_err("TINY"))?)
}
column_type::MYSQL_TYPE_SHORT | column_type::MYSQL_TYPE_YEAR => {
MySqlValue::Short(text.parse().map_err(|_| parse_err("SHORT"))?)
}
column_type::MYSQL_TYPE_LONG | column_type::MYSQL_TYPE_INT24 => {
MySqlValue::Long(text.parse().map_err(|_| parse_err("LONG"))?)
}
column_type::MYSQL_TYPE_LONGLONG => {
MySqlValue::LongLong(text.parse().map_err(|_| parse_err("LONGLONG"))?)
}
column_type::MYSQL_TYPE_FLOAT => {
MySqlValue::Float(text.parse().map_err(|_| parse_err("FLOAT"))?)
}
column_type::MYSQL_TYPE_DOUBLE
| column_type::MYSQL_TYPE_DECIMAL
| column_type::MYSQL_TYPE_NEWDECIMAL => {
MySqlValue::Double(text.parse().map_err(|_| parse_err("DOUBLE"))?)
}
_ => MySqlValue::Text(text),
})
}
fn parse_string_or_bytes_value(
data: &[u8],
col: &MySqlColumn,
) -> Result<MySqlValue, MySqlError> {
if Self::is_binary_payload_column(col) {
return Ok(MySqlValue::Bytes(data.to_vec()));
}
let text = std::str::from_utf8(data)
.map_err(|e| MySqlError::Protocol(format!("invalid UTF-8: {e}")))?;
Ok(MySqlValue::Text(text.to_string()))
}
#[inline]
fn is_binary_payload_column(col: &MySqlColumn) -> bool {
matches!(
col.column_type,
column_type::MYSQL_TYPE_GEOMETRY | column_type::MYSQL_TYPE_BIT
) || (col.charset == MYSQL_BINARY_CHARSET_ID
&& Self::is_string_like_column_type(col.column_type))
}
#[inline]
const fn is_string_like_column_type(column_type: u8) -> bool {
matches!(
column_type,
column_type::MYSQL_TYPE_VARCHAR
| column_type::MYSQL_TYPE_VAR_STRING
| column_type::MYSQL_TYPE_STRING
| column_type::MYSQL_TYPE_TINY_BLOB
| column_type::MYSQL_TYPE_MEDIUM_BLOB
| column_type::MYSQL_TYPE_LONG_BLOB
| column_type::MYSQL_TYPE_BLOB
)
}
#[deprecated(
note = "use execute_static_sql for trusted-literal SQL or the prepared-statement APIs for parameterized commands (br-asupersync-0fxbp6)"
)]
pub async fn execute(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
self.execute_unchecked_internal(cx, sql).await
}
async fn execute_unchecked_internal(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
if let Err(injection_error) = self.validate_sql_security(sql) {
return Outcome::Err(injection_error);
}
let was_closed_at_entry = self.inner.closed;
self.inner
.query_in_flight
.store(true, std::sync::atomic::Ordering::Release);
let apply_outcome = if Self::is_session_control_statement(sql) {
Outcome::Ok(())
} else {
self.apply_statement_timeout(cx).await
};
let result = match apply_outcome {
Outcome::Ok(()) => self.execute_unchecked_inner_impl(cx, sql).await,
Outcome::Err(err) => Outcome::Err(err),
Outcome::Cancelled(reason) => Outcome::Cancelled(reason),
Outcome::Panicked(payload) => Outcome::Panicked(payload),
};
if matches!(result, Outcome::Cancelled(_))
&& self.inner.closed
&& !was_closed_at_entry
&& !Self::is_session_control_statement(sql)
{
self.wire_cancel_in_drain(cx).await;
}
self.inner
.query_in_flight
.store(false, std::sync::atomic::Ordering::Release);
result
}
async fn execute_unchecked_inner_impl(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<u64, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
if let Err(e) = self.drain_abandoned_transaction().await {
return outcome_from_error(e);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_QUERY);
buf.write_bytes(sql.as_bytes());
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (data, seq) = match self.read_packet().await {
Ok(p) => p,
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
if data.is_empty() {
return Outcome::Err(MySqlError::Protocol("empty execute response".to_string()));
}
match data[0] {
0x00 => {
match Self::parse_ok_packet(&data) {
Ok(ok) => {
self.inner.status_flags = ok.status_flags;
self.inner.closed = false;
Outcome::Ok(ok.affected_rows)
}
Err(e) => {
self.inner.closed = false;
Outcome::Err(e)
}
}
}
0xFF => {
let err = Self::parse_error(&data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
Outcome::Err(err)
}
0xFB => Outcome::Err(MySqlError::Protocol(
"LOAD DATA LOCAL INFILE request rejected: client local infile is disabled by default"
.to_string(),
)),
_ => {
match self.read_result_set(cx, &data).await {
Ok(_) => {
self.inner.closed = false;
Outcome::Ok(0)
}
Err(MySqlError::Cancelled(r)) => Outcome::Cancelled(r),
Err(e) => outcome_from_error(e),
}
}
}
}
pub async fn begin(&mut self, cx: &Cx) -> Outcome<MySqlTransaction<'_>, MySqlError> {
trace_database_transaction(cx, "mysql", "begin", "start");
match self
.execute_unchecked_internal(cx, "START TRANSACTION")
.await
{
Outcome::Ok(_) => {
trace_database_transaction(cx, "mysql", "begin", "ok");
Outcome::Ok(MySqlTransaction {
conn: self,
finished: false,
isolation_level: None,
read_only: false,
obligation: reserve_transaction_obligation(cx),
})
}
Outcome::Err(e) => {
trace_database_transaction(cx, "mysql", "begin", "err");
outcome_from_error(e)
}
Outcome::Cancelled(r) => {
trace_database_transaction(cx, "mysql", "begin", "cancelled");
Outcome::Cancelled(r)
}
Outcome::Panicked(p) => {
trace_database_transaction(cx, "mysql", "begin", "panicked");
Outcome::Panicked(p)
}
}
}
pub async fn begin_with_isolation(
&mut self,
cx: &Cx,
level: IsolationLevel,
read_only: bool,
) -> Outcome<MySqlTransaction<'_>, MySqlError> {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "start");
let previous_level = match self
.query_unchecked_internal(cx, "SELECT @@SESSION.transaction_isolation AS isolation")
.await
{
Outcome::Ok(rows) => rows
.first()
.and_then(|r| r.get_str("isolation").ok())
.and_then(IsolationLevel::from_server_string),
Outcome::Err(e) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
return outcome_from_error(e);
}
Outcome::Cancelled(r) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "cancelled");
return Outcome::Cancelled(r);
}
Outcome::Panicked(p) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "panicked");
return Outcome::Panicked(p);
}
};
let set_sql = format!("SET SESSION TRANSACTION ISOLATION LEVEL {level}");
match self.execute_unchecked_internal(cx, &set_sql).await {
Outcome::Ok(_) => {}
Outcome::Err(e) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
return outcome_from_error(e);
}
Outcome::Cancelled(r) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "cancelled");
return Outcome::Cancelled(r);
}
Outcome::Panicked(p) => {
trace_database_transaction(cx, "mysql", "begin_with_isolation", "panicked");
return Outcome::Panicked(p);
}
}
self.inner.session_isolation_restore = previous_level.filter(|prev| *prev != level);
let access_mode = if read_only { "READ ONLY" } else { "READ WRITE" };
let start_sql = format!("START TRANSACTION {access_mode}");
match self.execute_unchecked_internal(cx, &start_sql).await {
Outcome::Ok(_) => {}
Outcome::Err(e) => {
self.restore_session_isolation(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
return outcome_from_error(e);
}
Outcome::Cancelled(r) => {
self.restore_session_isolation(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "cancelled");
return Outcome::Cancelled(r);
}
Outcome::Panicked(p) => {
self.restore_session_isolation(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "panicked");
return Outcome::Panicked(p);
}
}
let observed_level = match self
.query_unchecked_internal(cx, "SELECT @@SESSION.transaction_isolation AS isolation")
.await
{
Outcome::Ok(rows) => match rows
.first()
.and_then(|r| r.get_str("isolation").ok())
.map(str::to_string)
{
Some(s) => s,
None => {
self.rollback_isolated_begin_or_mark(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
return Outcome::Err(MySqlError::IsolationLevelMismatch {
requested: level,
observed: String::new(),
});
}
},
Outcome::Err(e) => {
self.rollback_isolated_begin_or_mark(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
return outcome_from_error(e);
}
Outcome::Cancelled(r) => {
self.rollback_isolated_begin_or_mark(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "cancelled");
return Outcome::Cancelled(r);
}
Outcome::Panicked(p) => {
self.rollback_isolated_begin_or_mark(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "panicked");
return Outcome::Panicked(p);
}
};
match IsolationLevel::from_server_string(&observed_level) {
Some(parsed) if parsed == level => {
let obligation = reserve_transaction_obligation(cx);
trace_database_transaction(cx, "mysql", "begin_with_isolation", "ok");
Outcome::Ok(MySqlTransaction {
conn: self,
finished: false,
isolation_level: Some(level),
read_only,
obligation,
})
}
_ => {
self.rollback_isolated_begin_or_mark(cx).await;
trace_database_transaction(cx, "mysql", "begin_with_isolation", "err");
Outcome::Err(MySqlError::IsolationLevelMismatch {
requested: level,
observed: observed_level,
})
}
}
}
async fn restore_session_isolation(&mut self, cx: &Cx) {
let Some(previous) = self.inner.session_isolation_restore.take() else {
return;
};
let restore_sql = format!("SET SESSION TRANSACTION ISOLATION LEVEL {previous}");
match self.execute_unchecked_internal(cx, &restore_sql).await {
Outcome::Ok(_) => {}
Outcome::Err(err) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(&format!(
"restoring the session isolation level failed; marking connection for orphan cleanup: {err:?}"
));
}
Outcome::Cancelled(reason) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(&format!(
"restoring the session isolation level was cancelled; marking connection for orphan cleanup: {reason}"
));
}
Outcome::Panicked(_) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(
"restoring the session isolation level panicked; marking connection for orphan cleanup",
);
}
}
}
async fn rollback_isolated_begin_or_mark(&mut self, cx: &Cx) {
const MASKED_ROLLBACK_POLLS: u32 = 32;
match crate::combinator::commit_section(
cx,
MASKED_ROLLBACK_POLLS,
self.execute_unchecked_internal(cx, "ROLLBACK"),
)
.await
{
Outcome::Ok(_) => {
self.restore_session_isolation(cx).await;
}
Outcome::Err(err) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(&format!(
"begin_with_isolation cleanup rollback failed; marking connection for orphan cleanup: {:?}",
err
));
}
Outcome::Cancelled(reason) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(&format!(
"begin_with_isolation cleanup rollback was cancelled; marking connection for orphan cleanup: {reason}"
));
}
Outcome::Panicked(_) => {
self.mark_unusable_after_cleanup_failure();
cx.trace(
"begin_with_isolation cleanup rollback panicked; marking connection for orphan cleanup",
);
}
}
}
fn mark_unusable_after_cleanup_failure(&mut self) {
self.inner.needs_rollback = true;
self.inner.closed = true;
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
}
pub async fn ping(&mut self, cx: &Cx) -> Outcome<(), MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_PING);
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (data, seq) = match self.read_packet().await {
Ok(p) => p,
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
match data.first() {
Some(0x00) => match Self::parse_ok_packet(&data) {
Ok(ok) => {
self.inner.status_flags = ok.status_flags;
self.inner.closed = false;
Outcome::Ok(())
}
Err(e) => {
self.inner.closed = false;
Outcome::Err(e)
}
},
Some(0xFF) => {
let err = Self::parse_error(&data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
Outcome::Err(err)
}
_ => Outcome::Err(MySqlError::Protocol("unexpected ping response".to_string())),
}
}
#[must_use]
pub fn server_version(&self) -> &str {
&self.inner.server_version
}
#[must_use]
pub fn connection_id(&self) -> u32 {
self.inner.connection_id
}
#[must_use]
pub fn in_transaction(&self) -> bool {
self.inner.status_flags & 0x0001 != 0 }
fn invalidate_prepared_statements_for_pool_return(&mut self) {
self.inner.prepared_statement_epoch = self.inner.prepared_statement_epoch.wrapping_add(1);
}
#[must_use]
pub fn prepared_cache_stats(&self) -> MySqlPreparedCacheStats {
self.inner.prepared_cache.stats()
}
pub async fn close(&mut self) -> Result<(), MySqlError> {
if self.inner.closed {
return Ok(());
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_QUIT);
let packet = buf.build_packet();
let _ = self.write_all(&packet.bytes).await;
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
self.inner.closed = true;
Ok(())
}
pub async fn prepare(&mut self, cx: &Cx, sql: &str) -> Outcome<MySqlStatement, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
if let Err(e) = self.drain_abandoned_transaction().await {
return outcome_from_error(e);
}
if let Some(cached) = self.inner.prepared_cache.get_and_touch(
sql,
self.inner.connection_id,
self.inner.prepared_statement_epoch,
) {
return Outcome::Ok(cached);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_STMT_PREPARE);
buf.write_bytes(sql.as_bytes());
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (response_data, seq) = match self.read_packet().await {
Ok((data, seq)) => (data, seq),
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
if response_data.is_empty() {
return Outcome::Err(MySqlError::InvalidPacket("Empty prepare response".into()));
}
if response_data[0] == 0xff {
let err = Self::parse_error(&response_data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
return Outcome::Err(err);
}
if response_data[0] != 0x00 {
return Outcome::Err(MySqlError::InvalidPacket("Invalid prepare response".into()));
}
if response_data.len() < 12 {
return Outcome::Err(MySqlError::InvalidPacket(
"Prepare response too short".into(),
));
}
let parsed_header = (|| {
let mut reader = PacketReader::new(&response_data[1..]);
let statement_id = reader.read_u32_le()?;
let column_count = reader.read_u16_le()?;
let param_count = reader.read_u16_le()?;
let _reserved = reader.read_byte()?; let _warning_count = reader.read_u16_le()?;
Ok((statement_id, column_count, param_count))
})();
let (statement_id, column_count, param_count) = match parsed_header {
Ok(header) => header,
Err(e) => {
return outcome_from_error(e);
}
};
let expects_metadata_eof = Self::expects_metadata_eof(self.inner.capabilities);
let mut params = Vec::new();
if param_count > 0 {
for _ in 0..param_count {
let (param_data, seq) = match self.read_packet().await {
Ok((data, seq)) => (data, seq),
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
let param = match Self::parse_column_definition(¶m_data) {
Ok(column) => column,
Err(e) => return outcome_from_error(e),
};
params.push(param);
}
if expects_metadata_eof {
if let Err(e) = self.read_metadata_eof("parameters").await {
return outcome_from_error(e);
}
}
}
let mut columns = Vec::new();
if column_count > 0 {
for _ in 0..column_count {
let (col_data, seq) = match self.read_packet().await {
Ok((data, seq)) => (data, seq),
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
let column = match Self::parse_column_definition(&col_data) {
Ok(column) => column,
Err(e) => return outcome_from_error(e),
};
columns.push(column);
}
if expects_metadata_eof {
if let Err(e) = self.read_metadata_eof("columns").await {
return outcome_from_error(e);
}
}
}
self.inner.closed = false;
let stmt = MySqlStatement {
statement_id,
owner_connection_id: self.inner.connection_id,
owner_prepared_statement_epoch: self.inner.prepared_statement_epoch,
param_count,
column_count,
params,
columns,
};
let evicted_statement_id = self
.inner
.prepared_cache
.insert_returning_evicted_id(sql.to_string(), stmt.clone());
if let Some(statement_id) = evicted_statement_id
&& let Err(err) = self.close_prepared_statement_id(statement_id).await
{
return outcome_from_error(err);
}
Outcome::Ok(stmt)
}
async fn close_prepared_statement_id(&mut self, statement_id: u32) -> Result<(), MySqlError> {
if self.inner.closed {
return Err(MySqlError::ConnectionClosed);
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_STMT_CLOSE);
buf.write_u32_le(statement_id);
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
return Err(e);
}
self.inner.sequence = 0;
self.inner.closed = false;
Ok(())
}
pub async fn query_prepared(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<Vec<MySqlRow>, MySqlError> {
let was_closed_at_entry = self.inner.closed;
self.inner
.query_in_flight
.store(true, std::sync::atomic::Ordering::Release);
let result = self.query_prepared_inner_impl(cx, stmt, params).await;
if matches!(result, Outcome::Cancelled(_)) && self.inner.closed && !was_closed_at_entry {
self.wire_cancel_in_drain(cx).await;
}
self.inner
.query_in_flight
.store(false, std::sync::atomic::Ordering::Release);
result
}
async fn query_prepared_inner_impl(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<Vec<MySqlRow>, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
if stmt.owner_connection_id != self.inner.connection_id {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"prepared statement belongs to connection {} but current connection is {}",
stmt.owner_connection_id, self.inner.connection_id
)));
}
if stmt.owner_prepared_statement_epoch != self.inner.prepared_statement_epoch {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"prepared statement belongs to pooled checkout epoch {} but current epoch is {}",
stmt.owner_prepared_statement_epoch, self.inner.prepared_statement_epoch
)));
}
if params.len() != stmt.param_count as usize {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"Expected {} parameters, got {}",
stmt.param_count,
params.len()
)));
}
if let Err(e) = self.drain_abandoned_transaction().await {
return outcome_from_error(e);
}
match self.apply_statement_timeout(cx).await {
Outcome::Ok(()) => {}
Outcome::Err(err) => return Outcome::Err(err),
Outcome::Cancelled(reason) => return Outcome::Cancelled(reason),
Outcome::Panicked(payload) => return Outcome::Panicked(payload),
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_STMT_EXECUTE);
buf.write_u32_le(stmt.statement_id);
buf.write_byte(0x00); buf.write_u32_le(1);
if let Err(e) = write_stmt_execute_params(&mut buf, params) {
return outcome_from_error(e);
}
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (response_data, seq) = match self.read_packet().await {
Ok((data, seq)) => (data, seq),
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
if response_data.is_empty() {
return Outcome::Err(MySqlError::InvalidPacket(
"Empty prepared query response".into(),
));
}
match response_data[0] {
0x00 => match Self::parse_ok_packet(&response_data) {
Ok(ok_packet) => {
self.inner.status_flags = ok_packet.status_flags;
self.inner.closed = false;
Outcome::Ok(Vec::new())
}
Err(e) => {
self.inner.closed = false;
outcome_from_error(e)
}
},
0xFF => {
let err = Self::parse_error(&response_data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
Outcome::Err(err)
}
_ => match self.read_binary_result_set(cx, &response_data).await {
Ok(rows) => {
self.inner.closed = false;
Outcome::Ok(rows)
}
Err(MySqlError::Cancelled(reason)) => Outcome::Cancelled(reason),
Err(e) => outcome_from_error(e),
},
}
}
pub async fn execute_prepared(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<u64, MySqlError> {
let was_closed_at_entry = self.inner.closed;
self.inner
.query_in_flight
.store(true, std::sync::atomic::Ordering::Release);
let result = self.execute_prepared_inner_impl(cx, stmt, params).await;
if matches!(result, Outcome::Cancelled(_)) && self.inner.closed && !was_closed_at_entry {
self.wire_cancel_in_drain(cx).await;
}
self.inner
.query_in_flight
.store(false, std::sync::atomic::Ordering::Release);
result
}
async fn execute_prepared_inner_impl(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<u64, MySqlError> {
if cx.checkpoint().is_err() {
return Outcome::Cancelled(
cx.cancel_reason()
.unwrap_or_else(|| CancelReason::user("cancelled")),
);
}
if self.inner.closed {
return Outcome::Err(MySqlError::ConnectionClosed);
}
if stmt.owner_connection_id != self.inner.connection_id {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"prepared statement belongs to connection {} but current connection is {}",
stmt.owner_connection_id, self.inner.connection_id
)));
}
if stmt.owner_prepared_statement_epoch != self.inner.prepared_statement_epoch {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"prepared statement belongs to pooled checkout epoch {} but current epoch is {}",
stmt.owner_prepared_statement_epoch, self.inner.prepared_statement_epoch
)));
}
if params.len() != stmt.param_count as usize {
return Outcome::Err(MySqlError::InvalidParameter(format!(
"Expected {} parameters, got {}",
stmt.param_count,
params.len()
)));
}
if let Err(e) = self.drain_abandoned_transaction().await {
return outcome_from_error(e);
}
match self.apply_statement_timeout(cx).await {
Outcome::Ok(()) => {}
Outcome::Err(err) => return Outcome::Err(err),
Outcome::Cancelled(reason) => return Outcome::Cancelled(reason),
Outcome::Panicked(payload) => return Outcome::Panicked(payload),
}
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_STMT_EXECUTE);
buf.write_u32_le(stmt.statement_id);
buf.write_byte(0x00); buf.write_u32_le(1);
if let Err(e) = write_stmt_execute_params(&mut buf, params) {
return outcome_from_error(e);
}
let packet = buf.build_packet();
self.inner.closed = true;
if let Err(e) = self.write_all(&packet.bytes).await {
return outcome_from_error(e);
}
self.inner.sequence = packet.next_sequence;
let (response_data, seq) = match self.read_packet().await {
Ok((data, seq)) => (data, seq),
Err(e) => return outcome_from_error(e),
};
self.inner.sequence = seq.wrapping_add(1);
if response_data.is_empty() {
return Outcome::Err(MySqlError::InvalidPacket("Empty execute response".into()));
}
if response_data[0] == 0xff {
let err = Self::parse_error(&response_data);
if matches!(&err, MySqlError::Server { .. }) {
self.inner.closed = false;
}
return Outcome::Err(err);
}
if response_data[0] == 0x00 {
let ok_packet = match Self::parse_ok_packet(&response_data) {
Ok(packet) => packet,
Err(e) => {
self.inner.closed = false;
return outcome_from_error(e);
}
};
self.inner.status_flags = ok_packet.status_flags;
self.inner.closed = false;
return Outcome::Ok(ok_packet.affected_rows);
}
Outcome::Err(MySqlError::InvalidPacket(
"Unexpected execute response".into(),
))
}
pub fn set_max_result_rows(&mut self, max: usize) {
self.inner.max_result_rows = max;
}
#[must_use]
pub fn max_result_rows(&self) -> usize {
self.inner.max_result_rows
}
async fn drain_abandoned_transaction(&mut self) -> Result<(), MySqlError> {
if !self.inner.needs_rollback {
return Ok(());
}
self.inner.closed = true;
self.raw_query_expect_ok("ROLLBACK", "implicit ROLLBACK")
.await?;
self.inner.needs_rollback = false;
if let Some(previous) = self.inner.session_isolation_restore.take() {
let restore_sql = format!("SET SESSION TRANSACTION ISOLATION LEVEL {previous}");
self.raw_query_expect_ok(&restore_sql, "session isolation restore")
.await?;
}
self.inner.closed = false;
Ok(())
}
async fn raw_query_expect_ok(&mut self, sql: &str, what: &str) -> Result<(), MySqlError> {
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_QUERY);
buf.write_bytes(sql.as_bytes());
let packet = buf.build_packet();
if let Err(e) = self.write_all(&packet.bytes).await {
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
return Err(e);
}
self.inner.sequence = packet.next_sequence;
let (data, seq) = match self.read_packet().await {
Ok(res) => res,
Err(e) => {
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
return Err(e);
}
};
self.inner.sequence = seq.wrapping_add(1);
match data.first() {
Some(0x00) => {
self.inner.status_flags = Self::parse_ok_packet(&data)?.status_flags;
Ok(())
}
Some(0xFF) => {
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
Err(Self::parse_error(&data))
}
_ => {
let _ = self.inner.stream.shutdown(std::net::Shutdown::Both);
Err(MySqlError::Protocol(format!(
"unexpected response to {what}"
)))
}
}
}
async fn write_all(&mut self, data: &[u8]) -> Result<(), MySqlError> {
let mut pos = 0;
while pos < data.len() {
let written = std::future::poll_fn(|cx| {
if crate::cx::Cx::with_current(|c| c.checkpoint().is_err()).unwrap_or(false) {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Interrupted,
"cancelled",
)));
}
Pin::new(&mut self.inner.stream).poll_write(cx, &data[pos..])
})
.await
.map_err(stream_io_error)?;
if written == 0 {
return Err(MySqlError::Io(io::Error::new(
io::ErrorKind::WriteZero,
"failed to write data",
)));
}
pos += written;
}
Ok(())
}
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), MySqlError> {
read_exact_from(&mut self.inner.stream, buf).await
}
async fn read_packet(&mut self) -> Result<(Vec<u8>, u8), MySqlError> {
let mut expected_seq = self.inner.sequence;
let mut last_seq;
let mut data = Vec::new();
loop {
let mut header = [0u8; 4];
self.read_exact(&mut header).await?;
let (len, seq) = Self::decode_packet_header(header, expected_seq)?;
last_seq = seq;
if len > 0 {
let start = data.len();
let new_len = start.saturating_add(len as usize);
if new_len > MAX_REASSEMBLED_PACKET_SIZE {
return Err(MySqlError::Protocol(format!(
"packet payload {new_len} exceeds maximum allowed {MAX_REASSEMBLED_PACKET_SIZE}"
)));
}
data.resize(new_len, 0);
self.read_exact(&mut data[start..]).await?;
}
expected_seq = expected_seq.wrapping_add(1);
if len < MAX_PACKET_SIZE {
return Ok((data, last_seq));
}
}
}
#[inline]
fn decode_packet_header(header: [u8; 4], expected_seq: u8) -> Result<(u32, u8), MySqlError> {
let len = u32::from(header[0]) | (u32::from(header[1]) << 8) | (u32::from(header[2]) << 16);
let seq = header[3];
if seq != expected_seq {
return Err(MySqlError::Protocol(format!(
"packet sequence mismatch: expected {expected_seq}, got {seq}"
)));
}
if len > MAX_PACKET_SIZE {
return Err(MySqlError::Protocol(format!(
"packet length {len} exceeds maximum allowed {MAX_PACKET_SIZE}"
)));
}
Ok((len, seq))
}
fn parse_error(data: &[u8]) -> MySqlError {
if data.is_empty() || data[0] != 0xFF {
return MySqlError::Protocol("not an error packet".to_string());
}
let mut reader = PacketReader::new(&data[1..]);
let code = match reader.read_u16_le() {
Ok(c) => c,
Err(e) => return e,
};
let sql_state = if reader.remaining() > 0 && data.get(reader.pos + 1) == Some(&b'#') {
reader.pos += 1; reader.read_bytes(5).map_or_else(
|_| "HY000".to_string(),
|state| std::str::from_utf8(state).unwrap_or("HY000").to_string(),
)
} else {
"HY000".to_string()
};
let message = std::str::from_utf8(reader.read_rest())
.unwrap_or("unknown error")
.to_string();
MySqlError::Server {
code,
sql_state,
message,
}
}
}
fn validate_auth_plugin_switch(
initial_plugin: &str,
switch_plugin: &str,
options: &MySqlConnectOptions,
) -> Result<(), MySqlError> {
let is_downgrade = matches!(
(initial_plugin, switch_plugin),
("caching_sha2_password", "mysql_native_password")
);
if is_downgrade && !options.insecure_allow_auth_switch_downgrade {
return Err(MySqlError::UnsupportedAuthPlugin(format!(
"auth switch downgrade from {initial_plugin} to {switch_plugin} rejected by default \
— set MySqlConnectOptions::insecure_allow_auth_switch_downgrade = true to opt in"
)));
}
Ok(())
}
pub trait ToSql: Sync {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError>;
fn mysql_type_code(&self) -> u8;
fn is_null(&self) -> bool {
false
}
fn is_unsigned(&self) -> bool {
false
}
}
trait StaticMySqlTypeInfo {
fn static_mysql_type_code() -> u8;
fn static_is_unsigned() -> bool {
false
}
}
impl ToSql for bool {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(vec![u8::from(*self)])
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for bool {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for i8 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok((*self as u8).to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
}
impl StaticMySqlTypeInfo for i8 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
}
impl ToSql for i16 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_SHORT
}
}
impl StaticMySqlTypeInfo for i16 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_SHORT
}
}
impl ToSql for i32 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_LONG
}
}
impl StaticMySqlTypeInfo for i32 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_LONG
}
}
impl ToSql for i64 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
}
impl StaticMySqlTypeInfo for i64 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
}
impl ToSql for u8 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for u8 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_TINY
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for u16 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_SHORT
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for u16 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_SHORT
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for u32 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_LONG
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for u32 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_LONG
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for u64 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for u64 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for usize {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok((*self as u64).to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
fn is_unsigned(&self) -> bool {
true
}
}
impl StaticMySqlTypeInfo for usize {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_LONGLONG
}
fn static_is_unsigned() -> bool {
true
}
}
impl ToSql for f32 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_FLOAT
}
}
impl StaticMySqlTypeInfo for f32 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_FLOAT
}
}
impl ToSql for f64 {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(self.to_le_bytes().to_vec())
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_DOUBLE
}
}
impl StaticMySqlTypeInfo for f64 {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_DOUBLE
}
}
impl ToSql for str {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(encode_lenenc_bytes(self.as_bytes()))
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_VAR_STRING
}
}
impl StaticMySqlTypeInfo for str {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_VAR_STRING
}
}
impl ToSql for String {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
self.as_str().to_sql()
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_VAR_STRING
}
}
impl StaticMySqlTypeInfo for String {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_VAR_STRING
}
}
impl ToSql for [u8] {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
Ok(encode_lenenc_bytes(self))
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_BLOB
}
}
impl StaticMySqlTypeInfo for [u8] {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_BLOB
}
}
impl ToSql for Vec<u8> {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
self.as_slice().to_sql()
}
fn mysql_type_code(&self) -> u8 {
mysql_type::MYSQL_TYPE_BLOB
}
}
impl StaticMySqlTypeInfo for Vec<u8> {
fn static_mysql_type_code() -> u8 {
mysql_type::MYSQL_TYPE_BLOB
}
}
impl<T: StaticMySqlTypeInfo + ?Sized> StaticMySqlTypeInfo for &T {
fn static_mysql_type_code() -> u8 {
T::static_mysql_type_code()
}
fn static_is_unsigned() -> bool {
T::static_is_unsigned()
}
}
impl<T: ToSql + StaticMySqlTypeInfo> ToSql for Option<T> {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
match self {
Some(value) => value.to_sql(),
None => Ok(vec![]),
}
}
fn mysql_type_code(&self) -> u8 {
match self {
Some(value) => value.mysql_type_code(),
None => T::static_mysql_type_code(),
}
}
fn is_null(&self) -> bool {
self.is_none()
}
fn is_unsigned(&self) -> bool {
match self {
Some(value) => value.is_unsigned(),
None => T::static_is_unsigned(),
}
}
}
impl<T: ToSql + ?Sized> ToSql for &T {
fn to_sql(&self) -> Result<Vec<u8>, MySqlError> {
(*self).to_sql()
}
fn mysql_type_code(&self) -> u8 {
(*self).mysql_type_code()
}
fn is_null(&self) -> bool {
(*self).is_null()
}
fn is_unsigned(&self) -> bool {
(*self).is_unsigned()
}
}
fn write_stmt_execute_params(
buf: &mut PacketBuffer,
params: &[&dyn ToSql],
) -> Result<(), MySqlError> {
if params.is_empty() {
return Ok(());
}
let mut null_bitmap = vec![0; params.len().div_ceil(8)];
for (idx, param) in params.iter().enumerate() {
if param.is_null() {
null_bitmap[idx / 8] |= 1 << (idx % 8);
}
}
buf.write_bytes(&null_bitmap);
buf.write_byte(0x01);
for param in params {
let mut type_field = u16::from(param.mysql_type_code());
if param.is_unsigned() {
type_field |= param_flag::UNSIGNED_LE_U16;
}
buf.write_u16_le(type_field);
}
for param in params {
if param.is_null() {
continue;
}
buf.write_bytes(¶m.to_sql()?);
}
Ok(())
}
fn encode_lenenc_bytes(data: &[u8]) -> Vec<u8> {
let mut buf = PacketBuffer::new();
buf.write_lenenc_int(u64::try_from(data.len()).unwrap_or(u64::MAX));
buf.write_bytes(data);
buf.buf
}
mod mysql_type {
pub const MYSQL_TYPE_TINY: u8 = 1;
pub const MYSQL_TYPE_SHORT: u8 = 2;
pub const MYSQL_TYPE_LONG: u8 = 3;
pub const MYSQL_TYPE_FLOAT: u8 = 4;
pub const MYSQL_TYPE_DOUBLE: u8 = 5;
pub const MYSQL_TYPE_LONGLONG: u8 = 8;
pub const MYSQL_TYPE_VAR_STRING: u8 = 253;
pub const MYSQL_TYPE_BLOB: u8 = 252;
}
mod param_flag {
pub const UNSIGNED_LE_U16: u16 = 0x80_00;
}
#[derive(Debug, Clone)]
pub struct MySqlStatement {
statement_id: u32,
owner_connection_id: u32,
owner_prepared_statement_epoch: u64,
param_count: u16,
column_count: u16,
params: Vec<MySqlColumn>,
columns: Vec<MySqlColumn>,
}
impl MySqlStatement {
#[must_use]
pub fn owner_connection_id(&self) -> u32 {
self.owner_connection_id
}
#[must_use]
pub fn owner_prepared_statement_epoch(&self) -> u64 {
self.owner_prepared_statement_epoch
}
#[must_use]
pub fn param_count(&self) -> u16 {
self.param_count
}
#[must_use]
pub fn column_count(&self) -> u16 {
self.column_count
}
#[must_use]
pub fn params(&self) -> &[MySqlColumn] {
&self.params
}
#[must_use]
pub fn columns(&self) -> &[MySqlColumn] {
&self.columns
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct MySqlPreparedCacheStats {
pub hits: u64,
pub misses: u64,
pub evictions: u64,
}
impl MySqlPreparedCacheStats {
#[must_use]
pub const fn lookups(&self) -> u64 {
self.hits + self.misses
}
#[must_use]
pub fn hit_ratio(&self) -> f64 {
let total = self.lookups();
if total == 0 {
0.0
} else {
self.hits as f64 / total as f64
}
}
}
struct MySqlPreparedStatementCache {
entries: HashMap<String, MySqlStatement>,
lru: VecDeque<String>,
cap: usize,
stats: MySqlPreparedCacheStats,
}
impl MySqlPreparedStatementCache {
fn new(cap: usize) -> Self {
Self {
entries: HashMap::with_capacity(cap.min(64)),
lru: VecDeque::with_capacity(cap.min(64)),
cap,
stats: MySqlPreparedCacheStats::default(),
}
}
fn stats(&self) -> MySqlPreparedCacheStats {
self.stats
}
fn get_and_touch(
&mut self,
sql: &str,
owner_connection_id: u32,
owner_prepared_statement_epoch: u64,
) -> Option<MySqlStatement> {
let Some(mut stmt) = self.entries.get(sql).cloned() else {
self.stats.misses += 1;
return None;
};
self.stats.hits += 1;
stmt.owner_connection_id = owner_connection_id;
stmt.owner_prepared_statement_epoch = owner_prepared_statement_epoch;
if let Some(pos) = self.lru.iter().position(|key| key == sql) {
if let Some(key) = self.lru.remove(pos) {
self.lru.push_back(key);
}
}
Some(stmt)
}
fn insert_returning_evicted_id(&mut self, sql: String, stmt: MySqlStatement) -> Option<u32> {
if self.cap == 0 {
self.stats.evictions += 1;
return Some(stmt.statement_id);
}
let mut evicted = None;
if let Some(old) = self.entries.remove(&sql) {
if let Some(pos) = self.lru.iter().position(|key| key == &sql) {
self.lru.remove(pos);
}
evicted = Some(old.statement_id);
} else if self.entries.len() >= self.cap
&& let Some(victim_sql) = self.lru.pop_front()
&& let Some(victim_stmt) = self.entries.remove(&victim_sql)
{
evicted = Some(victim_stmt.statement_id);
}
if evicted.is_some() {
self.stats.evictions += 1;
}
self.lru.push_back(sql.clone());
self.entries.insert(sql, stmt);
evicted
}
#[cfg(test)]
fn len(&self) -> usize {
self.entries.len()
}
}
#[cfg(feature = "test-internals")]
#[doc(hidden)]
#[derive(Debug, Clone, Copy)]
pub struct PreparedCacheBenchReport {
pub stats: MySqlPreparedCacheStats,
pub executions: u64,
pub prepares_issued: u64,
pub prepares_avoided: u64,
}
#[cfg(feature = "test-internals")]
impl PreparedCacheBenchReport {
#[must_use]
pub fn prepares_avoided_ratio(&self) -> f64 {
self.stats.hit_ratio()
}
}
#[cfg(feature = "test-internals")]
#[doc(hidden)]
#[must_use]
pub fn bench_prepared_cache_repeated_workload(
capacity: usize,
distinct_queries: usize,
repetitions: usize,
) -> PreparedCacheBenchReport {
let mut cache = MySqlPreparedStatementCache::new(capacity);
let queries: Vec<String> = (0..distinct_queries)
.map(|i| format!("SELECT c0, c1 FROM bench_t WHERE k = ? /* q{i} */"))
.collect();
let mut next_statement_id: u32 = 1;
let mut executions: u64 = 0;
for _ in 0..repetitions {
for sql in &queries {
executions += 1;
if cache.get_and_touch(sql, 1, 1).is_none() {
let stmt = MySqlStatement {
statement_id: next_statement_id,
owner_connection_id: 1,
owner_prepared_statement_epoch: 1,
param_count: 1,
column_count: 2,
params: Vec::new(),
columns: Vec::new(),
};
next_statement_id = next_statement_id.wrapping_add(1);
cache.insert_returning_evicted_id(sql.clone(), stmt);
}
}
}
let stats = cache.stats();
PreparedCacheBenchReport {
stats,
executions,
prepares_issued: stats.misses,
prepares_avoided: stats.hits,
}
}
pub struct MySqlConnectionManager {
options: MySqlConnectOptions,
}
impl fmt::Debug for MySqlConnectionManager {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MySqlConnectionManager")
.field("options", &self.options)
.finish()
}
}
impl MySqlConnectionManager {
#[must_use]
pub fn new(options: MySqlConnectOptions) -> Self {
Self { options }
}
#[must_use]
pub fn options(&self) -> &MySqlConnectOptions {
&self.options
}
}
impl crate::database::pool::AsyncConnectionManager for MySqlConnectionManager {
type Connection = MySqlConnection;
type Error = MySqlError;
async fn connect(&self, cx: &Cx) -> Outcome<Self::Connection, Self::Error> {
MySqlConnection::connect_with_options(cx, self.options.clone()).await
}
async fn is_valid(&self, _cx: &Cx, conn: &mut Self::Connection) -> bool {
!conn.inner.closed && !conn.in_transaction() && !conn.inner.needs_rollback
}
fn release_check(&self, conn: &mut Self::Connection) -> bool {
if conn.inner.closed || conn.in_transaction() || conn.inner.needs_rollback {
return false;
}
conn.invalidate_prepared_statements_for_pool_return();
true
}
}
pub struct MySqlTransaction<'a> {
conn: &'a mut MySqlConnection,
finished: bool,
isolation_level: Option<IsolationLevel>,
read_only: bool,
obligation: Option<ObligationToken<TransactionKind>>,
}
fn reserve_transaction_obligation(cx: &Cx) -> Option<ObligationToken<TransactionKind>> {
let region = cx.region_id();
if region.as_u64() == 0 {
None
} else {
Some(ObligationToken::reserve("db-transaction:mysql", region))
}
}
impl MySqlTransaction<'_> {
#[must_use]
pub const fn isolation_level(&self) -> Option<IsolationLevel> {
self.isolation_level
}
#[must_use]
pub const fn is_read_only(&self) -> bool {
self.read_only
}
#[must_use]
pub(crate) const fn requires_rollback_before_commit(&self) -> bool {
self.conn.inner.needs_rollback
}
pub(crate) fn poison_for_rollback(&mut self) {
self.conn.inner.needs_rollback = true;
}
pub async fn commit(mut self, cx: &Cx) -> Outcome<(), MySqlError> {
if self.finished {
trace_database_transaction(cx, "mysql", "commit", "already_finished");
return Outcome::Err(MySqlError::TransactionFinished);
}
trace_database_transaction(cx, "mysql", "commit", "start");
match self.conn.execute_unchecked_internal(cx, "COMMIT").await {
Outcome::Ok(_) => {
self.finished = true;
self.conn.restore_session_isolation(cx).await;
if let Some(token) = self.obligation.take() {
let _ = token.commit();
}
trace_database_transaction(cx, "mysql", "commit", "ok");
Outcome::Ok(())
}
Outcome::Err(e) => {
trace_database_transaction(cx, "mysql", "commit", "err");
outcome_from_error(e)
}
Outcome::Cancelled(r) => {
trace_database_transaction(cx, "mysql", "commit", "cancelled");
Outcome::Cancelled(r)
}
Outcome::Panicked(p) => {
trace_database_transaction(cx, "mysql", "commit", "panicked");
Outcome::Panicked(p)
}
}
}
pub async fn rollback(mut self, cx: &Cx) -> Outcome<(), MySqlError> {
if self.finished {
trace_database_transaction(cx, "mysql", "rollback", "already_finished");
return Outcome::Err(MySqlError::TransactionFinished);
}
trace_database_transaction(cx, "mysql", "rollback", "start");
match self.conn.execute_unchecked_internal(cx, "ROLLBACK").await {
Outcome::Ok(_) => {
self.finished = true;
self.conn.restore_session_isolation(cx).await;
if let Some(token) = self.obligation.take() {
let _ = token.abort();
}
trace_database_transaction(cx, "mysql", "rollback", "ok");
Outcome::Ok(())
}
Outcome::Err(e) => {
trace_database_transaction(cx, "mysql", "rollback", "err");
outcome_from_error(e)
}
Outcome::Cancelled(r) => {
trace_database_transaction(cx, "mysql", "rollback", "cancelled");
Outcome::Cancelled(r)
}
Outcome::Panicked(p) => {
trace_database_transaction(cx, "mysql", "rollback", "panicked");
Outcome::Panicked(p)
}
}
}
#[deprecated(
note = "use query_static_sql for trusted-literal SQL or the prepared-statement APIs for parameterized queries (br-asupersync-0fxbp6)"
)]
pub async fn query(&mut self, cx: &Cx, sql: &str) -> Outcome<Vec<MySqlRow>, MySqlError> {
self.query_unchecked_internal(cx, sql).await
}
async fn query_unchecked_internal(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
if self.finished {
return Outcome::Err(MySqlError::TransactionFinished);
}
self.conn.query_unchecked_internal(cx, sql).await
}
#[deprecated(
note = "use execute_static_sql for trusted-literal SQL or the prepared-statement APIs for parameterized commands (br-asupersync-0fxbp6)"
)]
pub async fn execute(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
self.execute_unchecked_internal(cx, sql).await
}
async fn execute_unchecked_internal(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
if self.finished {
return Outcome::Err(MySqlError::TransactionFinished);
}
self.conn.execute_unchecked_internal(cx, sql).await
}
pub async fn execute_static_sql(&mut self, cx: &Cx, sql: &str) -> Outcome<u64, MySqlError> {
self.execute_unchecked_internal(cx, sql).await
}
pub async fn query_static_sql(
&mut self,
cx: &Cx,
sql: &str,
) -> Outcome<Vec<MySqlRow>, MySqlError> {
self.query_unchecked_internal(cx, sql).await
}
pub async fn prepare(&mut self, cx: &Cx, sql: &str) -> Outcome<MySqlStatement, MySqlError> {
if self.finished {
return Outcome::Err(MySqlError::TransactionFinished);
}
self.conn.prepare(cx, sql).await
}
pub async fn execute_prepared(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<u64, MySqlError> {
if self.finished {
return Outcome::Err(MySqlError::TransactionFinished);
}
self.conn.execute_prepared(cx, stmt, params).await
}
pub async fn query_prepared(
&mut self,
cx: &Cx,
stmt: &MySqlStatement,
params: &[&dyn ToSql],
) -> Outcome<Vec<MySqlRow>, MySqlError> {
if self.finished {
return Outcome::Err(MySqlError::TransactionFinished);
}
self.conn.query_prepared(cx, stmt, params).await
}
}
impl Drop for MySqlTransaction<'_> {
fn drop(&mut self) {
if let Some(token) = self.obligation.take() {
let _ = token.abort();
}
if !self.finished {
self.poison_for_rollback();
}
}
}
#[doc(hidden)]
pub fn fuzz_parse_ok_packet_fields(data: &[u8]) -> Result<(u64, u16), MySqlError> {
MySqlConnection::parse_ok_packet(data).map(|packet| (packet.affected_rows, packet.status_flags))
}
#[doc(hidden)]
pub fn fuzz_parse_handshake_protocol_41(
data: &[u8],
connects_with_db: bool,
) -> Result<FuzzHandshakeProtocol41, MySqlError> {
const MIN_HANDSHAKE_SIZE: usize = 35;
if data.len() < MIN_HANDSHAKE_SIZE {
return Err(MySqlError::InvalidPacket(format!(
"handshake packet too short: {} bytes, minimum required: {}",
data.len(),
MIN_HANDSHAKE_SIZE
)));
}
let mut reader = PacketReader::new(data);
let protocol_version = reader.read_byte()?;
if protocol_version != 10 {
return Err(MySqlError::Protocol(format!(
"unsupported protocol version: {protocol_version}"
)));
}
let _server_version = reader.read_null_terminated()?;
let _connection_id = reader.read_u32_le()?;
let auth_data_1 = reader.read_bytes(8)?;
let _filler = reader.read_byte()?;
let cap_lower = reader.read_u16_le()?;
let _charset = reader.read_byte()?;
let _status_flags = reader.read_u16_le()?;
let cap_upper = reader.read_u16_le()?;
let server_capabilities = u32::from(cap_lower) | (u32::from(cap_upper) << 16);
let missing_required_caps = (capability::CLIENT_PROTOCOL_41
| capability::CLIENT_SECURE_CONNECTION)
& !server_capabilities;
if missing_required_caps != 0 {
let mut missing = Vec::new();
if missing_required_caps & capability::CLIENT_PROTOCOL_41 != 0 {
missing.push("CLIENT_PROTOCOL_41");
}
if missing_required_caps & capability::CLIENT_SECURE_CONNECTION != 0 {
missing.push("CLIENT_SECURE_CONNECTION");
}
return Err(MySqlError::Protocol(format!(
"server handshake missing required capabilities: {}",
missing.join(", ")
)));
}
let auth_data_len = reader.read_byte()?;
let _reserved = reader.read_bytes(10)?;
let mut auth_plugin_data_len = auth_data_1.len();
if server_capabilities & capability::CLIENT_SECURE_CONNECTION != 0 {
let part2_len = std::cmp::max(13, auth_data_len.saturating_sub(8)) as usize;
let auth_data_2 = reader.read_bytes(part2_len.min(reader.remaining()))?;
let end = if auth_data_2.last() == Some(&0) {
auth_data_2.len() - 1
} else {
auth_data_2.len()
};
auth_plugin_data_len += end;
}
let auth_plugin_name =
if server_capabilities & capability::CLIENT_PLUGIN_AUTH != 0 && reader.remaining() > 0 {
reader.read_null_terminated()?.to_string()
} else {
"mysql_native_password".to_string()
};
let client_capabilities =
MySqlConnection::client_handshake_response_capabilities(connects_with_db);
let negotiated_capabilities =
MySqlConnection::negotiated_capabilities(server_capabilities, client_capabilities);
Ok(FuzzHandshakeProtocol41 {
server_capabilities,
client_capabilities,
negotiated_capabilities,
auth_plugin_name,
auth_plugin_data_len,
})
}
#[doc(hidden)]
pub fn fuzz_parse_column_definition(data: &[u8]) -> Result<MySqlColumn, MySqlError> {
MySqlConnection::parse_column_definition(data)
}
#[doc(hidden)]
pub fn fuzz_decode_packet_header(
header: [u8; 4],
expected_seq: u8,
) -> Result<(u32, u8), MySqlError> {
MySqlConnection::decode_packet_header(header, expected_seq)
}
#[doc(hidden)]
#[must_use]
pub fn fuzz_parse_error_packet(data: &[u8]) -> MySqlError {
MySqlConnection::parse_error(data)
}
#[doc(hidden)]
pub fn fuzz_parse_text_row(
data: &[u8],
columns: &[MySqlColumn],
) -> Result<Vec<MySqlValue>, MySqlError> {
MySqlConnection::parse_text_row(data, columns)
}
#[doc(hidden)]
pub fn fuzz_parse_binary_row(
data: &[u8],
columns: &[MySqlColumn],
) -> Result<Vec<MySqlValue>, MySqlError> {
MySqlConnection::parse_binary_row(data, columns)
}
#[doc(hidden)]
pub fn fuzz_parse_data_row_or_terminator(
data: &[u8],
columns: &[MySqlColumn],
deprecate_eof: bool,
) -> Result<Option<Vec<MySqlValue>>, MySqlError> {
MySqlConnection::parse_data_row_or_terminator(data, columns, deprecate_eof)
}
#[doc(hidden)]
pub fn fuzz_build_stmt_execute_packet(
statement_id: u32,
params: &[&dyn ToSql],
) -> Result<Vec<u8>, MySqlError> {
let mut buf = PacketBuffer::new();
buf.set_sequence(0);
buf.write_byte(command::COM_STMT_EXECUTE);
buf.write_u32_le(statement_id);
buf.write_byte(0x00);
buf.write_u32_le(1);
write_stmt_execute_params(&mut buf, params)?;
Ok(buf.build_packet().bytes)
}
#[cfg(test)]
include!("mysql_tests.rs");
#[cfg(test)]
#[path = "mysql_load_data_infile_security_audit.rs"]
mod mysql_load_data_infile_security_audit;