use packet::{Column, InnerStmt};
use std::borrow::Borrow;
use std::cmp;
use std::collections::HashMap;
use std::fs;
use std::fmt;
use std::hash::BuildHasherDefault as BldHshrDflt;
use std::io;
use std::io::Read;
use std::io::Write as NewWrite;
use std::mem;
use std::ops::{
Deref,
DerefMut,
Index,
};
use std::path;
use std::str::from_utf8;
use std::sync::{Arc, Mutex};
use super::consts;
use super::consts::Command;
use super::consts::ColumnType;
use super::io::Read as MyRead;
use super::io::Write;
use super::io::Stream;
use super::error::Error::{
IoError,
MySqlError,
DriverError
};
use super::error::DriverError::{
CouldNotConnect,
UnsupportedProtocol,
Protocol41NotSet,
UnexpectedPacket,
MismatchedStmtParams,
NamedParamsForPositionalQuery,
SetupError,
ReadOnlyTransNotSupported,
};
use super::error::Result as MyResult;
use super::error::DriverError::SslNotSupported;
use named_params::parse_named_params;
use super::scramble::scramble;
use super::packet::{OkPacket, EOFPacket, ErrPacket, HandshakePacket, ServerVersion};
use super::value::{
Params,
Value,
from_value_opt,
from_value,
FromValue,
};
use super::value::Value::{NULL, Int, UInt, Float, Bytes, Date, Time};
use byteorder::LittleEndian as LE;
use byteorder::WriteBytesExt;
use fnv::FnvHasher;
pub mod pool;
mod opts;
mod stmt_cache;
use self::stmt_cache::StmtCache;
pub use self::opts::Opts;
pub use self::opts::OptsBuilder;
#[cfg(feature = "ssl")]
pub use self::opts::SslOpts;
pub trait GenericConnection {
fn query<T: AsRef<str>>(&mut self, query: T) -> MyResult<QueryResult>;
fn first<T: AsRef<str>>(&mut self, query: T) -> MyResult<Option<Row>>;
fn prepare<T: AsRef<str>>(&mut self, query: T) -> MyResult<Stmt>;
fn prep_exec<A, T>(&mut self, query: A, params: T) -> MyResult<QueryResult>
where A: AsRef<str>, T: Into<Params>;
fn first_exec<Q, P>(&mut self, query: Q, params: P) -> MyResult<Option<Row>>
where Q: AsRef<str>, P: Into<Params>;
}
#[derive(PartialEq, Eq, Clone, Copy, Debug)]
pub enum IsolationLevel {
ReadUncommitted,
ReadCommitted,
RepeatableRead,
Serializable,
}
impl fmt::Display for IsolationLevel {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match *self {
IsolationLevel::ReadUncommitted => write!(f, "READ UNCOMMITTED"),
IsolationLevel::ReadCommitted => write!(f, "READ COMMITTED"),
IsolationLevel::RepeatableRead => write!(f, "REPEATABLE READ"),
IsolationLevel::Serializable => write!(f, "SERIALIZABLE"),
}
}
}
#[derive(Debug)]
pub struct Transaction<'a> {
conn: ConnRef<'a>,
committed: bool,
rolled_back: bool,
restore_local_infile_handler: Option<LocalInfileHandler>,
}
impl<'a> Transaction<'a> {
fn new(conn: &'a mut Conn) -> Transaction<'a> {
let handler = conn.local_infile_handler.clone();
Transaction {
conn: ConnRef::ViaConnRef(conn),
committed: false,
rolled_back: false,
restore_local_infile_handler: handler,
}
}
fn new_pooled(conn: pool::PooledConn) -> Transaction<'a> {
let handler = conn.as_ref().local_infile_handler.clone();
Transaction {
conn: ConnRef::ViaPooledConn(conn),
committed: false,
rolled_back: false,
restore_local_infile_handler: handler,
}
}
pub fn query<T: AsRef<str>>(&mut self, query: T) -> MyResult<QueryResult> {
self.conn.query(query)
}
pub fn first<T: AsRef<str>>(&mut self, query: T) -> MyResult<Option<Row>> {
self.query(query).and_then(|result| {
for row in result {
return row.map(Some);
}
return Ok(None)
})
}
pub fn prepare<T: AsRef<str>>(&mut self, query: T) -> MyResult<Stmt> {
self.conn.prepare(query)
}
pub fn prep_exec<A: AsRef<str>, T: Into<Params>>(&mut self, query: A, params: T) -> MyResult<QueryResult> {
self.conn.prep_exec(query, params)
}
pub fn first_exec<Q, P>(&mut self, query: Q, params: P) -> MyResult<Option<Row>>
where Q: AsRef<str>,
P: Into<Params>,
{
self.prep_exec(query, params).and_then(|result| {
for row in result {
return row.map(Some);
}
return Ok(None)
})
}
pub fn commit(mut self) -> MyResult<()> {
self.conn.query("COMMIT")?;
self.committed = true;
Ok(())
}
pub fn rollback(mut self) -> MyResult<()> {
self.conn.query("ROLLBACK")?;
self.rolled_back = true;
Ok(())
}
pub fn set_local_infile_handler(&mut self, handler: Option<LocalInfileHandler>) {
self.conn.set_local_infile_handler(handler);
}
}
impl<'a> GenericConnection for Transaction<'a> {
fn query<T: AsRef<str>>(&mut self, query: T) -> MyResult<QueryResult> {
self.query(query)
}
fn first<T: AsRef<str>>(&mut self, query: T) -> MyResult<Option<Row>> {
self.first(query)
}
fn prepare<T: AsRef<str>>(&mut self, query: T) -> MyResult<Stmt> {
self.prepare(query)
}
fn prep_exec<A, T>(&mut self, query: A, params: T) -> MyResult<QueryResult>
where A: AsRef<str>, T: Into<Params> {
self.prep_exec(query, params)
}
fn first_exec<Q, P>(&mut self, query: Q, params: P) -> MyResult<Option<Row>>
where Q: AsRef<str>, P: Into<Params> {
self.first_exec(query, params)
}
}
impl<'a> Drop for Transaction<'a> {
fn drop(&mut self) {
if ! self.committed && ! self.rolled_back {
let _ = self.conn.query("ROLLBACK");
}
self.conn.local_infile_handler = self.restore_local_infile_handler.take();
}
}
#[derive(Debug)]
enum ConnRef<'a> {
ViaConnRef(&'a mut Conn),
ViaPooledConn(pool::PooledConn),
}
impl<'a> Deref for ConnRef<'a> {
type Target = Conn;
fn deref<'c>(&'c self) -> &'c Conn {
match *self {
ConnRef::ViaConnRef(ref conn_ref) => conn_ref,
ConnRef::ViaPooledConn(ref conn) => conn.as_ref(),
}
}
}
impl<'a> DerefMut for ConnRef<'a> {
fn deref_mut<'c>(&'c mut self) -> &'c mut Conn {
match *self {
ConnRef::ViaConnRef(ref mut conn_ref) => conn_ref,
ConnRef::ViaPooledConn(ref mut conn) => conn.as_mut(),
}
}
}
#[derive(Debug)]
pub struct Stmt<'a> {
stmt: InnerStmt,
conn: ConnRef<'a>,
}
impl<'a> Stmt<'a> {
fn new(stmt: InnerStmt, conn: &'a mut Conn) -> Stmt<'a> {
Stmt {
stmt: stmt,
conn: ConnRef::ViaConnRef(conn),
}
}
fn new_pooled(stmt: InnerStmt, pooled_conn: pool::PooledConn) -> Stmt<'a> {
Stmt {
stmt: stmt,
conn: ConnRef::ViaPooledConn(pooled_conn),
}
}
pub fn params_ref(&self) -> Option<&[Column]> {
self.stmt.params()
}
pub fn columns_ref(&self) -> Option<&[Column]> {
self.stmt.columns()
}
pub fn column_index<T: AsRef<str>>(&self, name: T) -> Option<usize> {
match self.stmt.columns() {
None => None,
Some(columns) => {
let name = name.as_ref().as_bytes();
for (i, c) in columns.iter().enumerate() {
if c.name() == name {
return Some(i)
}
}
None
}
}
}
pub fn execute<'s, T: Into<Params>>(&'s mut self, params: T) -> MyResult<QueryResult<'s>> {
self.conn.execute(&self.stmt, params)
}
pub fn first_exec<P>(&mut self, params: P) -> MyResult<Option<Row>>
where P: Into<Params>,
{
self.execute(params).and_then(|result| {
for row in result {
return row.map(Some);
}
return Ok(None)
})
}
fn prep_exec<T: Into<Params>>(mut self, params: T) -> MyResult<QueryResult<'a>> {
let (columns, ok_packet) = self.conn._execute(&self.stmt, params.into())?;
Ok(QueryResult::new(ResultConnRef::ViaStmt(self), columns, ok_packet, true))
}
}
impl<'a> Drop for Stmt<'a> {
fn drop(&mut self) {
if self.conn.stmt_cache.get_cap() == 0 {
let mut stmt_id = [0u8; 4];
let _ = (&mut stmt_id[..]).write_u32::<LE>(self.stmt.id());
let _ = self.conn.write_command_data(Command::COM_STMT_CLOSE, &stmt_id[..]);
}
}
}
#[derive(Clone, PartialEq)]
pub struct Row {
values: Vec<Option<Value>>,
columns: Arc<Vec<Column>>
}
impl fmt::Debug for Row {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut debug = f.debug_tuple("Row");
for val in self.values.iter() {
match *val {
Some(ref val) => {
debug.field(val);
},
None => {
debug.field(&"<taken>");
},
}
}
debug.finish()
}
}
impl Row {
#[doc(hidden)]
pub fn new(raw_row: Vec<Value>, columns: Arc<Vec<Column>>) -> Row {
Row {
values: raw_row.into_iter().map(|value| Some(value)).collect(),
columns: columns
}
}
pub fn len(&self) -> usize {
self.values.len()
}
pub fn as_ref(&self, index: usize) -> Option<&Value> {
self.values.get(index).and_then(|x| x.as_ref())
}
pub fn get<T, I>(&mut self, index: I) -> Option<T>
where T: FromValue,
I: ColumnIndex {
index.idx(&*self.columns).and_then(|idx| {
self.values.get(idx).and_then(|x| x.as_ref()).map(|x| from_value::<T>(x.clone()))
})
}
pub fn take<T, I>(&mut self, index: I) -> Option<T>
where T: FromValue,
I: ColumnIndex {
index.idx(&*self.columns).and_then(|idx| {
self.values.get_mut(idx).and_then(|x| x.take()).map(from_value::<T>)
})
}
pub fn unwrap(self) -> Vec<Value> {
self.values.into_iter()
.map(|x| x.expect("Can't unwrap row if some of columns was taken"))
.collect()
}
#[doc(hidden)]
pub fn place(&mut self, index: usize, value: Value) {
self.values[index] = Some(value);
}
}
impl Index<usize> for Row {
type Output = Value;
fn index<'a>(&'a self, index: usize) -> &'a Value {
self.values[index].as_ref().unwrap()
}
}
impl<'a> Index<&'a str> for Row {
type Output = Value;
fn index<'r>(&'r self, index: &'a str) -> &'r Value {
for (i, column) in self.columns.iter().enumerate() {
if column.name() == index.as_bytes() {
return self.values[i].as_ref().unwrap();
}
}
panic!("No such column: `{}`", index);
}
}
pub trait ColumnIndex {
fn idx(&self, columns: &Vec<Column>) -> Option<usize>;
}
impl ColumnIndex for usize {
fn idx(&self, columns: &Vec<Column>) -> Option<usize> {
if *self >= columns.len() {
None
} else {
Some(*self)
}
}
}
impl<'a> ColumnIndex for &'a str {
fn idx(&self, columns: &Vec<Column>) -> Option<usize> {
for (i, c) in columns.iter().enumerate() {
if c.name() == self.as_bytes() {
return Some(i);
}
}
None
}
}
#[derive(Clone)]
pub struct LocalInfileHandler(
Arc<Mutex<
for<'a> FnMut(&'a [u8], &'a mut LocalInfile) -> io::Result<()> + Send
>>
);
impl LocalInfileHandler {
pub fn new<F>(f: F) -> Self
where F: for<'a> FnMut(&'a [u8], &'a mut LocalInfile) -> io::Result<()> + Send + 'static
{
LocalInfileHandler(Arc::new(Mutex::new(f)))
}
}
impl PartialEq for LocalInfileHandler {
fn eq(&self, other: &LocalInfileHandler) -> bool {
(&*self.0 as *const _) == (&*other.0 as *const _)
}
}
impl Eq for LocalInfileHandler {}
impl fmt::Debug for LocalInfileHandler {
fn fmt(&self, f: &mut fmt::Formatter) -> Result<(), fmt::Error> {
write!(f, "LocalInfileHandler(...)")
}
}
#[derive(Debug)]
pub struct LocalInfile<'a> {
buffer: io::Cursor<Box<[u8]>>,
conn: &'a mut Conn,
}
impl<'a> io::Write for LocalInfile<'a> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let result = self.buffer.write(buf);
if let Ok(_) = result {
if self.buffer.position() as usize >= self.buffer.get_ref().len() {
self.flush()?;
}
}
result
}
fn flush(&mut self) -> io::Result<()> {
let n = self.buffer.position() as usize;
if n > 0 {
let range = &self.buffer.get_ref()[..n];
self.conn.write_packet(range).map_err(|e| {
io::Error::new(io::ErrorKind::Other, Box::new(e))
})?;
}
self.buffer.set_position(0);
Ok(())
}
}
#[derive(Debug)]
pub struct Conn {
opts: Opts,
stream: Option<Stream>,
stmt_cache: StmtCache,
server_version: ServerVersion,
affected_rows: u64,
last_insert_id: u64,
max_allowed_packet: usize,
capability_flags: consts::CapabilityFlags,
connection_id: u32,
status_flags: consts::StatusFlags,
seq_id: u8,
character_set: u8,
last_command: u8,
connected: bool,
has_results: bool,
local_infile_handler: Option<LocalInfileHandler>,
}
impl Conn {
fn empty<T: Into<Opts>>(opts: T) -> Conn {
let opts = opts.into();
Conn {
stmt_cache: StmtCache::new(opts.get_stmt_cache_size()),
opts: opts,
stream: None,
seq_id: 0u8,
capability_flags: consts::CapabilityFlags::empty(),
status_flags: consts::StatusFlags::empty(),
connection_id: 0u32,
character_set: 0u8,
affected_rows: 0u64,
last_insert_id: 0u64,
last_command: 0u8,
max_allowed_packet: consts::MAX_PAYLOAD_LEN,
connected: false,
has_results: false,
server_version: (0, 0, 0),
local_infile_handler: None
}
}
fn can_improved(&mut self) -> Option<Opts> {
if self.opts.get_prefer_socket() && self.opts.addr_is_loopback() {
if let Some(socket) = self.get_system_var("socket") {
if self.opts.get_socket().is_none() {
let mut socket_opts = OptsBuilder::from_opts(self.opts.clone());
let socket = from_value::<String>(socket);
socket_opts.socket(Some(socket));
return Some(socket_opts.into())
}
}
}
None
}
pub fn new<T: Into<Opts>>(opts: T) -> MyResult<Conn> {
let mut conn = Conn::empty(opts);
conn.connect_stream()?;
conn.connect()?;
let mut conn = {
if let Some(new_opts) = conn.can_improved() {
drop(conn);
let mut improved_conn = Conn::empty(new_opts);
improved_conn.connect_stream()?;
improved_conn.connect()?;
improved_conn
} else {
conn
}
};
for cmd in conn.opts.get_init().clone() {
conn.query(cmd)?;
}
return Ok(conn);
}
fn soft_reset(&mut self) -> MyResult<()> {
self.write_command(Command::COM_RESET_CONNECTION)?;
self.read_packet().and_then(|pld| {
match pld[0] {
0 => {
let ok = OkPacket::from_payload(&*pld)?;
self.handle_ok(&ok);
self.last_command = 0;
self.stmt_cache.clear();
Ok(())
},
_ => {
let err = ErrPacket::from_payload(&*pld, self.capability_flags)?;
Err(MySqlError(err.into()))
},
}
})
}
fn hard_reset(&mut self) -> MyResult<()> {
self.stream = None;
self.stmt_cache.clear();
self.seq_id = 0;
self.capability_flags = consts::CapabilityFlags::empty();
self.status_flags = consts::StatusFlags::empty();
self.connection_id = 0;
self.character_set = 0;
self.affected_rows = 0;
self.last_insert_id = 0;
self.last_command = 0;
self.max_allowed_packet = consts::MAX_PAYLOAD_LEN;
self.connected = false;
self.has_results = false;
self.connect_stream()?;
self.connect()
}
pub fn reset(&mut self) -> MyResult<()> {
if self.server_version > (5, 7, 2) {
match self.soft_reset() {
Ok(_) => Ok(()),
_ => self.hard_reset()
}
} else {
self.hard_reset()
}
}
fn get_mut_stream<'a>(&'a mut self) -> &'a mut Stream {
self.stream.as_mut().unwrap()
}
#[cfg(all(feature = "ssl", any(unix, target_os = "macos")))]
fn switch_to_ssl(&mut self) -> MyResult<()> {
if self.stream.is_some() {
let stream = self.stream.take().unwrap();
let stream = stream.make_secure(self.opts.get_verify_peer(),
self.opts.get_ip_or_hostname(),
self.opts.get_ssl_opts())?;
self.stream = Some(stream);
}
Ok(())
}
#[cfg(any(not(feature = "ssl"), target_os = "windows"))]
fn switch_to_ssl(&mut self) -> MyResult<()> {
unimplemented!();
}
fn connect_stream(&mut self) -> MyResult<()> {
let read_timeout = self.opts.get_read_timeout().clone();
let write_timeout = self.opts.get_write_timeout().clone();
let tcp_keepalive_time = self.opts.get_tcp_keepalive_time_ms().clone();
let tcp_connect_timeout = self.opts.get_tcp_connect_timeout();
let bind_address = self.opts.bind_address().cloned();
let stream = if let Some(ref socket) = *self.opts.get_socket() {
Stream::connect_socket(&*socket, read_timeout, write_timeout)?
} else if let Some(ref ip_or_hostname) = *self.opts.get_ip_or_hostname() {
let port = self.opts.get_tcp_port();
Stream::connect_tcp(&*ip_or_hostname,
port,
read_timeout,
write_timeout,
tcp_keepalive_time,
tcp_connect_timeout,
bind_address)?
} else {
return Err(DriverError(CouldNotConnect(None)));
};
self.stream = Some(stream);
return Ok(())
}
fn read_packet(&mut self) -> MyResult<Vec<u8>> {
let old_seq_id = self.seq_id;
let (data, seq_id) = self.get_mut_stream().as_mut().read_packet(old_seq_id)?;
self.seq_id = seq_id;
Ok(data)
}
fn drop_packet(&mut self) -> MyResult<()> {
let old_seq_id = self.seq_id;
let seq_id = self.get_mut_stream().as_mut().drop_packet(old_seq_id)?;
self.seq_id = seq_id;
Ok(())
}
fn write_packet(&mut self, data: &[u8]) -> MyResult<()> {
let seq_id = self.seq_id;
let max_allowed_packet = self.max_allowed_packet;
self.seq_id = self.get_mut_stream().as_mut().write_packet(data, seq_id, max_allowed_packet)?;
Ok(())
}
fn handle_handshake(&mut self, hp: &HandshakePacket) {
self.capability_flags = hp.capability_flags;
self.status_flags = hp.status_flags;
self.connection_id = hp.connection_id;
self.character_set = hp.character_set;
self.server_version = hp.server_version;
}
fn handle_ok(&mut self, op: &OkPacket) {
self.affected_rows = op.affected_rows;
self.last_insert_id = op.last_insert_id;
self.status_flags = op.status_flags;
}
fn handle_eof(&mut self, eof: &EOFPacket) {
self.status_flags = eof.status_flags;
}
fn do_handshake(&mut self) -> MyResult<()> {
self.read_packet().and_then(|pld| {
match pld[0] {
0xFF => {
let error_packet = ErrPacket::from_payload(pld.as_ref(),
self.capability_flags)?;
Err(MySqlError(error_packet.into()))
},
_ => {
let handshake = HandshakePacket::from_payload(pld.as_ref())?;
if handshake.protocol_version != 10u8 {
return Err(DriverError(UnsupportedProtocol(handshake.protocol_version)));
}
if !handshake.capability_flags.contains(consts::CLIENT_PROTOCOL_41) {
return Err(DriverError(Protocol41NotSet));
}
self.handle_handshake(&handshake);
if self.opts.get_ssl_opts().is_some() && self.stream.is_some() {
if self.stream.as_ref().unwrap().is_insecure() {
if !handshake.capability_flags.contains(consts::CLIENT_SSL) {
return Err(DriverError(SslNotSupported));
} else {
self.do_ssl_request(&handshake)?;
self.switch_to_ssl()?;
}
}
}
self.do_handshake_response(&handshake)
},
}
}).and_then(|_| {
self.read_packet()
}).and_then(|pld| {
match pld[0] {
0u8 => {
let ok = OkPacket::from_payload(pld.as_ref())?;
self.handle_ok(&ok);
Ok(())
},
0xffu8 => {
let err = ErrPacket::from_payload(pld.as_ref(),
self.capability_flags)?;
Err(MySqlError(err.into()))
},
_ => Err(DriverError(UnexpectedPacket))
}
})
}
fn get_client_flags(&self) -> consts::CapabilityFlags {
let mut client_flags = consts::CLIENT_PROTOCOL_41 |
consts::CLIENT_SECURE_CONNECTION |
consts::CLIENT_LONG_PASSWORD |
consts::CLIENT_TRANSACTIONS |
consts::CLIENT_LOCAL_FILES |
consts::CLIENT_MULTI_STATEMENTS |
consts::CLIENT_MULTI_RESULTS |
consts::CLIENT_PS_MULTI_RESULTS |
(self.capability_flags & consts::CLIENT_LONG_FLAG);
if let &Some(ref db_name) = self.opts.get_db_name() {
if db_name.len() > 0 {
client_flags.insert(consts::CLIENT_CONNECT_WITH_DB);
}
}
if self.stream.is_some() && self.stream.as_ref().unwrap().is_insecure() {
if self.opts.get_ssl_opts().is_some() {
client_flags.insert(consts::CLIENT_SSL);
}
}
client_flags
}
fn do_ssl_request(&mut self, hp: &HandshakePacket) -> MyResult<()> {
let client_flags = self.get_client_flags();
let mut buf = [0; 4 + 4 + 1 + 23];
{
let mut writer = &mut buf[..];
writer.write_u32::<LE>(client_flags.bits())?;
writer.write_all(&[0u8; 4])?;
writer.write_u8(hp.get_default_collation())?;
writer.write_all(&[0u8; 23])?;
}
self.write_packet(&buf[..])
}
fn do_handshake_response(&mut self, hp: &HandshakePacket) -> MyResult<()> {
let client_flags = self.get_client_flags();
let scramble_buf = if let &Some(ref pass) = self.opts.get_pass() {
scramble(&*hp.auth_plugin_data, pass.as_bytes())
} else {
None
};
let user_len = self.opts.get_user().as_ref().map(|x| x.as_bytes().len()).unwrap_or(0);
let db_name_len = self.opts.get_db_name().as_ref().map(|x| x.as_bytes().len()).unwrap_or(0);
let scramble_buf_len = if scramble_buf.is_some() { 20 } else { 0 };
let mut payload_len = 4 + 4 + 1 + 23 + user_len + 1 + 1 + scramble_buf_len;
if db_name_len > 0 {
payload_len += db_name_len + 1;
}
let mut buf = vec![0u8; payload_len];
{
let mut writer = &mut *buf;
writer.write_u32::<LE>(client_flags.bits())?;
writer.write_all(&[0u8; 4])?;
writer.write_u8(hp.get_default_collation())?;
writer.write_all(&[0u8; 23])?;
if let &Some(ref user) = self.opts.get_user() {
writer.write_all(user.as_bytes())?;
}
writer.write_u8(0u8)?;
writer.write_u8(scramble_buf_len as u8)?;
if let Some(scr) = scramble_buf {
writer.write_all(scr.as_ref())?;
}
if db_name_len > 0 {
let db_name = self.opts.get_db_name().as_ref().unwrap();
writer.write_all(db_name.as_bytes())?;
writer.write_u8(0u8)?;
}
}
self.write_packet(&*buf)
}
fn write_command(&mut self, cmd: consts::Command) -> MyResult<()> {
self.seq_id = 0u8;
self.last_command = cmd as u8;
self.write_packet(&[cmd as u8])
}
fn write_command_data(&mut self, cmd: consts::Command, data: &[u8]) -> MyResult<()> {
self.seq_id = 0u8;
self.last_command = cmd as u8;
let mut buf = vec![0u8; data.len() + 1];
{
let mut writer = &mut *buf;
let _ = writer.write_u8(cmd as u8);
let _ = writer.write_all(data);
}
self.write_packet(&*buf)
}
fn send_long_data(&mut self, stmt: &InnerStmt, params: &[Value], ids: Vec<u16>) -> MyResult<()> {
for &id in ids.iter() {
match params[id as usize] {
Bytes(ref x) => {
for chunk in x.chunks(self.max_allowed_packet - 7) {
let chunk_len = chunk.len() + 7;
let mut buf = vec![0u8; chunk_len];
{
let mut writer = &mut *buf;
writer.write_u32::<LE>(stmt.id())?;
writer.write_u16::<LE>(id)?;
writer.write_all(chunk)?;
}
self.write_command_data(Command::COM_STMT_SEND_LONG_DATA, &*buf)?;
}
},
_ => unreachable!(),
}
}
Ok(())
}
fn _execute(&mut self, stmt: &InnerStmt, params: Params) -> MyResult<(Vec<Column>, Option<OkPacket>)> {
let mut buf = [0u8; 4 + 1 + 4];
let mut data: Vec<u8>;
let out;
match params {
Params::Empty => {
if stmt.num_params() != 0 {
return Err(DriverError(MismatchedStmtParams(stmt.num_params(), 0)));
}
{
let mut writer = &mut buf[..];
writer.write_u32::<LE>(stmt.id())?;
writer.write_u8(0u8)?;
writer.write_u32::<LE>(1u32)?;
}
out = &buf[..];
},
Params::Positional(params) => {
if stmt.num_params() != params.len() as u16 {
return Err(DriverError(MismatchedStmtParams(stmt.num_params(), params.len())));
}
if let Some(sparams) = stmt.params() {
let (bitmap, values, large_ids) =
Value::to_bin_payload(sparams,
¶ms,
self.max_allowed_packet)?;
match large_ids {
Some(ids) => self.send_long_data(stmt, ¶ms, ids)?,
_ => ()
}
data = vec![0u8; 9 + bitmap.len() + 1 + params.len() * 2 + values.len()];
{
let mut writer = &mut *data;
writer.write_u32::<LE>(stmt.id())?;
writer.write_u8(0u8)?;
writer.write_u32::<LE>(1u32)?;
writer.write_all(bitmap.as_ref())?;
writer.write_u8(1u8)?;
for i in 0..params.len() {
match params[i] {
NULL => writer.write_all(&[sparams[i].column_type as u8, 0u8])?,
Bytes(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_VAR_STRING as u8, 0u8])?,
Int(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_LONGLONG as u8, 0u8])?,
UInt(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_LONGLONG as u8, 128u8])?,
Float(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_DOUBLE as u8, 0u8])?,
Date(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_DATETIME as u8, 0u8])?,
Time(..) =>
writer.write_all(&[ColumnType::MYSQL_TYPE_TIME as u8, 0u8])?
}
}
writer.write_all(values.as_ref())?;
}
out = &*data;
} else {
unreachable!();
}
},
Params::Named(_) => {
if let None = stmt.named_params() {
return Err(DriverError(NamedParamsForPositionalQuery));
}
let named_params = stmt.named_params().unwrap();
return self._execute(stmt, params.into_positional(named_params)?)
}
}
self.write_command_data(Command::COM_STMT_EXECUTE, out)?;
self.handle_result_set()
}
fn execute<'a, T: Into<Params>>(&'a mut self, stmt: &InnerStmt, params: T) -> MyResult<QueryResult<'a>> {
match self._execute(stmt, params.into()) {
Ok((columns, ok_packet)) => {
Ok(QueryResult::new(ResultConnRef::ViaConnRef(self), columns, ok_packet, true))
},
Err(err) => Err(err)
}
}
fn _start_transaction(&mut self,
consistent_snapshot: bool,
isolation_level: Option<IsolationLevel>,
readonly: Option<bool>) -> MyResult<()> {
if let Some(i_level) = isolation_level {
let _ = self.query(format!("SET TRANSACTION ISOLATION LEVEL {}", i_level))?;
}
if let Some(readonly) = readonly {
if self.server_version < (5, 6, 5) {
return Err(DriverError(ReadOnlyTransNotSupported));
}
let _ = if readonly {
self.query("SET TRANSACTION READ ONLY")?
} else {
self.query("SET TRANSACTION READ WRITE")?
};
}
let _ = if consistent_snapshot {
self.query("START TRANSACTION WITH CONSISTENT SNAPSHOT")?
} else {
self.query("START TRANSACTION")?
};
Ok(())
}
fn send_local_infile(&mut self, file_name: &[u8]) -> MyResult<Option<OkPacket>> {
{
let buffer_size = cmp::min(consts::MAX_PAYLOAD_LEN - 1, self.max_allowed_packet - 1);
let chunk = vec![0u8; buffer_size].into_boxed_slice();
let maybe_handler = self.local_infile_handler.clone().or_else(|| {
self.opts.get_local_infile_handler().clone()
});
let mut local_infile = LocalInfile {
buffer: io::Cursor::new(chunk),
conn: self
};
if let Some(handler) = maybe_handler {
let mut handler_fn = &mut *handler.0.lock().unwrap();
handler_fn(file_name, &mut local_infile)?;
} else {
let path = String::from_utf8_lossy(file_name);
let path = path.into_owned();
let path: path::PathBuf = path.into();
let mut file = fs::File::open(&path)?;
io::copy(&mut file, &mut local_infile)?;
};
local_infile.flush()?;
}
self.write_packet(&[])?;
let pld = self.read_packet()?;
if pld[0] == 0u8 {
let ok = OkPacket::from_payload(pld.as_ref())?;
self.handle_ok(&ok);
return Ok(Some(ok));
}
Ok(None)
}
fn handle_result_set(&mut self) -> MyResult<(Vec<Column>, Option<OkPacket>)> {
let pld = self.read_packet()?;
match pld[0] {
0x00 => {
let ok = OkPacket::from_payload(pld.as_ref())?;
self.handle_ok(&ok);
Ok((Vec::new(), Some(ok)))
},
0xfb => {
let mut reader = &pld[1..];
let mut file_name = Vec::with_capacity(reader.len());
reader.read_to_end(&mut file_name)?;
match self.send_local_infile(file_name.as_ref()) {
Ok(x) => Ok((Vec::new(), x)),
Err(err) => Err(err)
}
},
0xff => {
let err = ErrPacket::from_payload(pld.as_ref(), self.capability_flags)?;
Err(MySqlError(err.into()))
},
_ => {
let mut reader = &pld[..];
let column_count = reader.read_lenenc_int()?;
let mut columns: Vec<Column> = Vec::with_capacity(column_count as usize);
for _ in 0..column_count {
let pld = self.read_packet()?;
columns.push(Column::from_payload(pld)?);
}
self.read_packet()?;
self.has_results = true;
Ok((columns, None))
}
}
}
fn _query(&mut self, query: &str) -> MyResult<(Vec<Column>, Option<OkPacket>)> {
self.write_command_data(Command::COM_QUERY, query.as_bytes())?;
self.handle_result_set()
}
pub fn ping(&mut self) -> bool {
match self.write_command(Command::COM_PING) {
Ok(_) => {
self.drop_packet().is_ok()
},
_ => false
}
}
pub fn start_transaction<'a>(&'a mut self,
consistent_snapshot: bool,
isolation_level: Option<IsolationLevel>,
readonly: Option<bool>) -> MyResult<Transaction<'a>> {
let _ = self._start_transaction(consistent_snapshot, isolation_level, readonly)?;
Ok(Transaction::new(self))
}
pub fn query<T: AsRef<str>>(&mut self, query: T) -> MyResult<QueryResult> {
match self._query(query.as_ref()) {
Ok((columns, ok_packet)) => {
Ok(QueryResult::new(ResultConnRef::ViaConnRef(self), columns, ok_packet, false))
},
Err(err) => Err(err),
}
}
pub fn first<T: AsRef<str>>(&mut self, query: T) -> MyResult<Option<Row>> {
self.query(query).and_then(|result| {
for row in result {
return row.map(Some);
}
return Ok(None)
})
}
fn _true_prepare(&mut self,
query: &str,
named_params: Option<Vec<String>>) -> MyResult<InnerStmt> {
self.write_command_data(Command::COM_STMT_PREPARE, query.as_bytes())?;
let pld = self.read_packet()?;
match pld[0] {
0xff => {
let err = ErrPacket::from_payload(pld.as_ref(), self.capability_flags)?;
Err(MySqlError(err.into()))
},
_ => {
let mut stmt = InnerStmt::from_payload(pld.as_ref(), named_params)?;
if stmt.num_params() > 0 {
let mut params: Vec<Column> = Vec::with_capacity(stmt.num_params() as usize);
for _ in 0..stmt.num_params() {
let pld = self.read_packet()?;
params.push(Column::from_payload(pld)?);
}
stmt.set_params(Some(params));
self.read_packet()?;
}
if stmt.num_columns() > 0 {
let mut columns: Vec<Column> = Vec::with_capacity(stmt.num_columns() as usize);
for _ in 0..stmt.num_columns() {
let pld = self.read_packet()?;
columns.push(Column::from_payload(pld)?);
}
stmt.set_columns(Some(columns));
self.read_packet()?;
}
Ok(stmt)
}
}
}
fn _prepare(&mut self, query: &str, named_params: Option<Vec<String>>) -> MyResult<InnerStmt> {
if let Some(inner_st) = self.stmt_cache.get(query) {
let mut inner_st = inner_st.clone();
inner_st.set_named_params(named_params);
return Ok(inner_st);
}
let inner_st = self._true_prepare(query, named_params)?;
if self.stmt_cache.get_cap() > 0 {
if let Some(old_st) = self.stmt_cache.put(query.into(), inner_st.clone()) {
let mut stmt_id = [0u8; 4];
(&mut stmt_id[..]).write_u32::<LE>(old_st.id())?;
self.write_command_data(Command::COM_STMT_CLOSE, &stmt_id[..])?;
}
}
Ok(inner_st)
}
pub fn prepare<T: AsRef<str>>(&mut self, query: T) -> MyResult<Stmt> {
let query = query.as_ref();
let (named_params, real_query) = parse_named_params(query)?;
match self._prepare(real_query.borrow(), named_params) {
Ok(stmt) => Ok(Stmt::new(stmt, self)),
Err(err) => Err(err),
}
}
pub fn prep_exec<A, T>(&mut self, query: A, params: T) -> MyResult<QueryResult>
where A: AsRef<str>,
T: Into<Params> {
self.prepare(query)?.prep_exec(params.into())
}
pub fn first_exec<Q, P>(&mut self, query: Q, params: P) -> MyResult<Option<Row>>
where Q: AsRef<str>,
P: Into<Params>,
{
self.prep_exec(query, params).and_then(|result| {
for row in result {
return row.map(Some);
}
return Ok(None)
})
}
fn more_results_exists(&self) -> bool {
self.has_results
}
fn connect(&mut self) -> MyResult<()> {
if self.connected {
return Ok(());
}
self.do_handshake().and_then(|_| {
Ok(from_value_opt::<usize>(self.get_system_var("max_allowed_packet").unwrap_or(NULL))
.unwrap_or(0))
}).and_then(|max_allowed_packet| {
if max_allowed_packet == 0 {
Err(DriverError(SetupError))
} else {
self.max_allowed_packet = max_allowed_packet;
self.connected = true;
Ok(())
}
})
}
fn get_system_var(&mut self, name: &str) -> Option<Value> {
for row in self.query(format!("SELECT @@{};", name)).unwrap() {
match row {
Ok(mut r) => match r.len() {
0 => (),
_ => return r.take(0),
},
_ => (),
}
}
return None;
}
fn next_bin(&mut self, columns: &Vec<Column>) -> MyResult<Option<Vec<Value>>> {
if ! self.has_results {
return Ok(None);
}
let pld = match self.read_packet() {
Ok(pld) => pld,
Err(e) => {
self.has_results = false;
return Err(e);
}
};
let x = pld[0];
if x == 0xfe && pld.len() < 0xfe {
self.has_results = false;
let p = EOFPacket::from_payload(pld.as_ref())?;
self.handle_eof(&p);
return Ok(None);
}
let res = Value::from_bin_payload(pld.as_ref(), columns.as_ref());
match res {
Ok(p) => Ok(Some(p)),
Err(e) => {
self.has_results = false;
Err(IoError(e))
}
}
}
fn next_text(&mut self, col_count: usize) -> MyResult<Option<Vec<Value>>> {
if ! self.has_results {
return Ok(None);
}
let pld = match self.read_packet() {
Ok(pld) => pld,
Err(e) => {
self.has_results = false;
return Err(e);
}
};
let x = pld[0];
if (x == 0xfe || x == 0xff) && pld.len() < 0xfe {
self.has_results = false;
if x == 0xfe {
let p = EOFPacket::from_payload(pld.as_ref())?;
self.handle_eof(&p);
return Ok(None);
} else {
let p = ErrPacket::from_payload(pld.as_ref(), self.capability_flags);
match p {
Ok(p) => return Err(MySqlError(p.into())),
Err(err) => return Err(IoError(err))
}
}
}
let res = Value::from_payload(pld.as_ref(), col_count);
match res {
Ok(p) => Ok(Some(p)),
Err(err) => {
self.has_results = false;
Err(IoError(err))
}
}
}
fn has_stmt(&self, query: &str) -> bool {
self.stmt_cache.contains(query)
}
pub fn set_local_infile_handler(&mut self, handler: Option<LocalInfileHandler>) {
self.local_infile_handler = handler;
}
}
impl GenericConnection for Conn {
fn query<T: AsRef<str>>(&mut self, query: T) -> MyResult<QueryResult> {
self.query(query)
}
fn first<T: AsRef<str>>(&mut self, query: T) -> MyResult<Option<Row>> {
self.first(query)
}
fn prepare<T: AsRef<str>>(&mut self, query: T) -> MyResult<Stmt> {
self.prepare(query)
}
fn prep_exec<A, T>(&mut self, query: A, params: T) -> MyResult<QueryResult>
where A: AsRef<str>, T: Into<Params> {
self.prep_exec(query, params)
}
fn first_exec<Q, P>(&mut self, query: Q, params: P) -> MyResult<Option<Row>>
where Q: AsRef<str>, P: Into<Params> {
self.first_exec(query, params)
}
}
impl Drop for Conn {
fn drop(&mut self) {
let stmt_cache = mem::replace(&mut self.stmt_cache, StmtCache::new(0));
let mut stmt_id = [0u8; 4];
for (_, inner_st) in stmt_cache.into_iter() {
let _ = (&mut stmt_id[..]).write_u32::<LE>(inner_st.id());
let _ = self.write_command_data(Command::COM_STMT_CLOSE, &stmt_id[..]);
}
}
}
#[derive(Debug)]
enum ResultConnRef<'a> {
ViaConnRef(&'a mut Conn),
ViaStmt(Stmt<'a>)
}
impl<'a> Deref for ResultConnRef<'a> {
type Target = Conn;
fn deref<'c>(&'c self) -> &'c Conn {
match *self {
ResultConnRef::ViaConnRef(ref conn_ref) => conn_ref,
ResultConnRef::ViaStmt(ref stmt) => stmt.conn.deref(),
}
}
}
impl<'a> DerefMut for ResultConnRef<'a> {
fn deref_mut<'c>(&'c mut self) -> &'c mut Conn {
match *self {
ResultConnRef::ViaConnRef(ref mut conn_ref) => conn_ref,
ResultConnRef::ViaStmt(ref mut stmt) => stmt.conn.deref_mut(),
}
}
}
#[derive(Debug)]
pub struct QueryResult<'a> {
conn: ResultConnRef<'a>,
columns: Arc<Vec<Column>>,
ok_packet: Option<OkPacket>,
is_bin: bool,
}
impl<'a> QueryResult<'a> {
fn new(conn: ResultConnRef<'a>,
columns: Vec<Column>,
ok_packet: Option<OkPacket>,
is_bin: bool) -> QueryResult<'a>
{
QueryResult {
conn: conn,
columns: Arc::new(columns),
ok_packet: ok_packet,
is_bin: is_bin
}
}
fn handle_if_more_results(&mut self) -> Option<MyResult<Row>> {
if self.conn.status_flags.contains(consts::SERVER_MORE_RESULTS_EXISTS) {
match self.conn.handle_result_set() {
Ok((cols, ok_p)) => {
self.columns = Arc::new(cols);
self.ok_packet = ok_p;
None
},
Err(e) => return Some(Err(e)),
}
} else {
None
}
}
pub fn affected_rows(&self) -> u64 {
self.conn.affected_rows
}
pub fn last_insert_id(&self) -> u64 {
self.conn.last_insert_id
}
pub fn warnings(&self) -> u16 {
self.ok_packet.as_ref().map(|ok_p| ok_p.warnings).unwrap_or(0u16)
}
pub fn info(&self) -> Vec<u8> {
if self.ok_packet.is_some() {
self.ok_packet.as_ref().unwrap().info.clone()
} else {
Vec::with_capacity(0)
}
}
pub fn column_index<T: AsRef<str>>(&self, name: T) -> Option<usize> {
let name = name.as_ref().as_bytes();
for (i, c) in self.columns.iter().enumerate() {
if c.name() == name {
return Some(i)
}
}
None
}
pub fn column_indexes(&self) -> HashMap<String, usize, BldHshrDflt<FnvHasher>> {
let mut indexes = HashMap::default();
for (i, column) in self.columns.iter().enumerate() {
indexes.insert(from_utf8(column.name()).unwrap().to_string(), i);
}
indexes
}
pub fn columns_ref(&self) -> &[Column] {
self.columns.as_ref()
}
pub fn more_results_exists(&self) -> bool {
self.conn.has_results
}
}
impl<'a> Iterator for QueryResult<'a> {
type Item = MyResult<Row>;
fn next(&mut self) -> Option<MyResult<Row>> {
let values = if self.is_bin {
self.conn.next_bin(&self.columns)
} else {
self.conn.next_text(self.columns.len())
};
match values {
Ok(values) => {
match values {
Some(values) => Some(Ok(Row::new(values, self.columns.clone()))),
None => self.handle_if_more_results(),
}
},
Err(e) => Some(Err(e)),
}
}
}
impl<'a> Drop for QueryResult<'a> {
fn drop(&mut self) {
while self.conn.more_results_exists() {
while let Some(_) = self.next() {}
}
}
}
#[cfg(test)]
#[allow(non_snake_case)]
mod test {
use Opts;
use OptsBuilder;
static USER: &'static str = "root";
static PASS: &'static str = "password";
static ADDR: &'static str = "localhost";
static PORT: u16 = 3307;
#[cfg(all(feature = "ssl", target_os = "macos"))]
pub fn get_opts() -> Opts {
let pwd: String = ::std::env::var("MYSQL_SERVER_PASS").unwrap_or(PASS.to_string());
let port: u16 = ::std::env::var("MYSQL_SERVER_PORT").ok()
.map(|my_port| my_port.parse().ok().unwrap_or(PORT))
.unwrap_or(PORT);
let mut builder = OptsBuilder::default();
builder.user(Some(USER))
.pass(Some(pwd))
.ip_or_hostname(Some(ADDR))
.tcp_port(port)
.init(vec!["SET GLOBAL sql_mode = 'TRADITIONAL'"])
.verify_peer(true)
.ssl_opts(Some(Some(("tests/client.p12", "pass", vec!["tests/ca-cert.cer"]))));
builder.into()
}
#[cfg(all(feature = "ssl", not(target_os = "macos"), unix))]
pub fn get_opts() -> Opts {
let pwd: String = ::std::env::var("MYSQL_SERVER_PASS").unwrap_or(PASS.to_string());
let port: u16 = ::std::env::var("MYSQL_SERVER_PORT").ok()
.map(|my_port| my_port.parse().ok().unwrap_or(PORT))
.unwrap_or(PORT);
let mut builder = OptsBuilder::default();
builder.user(Some(USER))
.pass(Some(pwd))
.ip_or_hostname(Some(ADDR))
.tcp_port(port)
.init(vec!["SET GLOBAL sql_mode = 'TRADITIONAL'"])
.verify_peer(true)
.ssl_opts(Some(("tests/ca-cert.pem", None::<(String, String)>)));
builder.into()
}
#[cfg(any(not(feature = "ssl"), target_os = "windows"))]
pub fn get_opts() -> Opts {
let pwd: String = ::std::env::var("MYSQL_SERVER_PASS").unwrap_or(PASS.to_string());
let port: u16 = ::std::env::var("MYSQL_SERVER_PORT").ok()
.map(|my_port| my_port.parse().ok().unwrap_or(PORT))
.unwrap_or(PORT);
let mut builder = OptsBuilder::default();
builder.user(Some(USER))
.pass(Some(pwd))
.ip_or_hostname(Some(ADDR))
.tcp_port(port)
.init(vec!["SET GLOBAL sql_mode = 'TRADITIONAL'"]);
builder.into()
}
mod my_conn {
use std::iter;
use std::borrow::ToOwned;
use std::fs;
use std::io::Write;
use time::{Tm, now};
use Conn;
use DriverError::{
MissingNamedParameter,
NamedParamsForPositionalQuery,
};
use Error::DriverError;
use OptsBuilder;
use Params;
use LocalInfileHandler;
use super::super::super::value::{ToValue, from_value, from_row};
use super::super::super::value::Value::{NULL, Int, Bytes, Date};
use super::get_opts;
use super::super::Column;
#[test]
fn should_connect() {
let mut conn = Conn::new(get_opts()).unwrap();
let mode = conn.query("SELECT @@GLOBAL.sql_mode").unwrap().next().unwrap().unwrap().take(0).unwrap();
let mode = from_value::<String>(mode);
assert!(mode.contains("TRADITIONAL"));
assert!(conn.ping());
}
#[test]
fn should_connect_with_database() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.db_name(Some("mysql"));
let mut conn = Conn::new(opts).unwrap();
assert_eq!(conn.query("SELECT DATABASE()").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"mysql".to_vec())]);
}
#[test]
fn should_connect_by_hostname() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.db_name(Some("mysql"));
opts.ip_or_hostname(Some("localhost"));
let mut conn = Conn::new(opts).unwrap();
assert_eq!(conn.query("SELECT DATABASE()").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"mysql".to_vec())]);
}
#[test]
fn should_execute_queryes_and_parse_results() {
let mut conn = Conn::new(get_opts()).unwrap();
assert!(conn.query("CREATE TEMPORARY TABLE x.tbl(\
a TEXT,\
b INT,\
c INT UNSIGNED,\
d DATE,\
e FLOAT
)").is_ok());
assert!(conn.query("INSERT INTO x.tbl(a, b, c, d, e) VALUES (\
'hello',\
-123,\
123,\
'2014-05-05',\
123.123\
)").is_ok());
assert!(conn.query("INSERT INTO x.tbl(a, b, c, d, e) VALUES (\
'world',\
-321,\
321,\
'2014-06-06',\
321.321\
)").is_ok());
assert!(conn.query("SELECT * FROM unexisted").is_err());
assert!(conn.query("SELECT * FROM x.tbl").is_ok());
assert!(conn.query("UPDATE x.tbl SET a = 'foo'").is_ok());
assert_eq!(conn.affected_rows, 2);
assert!(conn.query("SELECT * FROM x.tbl WHERE a = 'bar'").unwrap().next().is_none());
for (i, row) in conn.query("SELECT * FROM x.tbl")
.unwrap().enumerate() {
let row = row.unwrap();
if i == 0 {
assert_eq!(row[0], Bytes(b"foo".to_vec()));
assert_eq!(row[1], Bytes(b"-123".to_vec()));
assert_eq!(row[2], Bytes(b"123".to_vec()));
assert_eq!(row[3], Bytes(b"2014-05-05".to_vec()));
assert_eq!(row[4], Bytes(b"123.123".to_vec()));
} else if i == 1 {
assert_eq!(row[0], Bytes(b"foo".to_vec()));
assert_eq!(row[1], Bytes(b"-321".to_vec()));
assert_eq!(row[2], Bytes(b"321".to_vec()));
assert_eq!(row[3], Bytes(b"2014-06-06".to_vec()));
assert_eq!(row[4], Bytes(b"321.321".to_vec()));
} else {
unreachable!();
}
}
}
#[test]
fn should_parse_large_text_result() {
let mut conn = Conn::new(get_opts()).unwrap();
assert_eq!(
conn.query("SELECT REPEAT('A', 20000000)").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(iter::repeat(b'A').take(20_000_000).collect())]
);
}
#[test]
fn should_execute_statements_and_parse_results() {
let mut conn = Conn::new(get_opts()).unwrap();
assert!(conn.query("CREATE TEMPORARY TABLE x.tbl(\
a TEXT,\
b INT,\
c INT UNSIGNED,\
d DATE,\
e DOUBLE\
)").is_ok());
let _ = conn.prepare("INSERT INTO x.tbl(a, b, c, d, e)\
VALUES (?, ?, ?, ?, ?)")
.and_then(|mut stmt| {
let tm = Tm { tm_year: 114, tm_mon: 4, tm_mday: 5, tm_hour: 0,
tm_min: 0, tm_sec: 0, tm_nsec: 0, ..now() };
let hello = b"hello".to_vec();
assert!(stmt.execute((&hello, -123, 123, tm.to_timespec(), 123.123f64)).is_ok());
assert!(stmt.execute(&[
&b"world".to_vec() as &ToValue,
&NULL as &ToValue,
&NULL as &ToValue,
&NULL as &ToValue,
&321.321f64 as &ToValue
][..]).is_ok());
Ok(())
}).unwrap();
let _ = conn.prepare("SELECT * from x.tbl").and_then(|mut stmt| {
for (i, row) in stmt.execute(()).unwrap().enumerate() {
let mut row = row.unwrap();
if i == 0 {
assert_eq!(row[0], Bytes(b"hello".to_vec()));
assert_eq!(row[1], Int(-123i64));
assert_eq!(row[2], Int(123i64));
assert_eq!(row[3], Date(2014u16, 5u8, 5u8, 0u8, 0u8, 0u8, 0u32));
assert_eq!(row.take::<f64, _>(4).unwrap(), 123.123f64);
} else if i == 1 {
assert_eq!(row[0], Bytes(b"world".to_vec()));
assert_eq!(row[1], NULL);
assert_eq!(row[2], NULL);
assert_eq!(row[3], NULL);
assert_eq!(row.take::<f64, _>(4).unwrap(), 321.321f64);
} else {
unreachable!();
}
}
Ok(())
}).unwrap();
let mut result = conn.prep_exec("SELECT ?, ?, ?", ("hello", 1, 1.1)).unwrap();
let row = result.next().unwrap();
let mut row = row.unwrap();
assert_eq!(row.take::<String, _>(0).unwrap(), "hello".to_string());
assert_eq!(row.take::<i8, _>(1).unwrap(), 1i8);
assert_eq!(row.take::<f32, _>(2).unwrap(), 1.1f32);
}
#[test]
fn should_parse_large_binary_result() {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT REPEAT('A', 20000000);").unwrap();
assert_eq!(
stmt.execute(()).unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(iter::repeat(b'A').take(20_000_000).collect())]
);
}
#[test]
fn should_start_commit_and_rollback_transactions() {
let mut conn = Conn::new(get_opts()).unwrap();
assert!(conn.query("CREATE TEMPORARY TABLE x.tbl(a INT)").is_ok());
let _ = conn.start_transaction(false, None, None).and_then(|mut t| {
assert!(t.query("INSERT INTO x.tbl(a) VALUES(1)").is_ok());
assert!(t.query("INSERT INTO x.tbl(a) VALUES(2)").is_ok());
assert!(t.commit().is_ok());
Ok(())
}).unwrap();
assert_eq!(
conn.query("SELECT COUNT(a) from x.tbl").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"2".to_vec())]
);
let _ = conn.start_transaction(false, None, None).and_then(|mut t| {
assert!(t.query("INSERT INTO tbl(a) VALUES(1)").is_err());
Ok(())
}).unwrap();
assert_eq!(
conn.query("SELECT COUNT(a) from x.tbl").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"2".to_vec())]
);
let _ = conn.start_transaction(false, None, None).and_then(|mut t| {
assert!(t.query("INSERT INTO x.tbl(a) VALUES(1)").is_ok());
assert!(t.query("INSERT INTO x.tbl(a) VALUES(2)").is_ok());
assert!(t.rollback().is_ok());
Ok(())
}).unwrap();
assert_eq!(
conn.query("SELECT COUNT(a) from x.tbl").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"2".to_vec())]
);
let _ = conn.start_transaction(false, None, None).and_then(|mut t| {
let _ = t.prepare("INSERT INTO x.tbl(a) VALUES(?)")
.and_then(|mut stmt| {
assert!(stmt.execute((3,)).is_ok());
assert!(stmt.execute((4,)).is_ok());
Ok(())
}).unwrap();
assert!(t.commit().is_ok());
Ok(())
}).unwrap();
assert_eq!(
conn.query("SELECT COUNT(a) from x.tbl").unwrap().next().unwrap().unwrap().unwrap(),
vec![Bytes(b"4".to_vec())]
);
let _ = conn.start_transaction(false, None, None). and_then(|mut t| {
t.prep_exec("INSERT INTO x.tbl(a) VALUES(?)", (5,)).unwrap();
t.prep_exec("INSERT INTO x.tbl(a) VALUES(?)", (6,)).unwrap();
Ok(())
}).unwrap();
}
#[test]
fn should_handle_LOCAL_INFILE() {
let mut conn = Conn::new(get_opts()).unwrap();
assert!(conn.query("CREATE TEMPORARY TABLE x.tbl(a TEXT)").is_ok());
let path = ::std::path::PathBuf::from("local_infile.txt");
{
let mut file = fs::File::create(&path).unwrap();
let _ = file.write(b"AAAAAA\n");
let _ = file.write(b"BBBBBB\n");
let _ = file.write(b"CCCCCC\n");
}
let query = format!("LOAD DATA LOCAL INFILE '{}' INTO TABLE x.tbl",
path.to_str().unwrap().to_owned());
conn.query(query).unwrap();
for (i, row) in conn.query("SELECT * FROM x.tbl")
.unwrap().enumerate() {
let row = row.unwrap();
match i {
0 => assert_eq!(row.unwrap(), vec!(Bytes(b"AAAAAA".to_vec()))),
1 => assert_eq!(row.unwrap(), vec!(Bytes(b"BBBBBB".to_vec()))),
2 => assert_eq!(row.unwrap(), vec!(Bytes(b"CCCCCC".to_vec()))),
_ => unreachable!()
}
}
let _ = fs::remove_file(&path);
}
#[test]
fn should_handle_LOCAL_INFILE_with_custom_handler() {
let mut conn = Conn::new(get_opts()).unwrap();
conn.query("CREATE TEMPORARY TABLE x.tbl(a TEXT)").unwrap();
conn.set_local_infile_handler(Some(
LocalInfileHandler::new(|_, stream| {
let mut cell_data = vec![b'Z'; 65535];
cell_data.push(b'\n');
for _ in 0..1536 {
stream.write_all(&*cell_data)?;
}
Ok(())
})
));
conn.query("LOAD DATA LOCAL INFILE 'file_name' INTO TABLE x.tbl").unwrap();
let count = conn.query("SELECT * FROM x.tbl").unwrap().map(|row| {
assert_eq!(from_row::<(Vec<u8>,)>(row.unwrap()).0.len(), 65535);
1
}).sum::<usize>();
assert_eq!(count, 1536);
}
#[test]
fn should_reset_connection() {
let mut conn = Conn::new(get_opts()).unwrap();
assert!(conn.query("CREATE TEMPORARY TABLE `db`.`test` \
(`test` VARCHAR(255) NULL);").is_ok());
assert!(conn.query("SELECT * FROM `db`.`test`;").is_ok());
assert!(conn.reset().is_ok());
assert!(conn.query("SELECT * FROM `db`.`test`;").is_err());
}
#[test]
fn should_connect_via_socket_for_127_0_0_1() {
let mut opts = OptsBuilder::from_opts(get_opts());
#[cfg(all(feature = "ssl", not(target_os = "windows")))]
opts.ssl_opts::<String, String, String>(None);
let conn = Conn::new(opts).unwrap();
let debug_format = format!("{:#?}", conn);
assert!(debug_format.contains("SocketStream"));
}
#[test]
fn should_connect_via_socket_localhost() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.ip_or_hostname(Some("localhost"));
#[cfg(all(feature = "ssl", not(target_os = "windows")))]
opts.ssl_opts::<String, String, String>(None);
let conn = Conn::new(opts).unwrap();
let debug_format = format!("{:?}", conn);
assert!(debug_format.contains("SocketStream"));
}
#[test]
#[cfg(all(feature = "ssl", any(target_os = "macos", unix)))]
fn should_connect_via_ssl() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
let conn = Conn::new(opts).unwrap();
let debug_format = format!("{:#?}", conn);
assert!(debug_format.contains("Secure stream"));
}
#[test]
fn should_handle_multi_resultset() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
opts.db_name(Some("mysql"));
let mut conn = Conn::new(opts).unwrap();
assert!(conn.query("DROP PROCEDURE IF EXISTS multi").is_ok());
assert!(conn.query(r#"CREATE PROCEDURE multi() BEGIN
SELECT 1;
SELECT 1;
END"#).is_ok());
for (i, row) in conn.query("CALL multi()")
.unwrap().enumerate() {
match i {
0 | 1 => assert_eq!(row.unwrap().unwrap(), vec![Bytes(b"1".to_vec())]),
_ => unreachable!(),
}
}
let mut result = conn.query("SELECT 1; SELECT 2; SELECT 3;").unwrap();
let mut i = 0;
while { i += 1; result.more_results_exists() } {
for row in result.by_ref() {
match i {
1 => assert_eq!(
row.unwrap().unwrap(),
vec![Bytes(b"1".to_vec())]
),
2 => assert_eq!(
row.unwrap().unwrap(),
vec![Bytes(b"2".to_vec())]
),
3 => assert_eq!(
row.unwrap().unwrap(),
vec![Bytes(b"3".to_vec())]
),
_ => unreachable!(),
}
}
}
assert_eq!(i, 4);
}
#[test]
fn should_work_with_named_params() {
let mut conn = Conn::new(get_opts()).unwrap();
{
let mut stmt = conn.prepare("SELECT :a, :b, :a, :c").unwrap();
let mut result = stmt.execute(params!{"a" => 1, "b" => 2, "c" => 3}).unwrap();
let row = result.next().unwrap().unwrap();
assert_eq!((1, 2, 1, 3), from_row(row));
}
let mut result = conn.prep_exec("SELECT :a, :b, :a + :b, :c", params!{
"a" => 1,
"b" => 2,
"c" => 3,
}).unwrap();
let row = result.next().unwrap().unwrap();
assert_eq!((1, 2, 3, 3), from_row(row));
}
#[test]
fn should_return_error_on_missing_named_parameter() {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT :a, :b, :a, :c, :d").unwrap();
let result = stmt.execute(params!{"a" => 1, "b" => 2, "c" => 3,});
match result {
Err(DriverError(MissingNamedParameter(ref x))) if x == "d" => (),
_ => assert!(false),
}
}
#[test]
fn should_return_error_on_named_params_for_positional_statement() {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT ?, ?, ?, ?, ?").unwrap();
let result = stmt.execute(params!{"a" => 1, "b" => 2, "c" => 3,});
match result {
Err(DriverError(NamedParamsForPositionalQuery)) => (),
_ => assert!(false),
}
}
#[test]
#[should_panic]
fn should_panic_on_named_param_redefinition() {
let _: Params = params!{"a" => 1, "b" => 2, "a" => 3}.into();
}
#[test]
fn should_parse_column_from_payload() {
let payload1 = b"\x03def\x06schema\x05table\x09org_table\x04name\x08org_name\
\x0c\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00".to_vec();
let payload2 = b"\x03def\x06schema\x05table\x09org_table\x04name\x08org_name\
\x0c\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x07default".to_vec();
let out1 = Column::from_payload(payload1).unwrap();
let out2 = Column::from_payload(payload2).unwrap();
assert_eq!(out1.schema(), b"schema");
assert_eq!(out1.table(), b"table");
assert_eq!(out1.org_table(), b"org_table");
assert_eq!(out1.name(), b"name");
assert_eq!(out1.org_name(), b"org_name");
assert_eq!(out1.default_values(), None);
assert_eq!(out2.schema(), b"schema");
assert_eq!(out2.table(), b"table");
assert_eq!(out2.org_table(), b"org_table");
assert_eq!(out2.name(), b"name");
assert_eq!(out2.org_name(), b"org_name");
assert_eq!(out2.default_values(), Some(&b"default"[..]));
}
#[test]
fn should_handle_tcp_connect_timeout() {
use error::Error::DriverError;
use error::DriverError::ConnectTimeout;
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
opts.tcp_connect_timeout(Some(::std::time::Duration::from_millis(1000)));
assert!(Conn::new(opts).unwrap().ping());
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
opts.tcp_connect_timeout(Some(::std::time::Duration::from_millis(1000)));
opts.ip_or_hostname(Some("192.168.255.255"));
match Conn::new(opts).unwrap_err() {
DriverError(ConnectTimeout) => {},
err => panic!("Unexpected error: {}", err),
}
}
#[test]
#[cfg(not(feature = "ssl"))]
fn should_bind_before_connect() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
opts.ip_or_hostname(Some("127.0.0.1"));
opts.bind_address(Some(([127, 0, 0, 1], 27272)));
let conn = Conn::new(opts).unwrap();
let debug_format: String = format!("{:?}", conn);
assert!(debug_format.contains("addr: V4(127.0.0.1:27272)"));
}
#[test]
#[cfg(not(feature = "ssl"))]
fn should_bind_before_connect_with_timeout() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.prefer_socket(false);
opts.ip_or_hostname(Some("127.0.0.1"));
opts.bind_address(Some(([127, 0, 0, 1], 27273)));
opts.tcp_connect_timeout(Some(::std::time::Duration::from_millis(1000)));
let mut conn = Conn::new(opts).unwrap();
assert!(conn.ping());
let debug_format: String = format!("{:?}", conn);
assert!(debug_format.contains("addr: V4(127.0.0.1:27273)"));
}
#[test]
fn should_not_cache_statements_if_stmt_cache_size_is_zero() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.stmt_cache_size(0);
let mut conn = Conn::new(opts).unwrap();
conn.prepare("DO 1").unwrap();
conn.prepare("DO 2").unwrap();
conn.prepare("DO 3").unwrap();
let row = conn.first("SHOW SESSION STATUS LIKE 'Com_stmt_close';").unwrap().unwrap();
assert_eq!(from_row::<(String, usize)>(row).1, 3);
}
#[test]
fn should_hold_stmt_cache_size_bound() {
let mut opts = OptsBuilder::from_opts(get_opts());
opts.stmt_cache_size(3);
let mut conn = Conn::new(opts).unwrap();
conn.prepare("DO 1").unwrap();
conn.prepare("DO 2").unwrap();
conn.prepare("DO 3").unwrap();
conn.prepare("DO 1").unwrap();
conn.prepare("DO 4").unwrap();
conn.prepare("DO 3").unwrap();
conn.prepare("DO 5").unwrap();
conn.prepare("DO 6").unwrap();
let row = conn.first("SHOW SESSION STATUS LIKE 'Com_stmt_close';").unwrap().unwrap();
assert_eq!(from_row::<(String, usize)>(row).1, 3);
let order = conn.stmt_cache.iter().collect::<Vec<&String>>();
assert_eq!(order, &["DO 3", "DO 5", "DO 6"]);
}
#[test]
fn should_handle_json_columns() {
#[cfg(feature = "rustc_serialize")]
use rustc_serialize::json::Json;
#[cfg(not(feature = "rustc_serialize"))]
use serde_json::Value as Json;
#[cfg(not(feature = "rustc_serialize"))]
use std::str::FromStr;
use Serialized;
use Deserialized;
#[cfg(feature = "rustc_serialize")]
#[derive(RustcDecodable, RustcEncodable, Debug, Eq, PartialEq)]
pub struct DecTest {
foo: String,
quux: (u64, String),
}
#[cfg(not(feature = "rustc_serialize"))]
#[derive(Serialize, Deserialize, Debug, Eq, PartialEq)]
pub struct DecTest {
foo: String,
quux: (u64, String),
}
let decodable = DecTest {
foo: "bar".into(),
quux: (42, "hello".into()),
};
let mut conn = Conn::new(get_opts()).unwrap();
if conn.query("CREATE TEMPORARY TABLE x.tbl(a VARCHAR(32), b JSON)").is_err() {
conn.query("CREATE TEMPORARY TABLE x.tbl(a VARCHAR(32), b TEXT)").unwrap();
}
conn.prep_exec(
r#"INSERT INTO x.tbl VALUES ('hello', ?)"#,
(Serialized(&decodable), )
).unwrap();
let row = conn.first("SELECT a, b FROM x.tbl").unwrap().unwrap();
let (a, b): (String, Json) = from_row(row);
assert_eq!((a, b), ("hello".into(), Json::from_str(r#"{"foo": "bar", "quux": [42, "hello"]}"#).unwrap()));
let row = conn.first_exec("SELECT a, b FROM x.tbl WHERE a = ?", ("hello", )).unwrap().unwrap();
let (a, Deserialized(b)) = from_row(row);
assert_eq!((a, b), (String::from("hello"), decodable));
}
}
#[cfg(feature = "nightly")]
mod bench {
use test;
use super::get_opts;
use super::super::{Conn};
use super::super::super::value::Value::NULL;
#[bench]
fn simple_exec(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
bencher.iter(|| { let _ = conn.query("DO 1"); })
}
#[bench]
fn prepared_exec(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("DO 1").unwrap();
bencher.iter(|| { let _ = stmt.execute(()); })
}
#[bench]
fn prepare_and_exec(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
bencher.iter(|| {
let mut stmt = conn.prepare("SELECT ?").unwrap();
let _ = stmt.execute((0,)).unwrap();
})
}
#[bench]
fn simple_query_row(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
bencher.iter(|| { let _ = conn.query("SELECT 1"); })
}
#[bench]
fn simple_prepared_query_row(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT 1").unwrap();
bencher.iter(|| { let _ = stmt.execute(()); })
}
#[bench]
fn simple_prepared_query_row_with_param(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT ?").unwrap();
bencher.iter(|| { let _ = stmt.execute((0,)); })
}
#[bench]
fn simple_prepared_query_row_with_named_param(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT :a").unwrap();
bencher.iter(|| { let _ = stmt.execute(params!{"a" => 0}); })
}
#[bench]
fn simple_prepared_query_row_with_5_params(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT ?, ?, ?, ?, ?").unwrap();
let params = (42i8, b"123456".to_vec(), 1.618f64, NULL, 1i8);
bencher.iter(|| { let _ = stmt.execute(¶ms); })
}
#[bench]
fn simple_prepared_query_row_with_5_named_params(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT :one, :two, :three, :four, :five").unwrap();
bencher.iter(|| {
let _ = stmt.execute(params!{
"one" => 42i8,
"two" => b"123456",
"three" => 1.618f64,
"four" => NULL,
"five" => 1i8,
});
})
}
#[bench]
fn select_large_string(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
bencher.iter(|| { let _ = conn.query("SELECT REPEAT('A', 10000)"); })
}
#[bench]
fn select_prepared_large_string(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
let mut stmt = conn.prepare("SELECT REPEAT('A', 10000)").unwrap();
bencher.iter(|| { let _ = stmt.execute(()); })
}
#[bench]
fn many_small_rows(bencher: &mut test::Bencher) {
let mut conn = Conn::new(get_opts()).unwrap();
conn.query("CREATE TEMPORARY TABLE x.x (id INT)").unwrap();
for _ in 0..512 {
conn.query("INSERT INTO x.x VALUES (256)").unwrap();
}
let mut stmt = conn.prepare("SELECT * FROM x.x").unwrap();
bencher.iter(|| {
let _ = stmt.execute(());
});
}
}
}