Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 31 additions & 5 deletions sqlx-mysql/src/connection/executor.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
use super::MySqlStream;
use crate::connection::stream::Waiting;
use crate::connection::stream::{PrepareProgress, Waiting};
use crate::error::Error;
use crate::executor::{Execute, Executor};
use crate::ext::ustr::UStr;
Expand Down Expand Up @@ -33,6 +33,13 @@ impl MySqlConnection {
// https://dev.mysql.com/doc/internals/en/com-stmt-prepare.html
// https://dev.mysql.com/doc/internals/en/com-stmt-prepare-response.html#packet-COM_STMT_PREPARE_OK

// Queued before the request, so the drain can finish and close a
// response a cancelled prepare left behind.
self.inner
.stream
.waiting
.push_back(Waiting::Prepare(PrepareProgress::Header));

self.inner
.stream
.send_packet(Prepare { query: sql })
Expand Down Expand Up @@ -70,6 +77,10 @@ impl MySqlConnection {
column_names: Arc::new(column_names),
};

// Queued so the drain closes the statement if the caller is
// dropped before caching or closing it.
self.inner.stream.close_pending.push(id);

Ok((id, metadata))
}

Expand All @@ -85,11 +96,15 @@ impl MySqlConnection {
let (id, metadata) = self.prepare_statement(sql).await?;

// in case of the cache being full, close the least recently used statement
if let Some((id, _)) = self
let evicted = self
.inner
.cache_statement
.insert(sql, (id, metadata.clone()))
{
.insert(sql, (id, metadata.clone()));

// Cached, so no longer closed by the drain.
self.inner.stream.close_pending.pop();

if let Some((id, _)) = evicted {
self.inner
.stream
.send_packet(StmtClose { statement: id })
Expand All @@ -110,7 +125,6 @@ impl MySqlConnection {
let mut logger = QueryLogger::new(sql, self.inner.log_settings.clone());

self.inner.stream.wait_until_ready().await?;
self.inner.stream.waiting.push_back(Waiting::Result);

Ok(try_stream! {
let sql = logger.sql().as_str();
Expand All @@ -136,6 +150,10 @@ impl MySqlConnection {
);
}

// Queued only once the request is sent, so a cancelled
// prepare or a failed parameter-count check leaves none.
self.inner.stream.waiting.push_back(Waiting::Result);

// https://dev.mysql.com/doc/internals/en/com-stmt-execute.html
self.inner.stream
.send_packet(StatementExecute {
Expand All @@ -151,6 +169,7 @@ impl MySqlConnection {
.await?;

if arguments.types.len() != metadata.parameters {
// Not cached and not closed here; the drain closes it.
return Err(
err_protocol!(
"prepared statement expected {} parameters but {} parameters were provided",
Expand All @@ -160,6 +179,8 @@ impl MySqlConnection {
);
}

self.inner.stream.waiting.push_back(Waiting::Result);

// https://dev.mysql.com/doc/internals/en/com-stmt-execute.html
self.inner.stream
.send_packet(StatementExecute {
Expand All @@ -168,11 +189,14 @@ impl MySqlConnection {
})
.await?;

self.inner.stream.close_pending.pop();
self.inner.stream.send_packet(StmtClose { statement: id }).await?;

MySqlValueFormat::Binary
}
} else {
self.inner.stream.waiting.push_back(Waiting::Result);

// https://dev.mysql.com/doc/internals/en/com-query.html
self.inner.stream.send_packet(Query(sql)).await?;

Expand Down Expand Up @@ -335,6 +359,7 @@ impl<'c> Executor<'c> for &'c mut MySqlConnection {
} else {
let (id, metadata) = self.prepare_statement(sql.as_str()).await?;

self.inner.stream.close_pending.pop();
self.inner
.stream
.send_packet(StmtClose { statement: id })
Expand Down Expand Up @@ -365,6 +390,7 @@ impl<'c> Executor<'c> for &'c mut MySqlConnection {

let (id, metadata) = self.prepare_statement(sql.as_str()).await?;

self.inner.stream.close_pending.pop();
self.inner
.stream
.send_packet(StmtClose { statement: id })
Expand Down
134 changes: 124 additions & 10 deletions sqlx-mysql/src/connection/stream.rs
Original file line number Diff line number Diff line change
@@ -1,13 +1,14 @@
use std::collections::VecDeque;
use std::ops::{Deref, DerefMut};
use std::ops::{ControlFlow, Deref, DerefMut};

use bytes::{Buf, Bytes, BytesMut};
use bytes::{Bytes, BytesMut};

use crate::error::Error;
use crate::io::MySqlBufExt;
use crate::io::{ProtocolDecode, ProtocolEncode};
use crate::net::{BufferedSocket, Socket};
use crate::protocol::response::{EofPacket, ErrPacket, OkPacket, Status};
use crate::protocol::statement::{PrepareOk, StmtClose};
use crate::protocol::{Capabilities, Packet};
use crate::{MySqlConnectOptions, MySqlDatabaseError};

Expand All @@ -18,6 +19,8 @@ pub struct MySqlStream<S = Box<dyn Socket>> {
pub(super) capabilities: Capabilities,
pub(crate) sequence_id: u8,
pub(crate) waiting: VecDeque<Waiting>,
// statements a dropped operation left open; closed by the drain
pub(crate) close_pending: Vec<u32>,
pub(crate) is_tls: bool,
}

Expand All @@ -28,6 +31,21 @@ pub(crate) enum Waiting {

// waiting for a row within a result set
Row,

// waiting for (the rest of) a COM_STMT_PREPARE response
Prepare(PrepareProgress),
}

#[derive(Debug, PartialEq, Eq)]
pub(crate) enum PrepareProgress {
// the stmt-prepare-ok header has not been read yet
Header,

// the header was read; this many definition and EOF packets follow
Tail {
statement_id: u32,
packets_left: u32,
},
}

impl<S: Socket> MySqlStream<S> {
Expand All @@ -54,6 +72,7 @@ impl<S: Socket> MySqlStream<S> {

Self {
waiting: VecDeque::new(),
close_pending: Vec::new(),
capabilities,
server_version: (0, 0, 0),
sequence_id: 0,
Expand All @@ -63,11 +82,30 @@ impl<S: Socket> MySqlStream<S> {
}

pub(crate) async fn wait_until_ready(&mut self) -> Result<(), Error> {
// Close the statements a dropped operation left open.
for statement in std::mem::take(&mut self.close_pending) {
self.sequence_id = 0;
self.write_packet(StmtClose { statement })?;
}

if !self.socket.write_buffer().is_empty() {
self.socket.flush().await?;
}

while !self.waiting.is_empty() {
while matches!(self.waiting.front(), Some(Waiting::Prepare(_))) {
// The rest of a response a cancelled prepare left behind.
let (_, prepared) = self.recv_packet_tracked().await?;

// Close the statement in the poll that completed it, so a
// cancellation cannot lose the id.
if let Some(statement) = prepared {
self.sequence_id = 0;
self.write_packet(StmtClose { statement })?;
self.socket.flush().await?;
}
}

while self.waiting.front() == Some(&Waiting::Row) {
let packet = self.recv_packet().await?;

Expand Down Expand Up @@ -123,16 +161,37 @@ impl<S: Socket> MySqlStream<S> {
// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_basic_packets.html
// https://mariadb.com/kb/en/library/0-packet/#standard-packet

let mut header: Bytes = self.socket.read(4).await?;
// One `try_read` takes header and payload, so a cancelled read leaves
// the stream at a part boundary. Payloads of 0xFFFFFF bytes or more
// span several parts.
const HEADER_LEN: usize = 4;

let (sequence_id, payload) = self
.socket
.try_read(|buf| {
if buf.len() < HEADER_LEN {
return Ok(ControlFlow::Continue(HEADER_LEN));
}

// cannot overflow
#[allow(clippy::cast_possible_truncation)]
let packet_size = header.get_uint_le(3) as usize;
let sequence_id = header.get_u8();
// cannot overflow
#[allow(clippy::cast_possible_truncation)]
let packet_size = u32::from_le_bytes([buf[0], buf[1], buf[2], 0]) as usize;

self.sequence_id = sequence_id.wrapping_add(1);
let frame_len = HEADER_LEN + packet_size;

if buf.len() < frame_len {
return Ok(ControlFlow::Continue(frame_len));
}

let mut frame = buf.split_to(frame_len);
let sequence_id = frame[3];
let payload = frame.split_off(HEADER_LEN).freeze();

let payload: Bytes = self.socket.read(packet_size).await?;
Ok(ControlFlow::Break((sequence_id, payload)))
})
.await?;

self.sequence_id = sequence_id.wrapping_add(1);

// TODO: packet compression

Expand All @@ -142,6 +201,14 @@ impl<S: Socket> MySqlStream<S> {
// receive the next packet from the database server
// may block (async) on more data from the server
pub(crate) async fn recv_packet(&mut self) -> Result<Packet<Bytes>, Error> {
let (packet, _) = self.recv_packet_tracked().await?;

Ok(packet)
}

/// Like `recv_packet`, plus the statement id of a COM_STMT_PREPARE
/// response this packet completed.
async fn recv_packet_tracked(&mut self) -> Result<(Packet<Bytes>, Option<u32>), Error> {
let payload = self.recv_packet_part().await?;
let payload = if payload.len() < 0xFF_FF_FF {
payload
Expand Down Expand Up @@ -174,7 +241,53 @@ impl<S: Socket> MySqlStream<S> {
);
}

Ok(Packet(payload))
let prepared = self.note_prepare_response_packet(&payload)?;

Ok((Packet(payload), prepared))
}

/// Counts the packets of a pending COM_STMT_PREPARE response and returns
/// its statement id once complete. Only stmt-prepare-ok says how many
/// packets follow, so every read path counts here.
fn note_prepare_response_packet(&mut self, payload: &Bytes) -> Result<Option<u32>, Error> {
let capabilities = self.capabilities;

let Some(Waiting::Prepare(progress)) = self.waiting.front_mut() else {
return Ok(None);
};

let (statement_id, packets_left) = match *progress {
PrepareProgress::Header => {
let ok = PrepareOk::decode_with(payload.clone(), capabilities)?;

let eof_packets = if capabilities.contains(Capabilities::DEPRECATE_EOF) {
0
} else {
u32::from(ok.params > 0) + u32::from(ok.columns > 0)
};

let packets_left = u32::from(ok.params) + u32::from(ok.columns) + eof_packets;

(ok.statement_id, packets_left)
}
PrepareProgress::Tail {
statement_id,
packets_left,
} => (statement_id, packets_left - 1),
};

if packets_left == 0 {
self.waiting.pop_front();

return Ok(Some(statement_id));
}

*progress = PrepareProgress::Tail {
statement_id,
packets_left,
};

Ok(None)
}

pub(crate) async fn recv<'de, T>(&mut self) -> Result<T, Error>
Expand Down Expand Up @@ -215,6 +328,7 @@ impl<S: Socket> MySqlStream<S> {
capabilities: self.capabilities,
sequence_id: self.sequence_id,
waiting: self.waiting,
close_pending: self.close_pending,
is_tls: self.is_tls,
}
}
Expand Down
1 change: 1 addition & 0 deletions sqlx-mysql/src/connection/tls.rs
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ impl WithSocket for MapStream {
capabilities: self.capabilities,
sequence_id: self.sequence_id,
waiting: self.waiting,
close_pending: Vec::new(),
is_tls: true,
}
}
Expand Down
Loading
Loading