#![deny(missing_docs)]
extern crate byteorder;
extern crate chrono;
extern crate mysql_common as myc;
#[macro_use]
extern crate nom;
use std::collections::HashMap;
use std::io;
use std::io::prelude::*;
use std::iter;
use std::net;
pub use myc::constants::{ColumnFlags, ColumnType, StatusFlags};
mod commands;
mod errorcodes;
mod packet;
mod params;
mod resultset;
mod value;
mod writers;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Column {
pub table: String,
pub column: String,
pub coltype: ColumnType,
pub colflags: ColumnFlags,
}
pub use errorcodes::ErrorKind;
pub use params::{ParamParser, ParamValue, Params};
pub use resultset::{QueryResultWriter, RowWriter, StatementMetaWriter};
pub use value::{ToMysqlValue, Value, ValueInner};
pub trait MysqlShim<W: Write> {
type Error: From<io::Error>;
fn on_prepare(&mut self, query: &str, info: StatementMetaWriter<W>) -> Result<(), Self::Error>;
fn on_execute(
&mut self,
id: u32,
params: ParamParser,
results: QueryResultWriter<W>,
) -> Result<(), Self::Error>;
fn on_close(&mut self, stmt: u32);
fn on_query(&mut self, query: &str, results: QueryResultWriter<W>) -> Result<(), Self::Error>;
}
pub struct MysqlIntermediary<B, R: Read, W: Write> {
shim: B,
reader: packet::PacketReader<R>,
writer: packet::PacketWriter<W>,
}
impl<B: MysqlShim<net::TcpStream>> MysqlIntermediary<B, net::TcpStream, net::TcpStream> {
pub fn run_on_tcp(shim: B, stream: net::TcpStream) -> Result<(), B::Error> {
let w = stream.try_clone()?;
MysqlIntermediary::run_on(shim, stream, w)
}
}
impl<B: MysqlShim<S>, S: Read + Write + Clone> MysqlIntermediary<B, S, S> {
pub fn run_on_stream(shim: B, stream: S) -> Result<(), B::Error> {
MysqlIntermediary::run_on(shim, stream.clone(), stream)
}
}
#[derive(Default)]
struct StatementData {
long_data: HashMap<u16, Vec<u8>>,
bound_types: Vec<(myc::constants::ColumnType, bool)>,
params: u16,
}
impl<B: MysqlShim<W>, R: Read, W: Write> MysqlIntermediary<B, R, W> {
pub fn run_on(shim: B, reader: R, writer: W) -> Result<(), B::Error> {
let r = packet::PacketReader::new(reader);
let w = packet::PacketWriter::new(writer);
let mut mi = MysqlIntermediary {
shim,
reader: r,
writer: w,
};
mi.init()?;
mi.run()
}
fn init(&mut self) -> Result<(), B::Error> {
self.writer.write_all(&[10])?;
self.writer.write_all(&b"5.1.10-alpha-msql-proxy\0"[..])?;
self.writer.write_all(&[0x08, 0x00, 0x00, 0x00])?; self.writer.write_all(&b";X,po_k}\0"[..])?; self.writer.write_all(&[0x00, 0x42])?; self.writer.write_all(&[0x21])?; self.writer.write_all(&[0x00, 0x00])?; self.writer.write_all(&[0x00, 0x00])?; self.writer.write_all(&[0x00])?; self.writer.write_all(&[0x00; 6][..])?; self.writer.write_all(&[0x00; 4][..])?; self.writer.write_all(&b">o6^Wz!/kM}N\0"[..])?; self.writer.flush()?;
{
let (seq, handshake) = self.reader.next()?.unwrap();
let _handshake = commands::client_handshake(&handshake).unwrap().1;
self.writer.set_seq(seq + 1);
}
writers::write_ok_packet(&mut self.writer, 0, 0, StatusFlags::empty())?;
self.writer.flush()?;
Ok(())
}
fn run(mut self) -> Result<(), B::Error> {
use commands::Command;
let mut stmts: HashMap<u32, _> = HashMap::new();
while let Some((seq, packet)) = self.reader.next()? {
self.writer.set_seq(seq + 1);
let cmd = commands::parse(&packet).unwrap().1;
match cmd {
Command::Query(q) => {
let w = QueryResultWriter::new(&mut self.writer, false);
if q.starts_with(b"SELECT @@") || q.starts_with(b"select @@") {
let var = &q[b"SELECT @@".len()..];
match var {
b"max_allowed_packet" => {
let cols = &[Column {
table: String::new(),
column: "@@max_allowed_packet".to_owned(),
coltype: myc::constants::ColumnType::MYSQL_TYPE_SHORT,
colflags: myc::constants::ColumnFlags::UNSIGNED_FLAG,
}];
let mut w = w.start(cols)?;
w.write_row(iter::once(1024u16))?;
w.finish()?;
}
_ => {
w.completed(0, 0)?;
}
}
} else {
self.shim.on_query(
::std::str::from_utf8(q)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?,
w,
)?;
}
}
Command::Prepare(q) => {
let w = StatementMetaWriter {
writer: &mut self.writer,
stmts: &mut stmts,
};
self.shim.on_prepare(
::std::str::from_utf8(q)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?,
w,
)?;
}
Command::Execute { stmt, params } => {
let state = stmts.get_mut(&stmt).ok_or(io::Error::new(
io::ErrorKind::InvalidData,
format!("asked to execute unknown statement {}", stmt),
))?;
{
let params = params::ParamParser::new(params, state);
let w = QueryResultWriter::new(&mut self.writer, true);
self.shim.on_execute(stmt, params, w)?;
}
state.long_data.clear();
}
Command::SendLongData { stmt, param, data } => {
stmts
.get_mut(&stmt)
.ok_or(io::Error::new(
io::ErrorKind::InvalidData,
format!("got long data packet for unknown statement {}", stmt),
))?
.long_data
.entry(param)
.or_insert_with(Vec::new)
.extend(data);
}
Command::Close(stmt) => {
self.shim.on_close(stmt);
stmts.remove(&stmt);
}
Command::ListFields(_) => {
let cols = &[Column {
table: String::new(),
column: "not implemented".to_owned(),
coltype: myc::constants::ColumnType::MYSQL_TYPE_SHORT,
colflags: myc::constants::ColumnFlags::UNSIGNED_FLAG,
}];
writers::write_column_definitions(cols, &mut self.writer, true)?;
}
Command::Init(_) | Command::Ping => {
writers::write_ok_packet(&mut self.writer, 0, 0, StatusFlags::empty())?;
}
Command::Quit => {
break;
}
}
self.writer.flush()?;
}
Ok(())
}
}