Skip to content
Merged
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
40 changes: 18 additions & 22 deletions src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -192,9 +192,8 @@ impl Client {
pub fn version(&self) -> Result<Vec<(String, String)>, 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)
}
Expand All @@ -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(());
}
Expand All @@ -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(());
}
Expand All @@ -239,7 +238,7 @@ impl Client {
/// ```
pub fn get<V: FromMemcacheValueExt>(&self, key: &str) -> Result<Option<V>, 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.
Expand Down Expand Up @@ -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);
}
Expand All @@ -284,7 +282,7 @@ impl Client {
/// ```
pub fn set<V: ToMemcacheValue<Stream>>(&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.
Expand All @@ -308,7 +306,7 @@ impl Client {
cas_id: u64,
) -> Result<bool, MemcacheError> {
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.
Expand All @@ -324,7 +322,7 @@ impl Client {
/// ```
pub fn add<V: ToMemcacheValue<Stream>>(&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.
Expand All @@ -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.
Expand All @@ -363,7 +361,7 @@ impl Client {
/// ```
pub fn append<V: ToMemcacheValue<Stream>>(&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.
Expand All @@ -381,7 +379,7 @@ impl Client {
/// ```
pub fn prepend<V: ToMemcacheValue<Stream>>(&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.
Expand All @@ -395,7 +393,7 @@ impl Client {
/// ```
pub fn delete(&self, key: &str) -> Result<bool, MemcacheError> {
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.
Expand All @@ -409,7 +407,7 @@ impl Client {
/// ```
pub fn increment(&self, key: &str, amount: u64) -> Result<u64, MemcacheError> {
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.
Expand All @@ -423,7 +421,7 @@ impl Client {
/// ```
pub fn decrement(&self, key: &str, amount: u64) -> Result<u64, MemcacheError> {
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.
Expand All @@ -439,7 +437,7 @@ impl Client {
/// ```
pub fn touch(&self, key: &str, expiration: u32) -> Result<bool, MemcacheError> {
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.
Expand All @@ -452,9 +450,7 @@ impl Client {
pub fn stats(&self) -> Result<Vec<(String, Stats)>, MemcacheError> {
let mut result: Vec<(String, HashMap<String, String>)> = 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);
Expand Down
37 changes: 32 additions & 5 deletions src/connection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<String>,
broken: bool,
}

impl DerefMut for Connection {
Expand Down Expand Up @@ -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<T>(
pool: &Pool<ConnectionManager>,
f: impl FnOnce(&mut Connection) -> Result<T, MemcacheError>,
) -> Result<T, MemcacheError> {
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),
Expand Down Expand Up @@ -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<Self, MemcacheError> {
let transport = Transport::from_url(url)?;
let is_ascii = url.query_pairs().any(|(ref k, ref v)| k == "protocol" && v == "ascii");
Expand Down Expand Up @@ -254,6 +280,7 @@ impl Connection {
Ok(Connection {
url: Arc::new(url.to_string()),
protocol: protocol,
broken: false,
})
}
}
Expand Down
75 changes: 75 additions & 0 deletions tests/test_broken_connection.rs
Original file line number Diff line number Diff line change
@@ -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::<BigEndian>(0).unwrap();
packet.write_u8(extras.len() as u8).unwrap();
packet.write_u8(0).unwrap();
packet.write_u16::<BigEndian>(0).unwrap();
packet
.write_u32::<BigEndian>((extras.len() + value.len()) as u32)
.unwrap();
packet.write_u32::<BigEndian>(0).unwrap();
packet.write_u64::<BigEndian>(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::<String>("slow").is_err());
thread::sleep(Duration::from_millis(600));

let value: Option<String> = client.get("fast").unwrap();
assert_eq!(value, Some("value-of-fast".into()));
assert_eq!(client.delete("fast").unwrap(), true);
let value: Option<String> = client.get("other").unwrap();
assert_eq!(value, Some("value-of-other".into()));
}
Loading