From 7a5ea42bbd64c99c3f2e36e2a9e13a20644b2c8e Mon Sep 17 00:00:00 2001 From: An Long Date: Tue, 8 Sep 2026 23:04:00 +0900 Subject: [PATCH] =?UTF-8?q?=F0=9F=90=9B=20Drop=20pooled=20connections=20af?= =?UTF-8?q?ter=20IO=20errors=20or=20malformed=20responses?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/client.rs | 40 ++++++++---------- src/connection.rs | 37 +++++++++++++--- tests/test_broken_connection.rs | 75 +++++++++++++++++++++++++++++++++ 3 files changed, 125 insertions(+), 27 deletions(-) create mode 100644 tests/test_broken_connection.rs diff --git a/src/client.rs b/src/client.rs index a06e5b4..c4e5957 100644 --- a/src/client.rs +++ b/src/client.rs @@ -5,7 +5,7 @@ use std::time::Duration; use url::Url; -use crate::connection::ConnectionManager; +use crate::connection::{ConnectionManager, with_connection}; use crate::error::{ClientError, MemcacheError}; use crate::protocol::{Protocol, ProtocolTrait}; use crate::stream::Stream; @@ -192,9 +192,8 @@ impl Client { pub fn version(&self) -> Result, MemcacheError> { let mut result = Vec::with_capacity(self.connections.len()); for connection in self.connections.iter() { - let mut connection = connection.get()?; - let url = connection.get_url(); - result.push((url, connection.version()?)); + let (url, version) = with_connection(connection, |c| Ok((c.get_url(), c.version()?)))?; + result.push((url, version)); } Ok(result) } @@ -209,7 +208,7 @@ impl Client { /// ``` pub fn flush(&self) -> Result<(), MemcacheError> { for connection in self.connections.iter() { - connection.get()?.flush()?; + with_connection(connection, |c| c.flush())?; } return Ok(()); } @@ -224,7 +223,7 @@ impl Client { /// ``` pub fn flush_with_delay(&self, delay: u32) -> Result<(), MemcacheError> { for connection in self.connections.iter() { - connection.get()?.flush_with_delay(delay)?; + with_connection(connection, |c| c.flush_with_delay(delay))?; } return Ok(()); } @@ -239,7 +238,7 @@ impl Client { /// ``` pub fn get(&self, key: &str) -> Result, MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.get(key); + return with_connection(&self.get_connection(key), |c| c.get(key)); } /// Get multiple keys from memcached server. Using this function instead of calling `get` multiple times can reduce network workloads. @@ -267,8 +266,7 @@ impl Client { array.push(key); } for (&connection_index, keys) in con_keys.iter() { - let connection = self.connections[connection_index].clone(); - result.extend(connection.get()?.gets(keys)?); + result.extend(with_connection(&self.connections[connection_index], |c| c.gets(keys))?); } return Ok(result); } @@ -284,7 +282,7 @@ impl Client { /// ``` pub fn set>(&self, key: &str, value: V, expiration: u32) -> Result<(), MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.set(key, value, expiration); + return with_connection(&self.get_connection(key), |c| c.set(key, value, expiration)); } /// Compare and swap a key with the associate value into memcached server with expiration seconds. @@ -308,7 +306,7 @@ impl Client { cas_id: u64, ) -> Result { check_key_len(key)?; - self.get_connection(key).get()?.cas(key, value, expiration, cas_id) + with_connection(&self.get_connection(key), |c| c.cas(key, value, expiration, cas_id)) } /// Add a key with associate value into memcached server with expiration seconds. @@ -324,7 +322,7 @@ impl Client { /// ``` pub fn add>(&self, key: &str, value: V, expiration: u32) -> Result<(), MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.add(key, value, expiration); + return with_connection(&self.get_connection(key), |c| c.add(key, value, expiration)); } /// Replace a key with associate value into memcached server with expiration seconds. @@ -345,7 +343,7 @@ impl Client { expiration: u32, ) -> Result<(), MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.replace(key, value, expiration); + return with_connection(&self.get_connection(key), |c| c.replace(key, value, expiration)); } /// Append value to the key. @@ -363,7 +361,7 @@ impl Client { /// ``` pub fn append>(&self, key: &str, value: V) -> Result<(), MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.append(key, value); + return with_connection(&self.get_connection(key), |c| c.append(key, value)); } /// Prepend value to the key. @@ -381,7 +379,7 @@ impl Client { /// ``` pub fn prepend>(&self, key: &str, value: V) -> Result<(), MemcacheError> { check_key_len(key)?; - return self.get_connection(key).get()?.prepend(key, value); + return with_connection(&self.get_connection(key), |c| c.prepend(key, value)); } /// Delete a key from memcached server. @@ -395,7 +393,7 @@ impl Client { /// ``` pub fn delete(&self, key: &str) -> Result { check_key_len(key)?; - return self.get_connection(key).get()?.delete(key); + return with_connection(&self.get_connection(key), |c| c.delete(key)); } /// Increment the value with amount. @@ -409,7 +407,7 @@ impl Client { /// ``` pub fn increment(&self, key: &str, amount: u64) -> Result { check_key_len(key)?; - return self.get_connection(key).get()?.increment(key, amount); + return with_connection(&self.get_connection(key), |c| c.increment(key, amount)); } /// Decrement the value with amount. @@ -423,7 +421,7 @@ impl Client { /// ``` pub fn decrement(&self, key: &str, amount: u64) -> Result { check_key_len(key)?; - return self.get_connection(key).get()?.decrement(key, amount); + return with_connection(&self.get_connection(key), |c| c.decrement(key, amount)); } /// Set a new expiration time for a exist key. @@ -439,7 +437,7 @@ impl Client { /// ``` pub fn touch(&self, key: &str, expiration: u32) -> Result { check_key_len(key)?; - return self.get_connection(key).get()?.touch(key, expiration); + return with_connection(&self.get_connection(key), |c| c.touch(key, expiration)); } /// Get all servers' statistics. @@ -452,9 +450,7 @@ impl Client { pub fn stats(&self) -> Result, MemcacheError> { let mut result: Vec<(String, HashMap)> = vec![]; for connection in self.connections.iter() { - let mut connection = connection.get()?; - let stats_info = connection.stats()?; - let url = connection.get_url(); + let (url, stats_info) = with_connection(connection, |c| Ok((c.get_url(), c.stats()?)))?; result.push((url, stats_info)); } return Ok(result); diff --git a/src/connection.rs b/src/connection.rs index d8af6e0..bad6f7d 100644 --- a/src/connection.rs +++ b/src/connection.rs @@ -6,19 +6,20 @@ use std::sync::Arc; use std::time::Duration; use url::Url; -use crate::error::MemcacheError; +use crate::error::{MemcacheError, ServerError}; use crate::protocol::{AsciiProtocol, BinaryProtocol, Protocol, ProtocolTrait}; use crate::stream::Stream; use crate::stream::UdpStream; #[cfg(feature = "tls")] use crate::tls::{self, TlsConfig, VerifyMode}; -use r2d2::ManageConnection; +use r2d2::{ManageConnection, Pool}; /// A connection to the memcached server pub struct Connection { pub protocol: Protocol, pub url: Arc, + broken: bool, } impl DerefMut for Connection { @@ -65,12 +66,25 @@ impl ManageConnection for ConnectionManager { conn.version().map(|_| ()) } - fn has_broken(&self, _conn: &mut Self::Connection) -> bool { - // TODO: fix this - false + fn has_broken(&self, conn: &mut Self::Connection) -> bool { + conn.broken } } +/// Run `f` on a connection taken from `pool`, flagging the connection as +/// broken when the error means it can no longer be reused. +pub(crate) fn with_connection( + pool: &Pool, + f: impl FnOnce(&mut Connection) -> Result, +) -> Result { + let mut connection = pool.get()?; + let result = f(&mut connection); + if let Err(err) = &result { + connection.mark_broken_on(err); + } + result +} + enum Transport { Tcp(TcpOptions), Udp(UdpOptions), @@ -225,6 +239,18 @@ impl Connection { self.url.to_string() } + /// Flag the connection so the pool drops it if `err` may have left unread + /// data on the stream, otherwise later commands would read stale responses. + fn mark_broken_on(&mut self, err: &MemcacheError) { + if matches!( + err, + MemcacheError::IOError(_) + | MemcacheError::ServerError(ServerError::BadMagic(_) | ServerError::BadResponse(_)) + ) { + self.broken = true; + } + } + pub(crate) fn connect(url: &Url) -> Result { let transport = Transport::from_url(url)?; let is_ascii = url.query_pairs().any(|(ref k, ref v)| k == "protocol" && v == "ascii"); @@ -254,6 +280,7 @@ impl Connection { Ok(Connection { url: Arc::new(url.to_string()), protocol: protocol, + broken: false, }) } } diff --git a/tests/test_broken_connection.rs b/tests/test_broken_connection.rs new file mode 100644 index 0000000..5e6c840 --- /dev/null +++ b/tests/test_broken_connection.rs @@ -0,0 +1,75 @@ +use std::io::{Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::thread; +use std::time::Duration; + +use byteorder::{BigEndian, ByteOrder, WriteBytesExt}; + +fn write_response(stream: &mut TcpStream, opcode: u8, extras: &[u8], value: &[u8]) { + let mut packet = Vec::new(); + packet.write_u8(0x81).unwrap(); + packet.write_u8(opcode).unwrap(); + packet.write_u16::(0).unwrap(); + packet.write_u8(extras.len() as u8).unwrap(); + packet.write_u8(0).unwrap(); + packet.write_u16::(0).unwrap(); + packet + .write_u32::((extras.len() + value.len()) as u32) + .unwrap(); + packet.write_u32::(0).unwrap(); + packet.write_u64::(0).unwrap(); + packet.extend_from_slice(extras); + packet.extend_from_slice(value); + stream.write_all(&packet).unwrap(); +} + +fn serve(mut stream: TcpStream) { + let mut header = [0u8; 24]; + while stream.read_exact(&mut header).is_ok() { + let opcode = header[1]; + let key_length = BigEndian::read_u16(&header[2..]) as usize; + let extras_length = header[4] as usize; + let body_length = BigEndian::read_u32(&header[8..]) as usize; + let mut body = vec![0u8; body_length]; + stream.read_exact(&mut body).unwrap(); + let key = &body[extras_length..extras_length + key_length]; + match opcode { + 0x0b => write_response(&mut stream, opcode, &[], b"1.6.45"), + 0x00 => { + if key.starts_with(b"slow") { + thread::sleep(Duration::from_millis(500)); + } + let value = [b"value-of-", key].concat(); + write_response(&mut stream, opcode, &[0, 0, 0, 0], &value); + } + _ => write_response(&mut stream, opcode, &[], &[]), + } + } +} + +#[test] +fn test_connection_dropped_after_read_timeout() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + thread::spawn(move || { + for stream in listener.incoming() { + thread::spawn(move || serve(stream.unwrap())); + } + }); + + let client = memcache::Client::builder() + .add_server(format!("memcache://127.0.0.1:{}", port)) + .unwrap() + .with_read_timeout(Duration::from_millis(100)) + .build() + .unwrap(); + + assert!(client.get::("slow").is_err()); + thread::sleep(Duration::from_millis(600)); + + let value: Option = client.get("fast").unwrap(); + assert_eq!(value, Some("value-of-fast".into())); + assert_eq!(client.delete("fast").unwrap(), true); + let value: Option = client.get("other").unwrap(); + assert_eq!(value, Some("value-of-other".into())); +}