diff --git a/src/client.rs b/src/client.rs index 6715f43..d78f43d 100644 --- a/src/client.rs +++ b/src/client.rs @@ -117,10 +117,7 @@ impl Client { let pool = builder.build(ConnectionManager::new(parsed))?; connections.push(pool); } - Ok(Client { - connections, - hash_function: default_hash_function, - }) + Self::with_pools(connections) } pub fn with_pool(pool: Pool) -> Result { @@ -131,6 +128,9 @@ impl Client { } pub fn with_pools(pools: Vec>) -> Result { + if pools.is_empty() { + return Err(MemcacheError::BadURL("No servers specified".to_string())); + } Ok(Client { connections: pools, hash_function: default_hash_function, @@ -631,6 +631,12 @@ mod tests { assert!(client.is_err()); } + #[test] + fn build_client_no_pools() { + assert!(super::Client::with_pools(vec![]).is_err()); + assert!(super::Client::with_pool_size(Vec::::new(), 1).is_err()); + } + #[test] fn build_client_no_url() { let client = super::Client::builder().build(); diff --git a/src/protocol/ascii.rs b/src/protocol/ascii.rs index 6659d93..62326d6 100644 --- a/src/protocol/ascii.rs +++ b/src/protocol/ascii.rs @@ -161,40 +161,18 @@ impl ProtocolTrait for AsciiProtocol { ("get", false) }; - write!(self.reader.get_mut(), "{} {}\r\n", command, key)?; - self.reader.get_mut().flush()?; - - if let Some((k, v)) = self.parse_get_response(has_cas)? { - if k != key { - Err(ServerError::BadResponse(Cow::Borrowed( - "key doesn't match in the response", - )))? - } else if self.parse_get_response::(has_cas)?.is_none() { - Ok(Some(v)) - } else { - Err(ServerError::BadResponse(Cow::Borrowed("Expected end of get response")))? - } - } else { - Ok(None) + let mut values = self.get_values(command, has_cas, &[key])?; + let value = values.remove(key); + if !values.is_empty() { + Err(ServerError::BadResponse(Cow::Borrowed( + "key doesn't match in the response", + )))? } + Ok(value) } fn gets(&mut self, keys: &[&str]) -> Result, MemcacheError> { - write!(self.reader.get_mut(), "gets {}\r\n", keys.join(" "))?; - self.reader.get_mut().flush()?; - - let mut result: HashMap = HashMap::with_capacity(keys.len()); - // there will be atmost keys.len() "VALUE <...>" responses and one END response - for _ in 0..=keys.len() { - match self.parse_get_response(true)? { - Some((key, value)) => { - result.insert(key, value); - } - None => return Ok(result), - } - } - - Err(ServerError::BadResponse(Cow::Borrowed("Expected end of gets response")))? + self.get_values("gets", true, keys) } fn cas>( @@ -420,6 +398,29 @@ impl AsciiProtocol { }) } + fn get_values( + &mut self, + command: &str, + has_cas: bool, + keys: &[&str], + ) -> Result, MemcacheError> { + write!(self.reader.get_mut(), "{} {}\r\n", command, keys.join(" "))?; + self.reader.get_mut().flush()?; + + let mut result: HashMap = HashMap::with_capacity(keys.len()); + // there will be atmost keys.len() "VALUE <...>" responses and one END response + for _ in 0..=keys.len() { + match self.parse_get_response(has_cas)? { + Some((key, value)) => { + result.insert(key, value); + } + None => return Ok(result), + } + } + + Err(ServerError::BadResponse(Cow::Borrowed("Expected end of gets response")))? + } + fn parse_get_response( &mut self, has_cas: bool, @@ -520,4 +521,34 @@ mod tests { ); assert!(capped_line_reader.read_line(|x| Ok(x.to_string())).is_err()); } + + #[test] + fn get_drains_the_response_on_key_mismatch() { + use std::io::BufRead; + use std::net::{TcpListener, TcpStream}; + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr = listener.local_addr().unwrap(); + std::thread::spawn(move || { + let (mut socket, _) = listener.accept().unwrap(); + let mut reader = std::io::BufReader::new(socket.try_clone().unwrap()); + let mut line = String::new(); + reader.read_line(&mut line).unwrap(); + socket.write_all(b"VALUE other 0 1\r\nx\r\nEND\r\n").unwrap(); + line.clear(); + reader.read_line(&mut line).unwrap(); + socket.write_all(b"VALUE foo 0 3\r\nbar\r\nEND\r\n").unwrap(); + }); + + let stream = TcpStream::connect(addr).unwrap(); + stream + .set_read_timeout(Some(std::time::Duration::from_secs(1))) + .unwrap(); + let mut protocol = AsciiProtocol::new(Stream::Tcp(stream)); + assert!(matches!( + protocol.get::("foo"), + Err(MemcacheError::ServerError(ServerError::BadResponse(_))) + )); + assert_eq!(protocol.get::("foo").unwrap(), Some("bar".to_string())); + } }