From ceaac705d6879da76d9d53005b35451e39859ae4 Mon Sep 17 00:00:00 2001 From: Steven Fackler Date: Sat, 24 Aug 2013 23:47:36 -0400 Subject: [PATCH] Parameter binding support! --- src/lib.rs | 159 +++++++++++++++++++++++++++++++++++++++++++++---- src/message.rs | 8 +-- src/test.rs | 33 ++++++++-- 3 files changed, 180 insertions(+), 20 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index 50700a61..1de2fdd8 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -46,7 +46,7 @@ impl PostgresConnection { loop { match conn.read_message() { ParameterStatus(param, value) => - printfln!("Param %s = %s", param, value), + info!("Param %s = %s", param, value), BackendKeyData(*) => (), ReadyForQuery(*) => break, resp => fail!("Bad response: %?", resp.to_str()) @@ -173,9 +173,15 @@ impl<'self> PostgresStatement<'self> { self.num_params } - fn execute(&self, portal_name: &str) { + fn execute(&self, portal_name: &str, params: &[&ToSql]) { + if self.num_params != params.len() { + fail!("Expected %u params but got %u", self.num_params, + params.len()); + } + let formats = []; - let values = []; + let values: ~[Option<~[u8]>] = params.iter().map(|val| val.to_sql()) + .collect(); let result_formats = []; self.conn.write_message(&Bind(portal_name, self.name.as_slice(), @@ -190,8 +196,9 @@ impl<'self> PostgresStatement<'self> { } } - pub fn update(&self) -> uint { - self.execute(""); + pub fn update(&self, params: &[&ToSql]) -> uint { + // The unnamed portal is automatically cleaned up at sync time + self.execute("", params); let mut num = 0; loop { @@ -216,12 +223,12 @@ impl<'self> PostgresStatement<'self> { num } - pub fn query<'a>(&'a self) -> PostgresResult<'a> { + pub fn query<'a>(&'a self, params: &[&ToSql]) -> PostgresResult<'a> { let id = self.next_portal_id.take(); let portal_name = ifmt!("{:s}_portal_{}", self.name.as_slice(), id); self.next_portal_id.put_back(id + 1); - self.execute(portal_name); + self.execute(portal_name, params); let mut data = ~[]; loop { @@ -313,9 +320,139 @@ pub trait FromSql { fn from_sql(raw: &Option<~[u8]>) -> Self; } -impl FromSql for int { - fn from_sql(raw: &Option<~[u8]>) -> int { - FromStr::from_str(str::from_bytes_slice(raw.get_ref().as_slice())) - .unwrap() +macro_rules! from_str_impl( + ($t:ty) => ( + impl FromSql for Option<$t> { + fn from_sql(raw: &Option<~[u8]>) -> Option<$t> { + match *raw { + None => None, + Some(ref buf) => { + let s = str::from_bytes_slice(buf.as_slice()); + Some(FromStr::from_str(s).unwrap()) + } + } + } + } + ) +) + +macro_rules! from_option_impl( + ($t:ty) => ( + impl FromSql for $t { + fn from_sql(raw: &Option<~[u8]>) -> $t { + FromSql::from_sql::>(raw).unwrap() + } + } + ) +) + +from_str_impl!(int) +from_option_impl!(int) +from_str_impl!(i8) +from_option_impl!(i8) +from_str_impl!(i16) +from_option_impl!(i16) +from_str_impl!(i32) +from_option_impl!(i32) +from_str_impl!(i64) +from_option_impl!(i64) +from_str_impl!(uint) +from_option_impl!(uint) +from_str_impl!(u8) +from_option_impl!(u8) +from_str_impl!(u16) +from_option_impl!(u16) +from_str_impl!(u32) +from_option_impl!(u32) +from_str_impl!(u64) +from_option_impl!(u64) +from_str_impl!(float) +from_option_impl!(float) +from_str_impl!(f32) +from_option_impl!(f32) +from_str_impl!(f64) +from_option_impl!(f64) + +impl FromSql for Option<~str> { + fn from_sql(raw: &Option<~[u8]>) -> Option<~str> { + do raw.chain_ref |buf| { + Some(str::from_bytes(buf.as_slice())) + } + } +} +from_option_impl!(~str) + +pub trait ToSql { + fn to_sql(&self) -> Option<~[u8]>; +} + +macro_rules! to_str_impl( + ($t:ty) => ( + impl ToSql for $t { + fn to_sql(&self) -> Option<~[u8]> { + Some(self.to_str().into_bytes()) + } + } + ) +) + +macro_rules! to_option_impl( + ($t:ty) => ( + impl ToSql for Option<$t> { + fn to_sql(&self) -> Option<~[u8]> { + do self.chain |val| { + val.to_sql() + } + } + } + ) +) + +to_str_impl!(int) +to_option_impl!(int) +to_str_impl!(i8) +to_option_impl!(i8) +to_str_impl!(i16) +to_option_impl!(i16) +to_str_impl!(i32) +to_option_impl!(i32) +to_str_impl!(i64) +to_option_impl!(i64) +to_str_impl!(uint) +to_option_impl!(uint) +to_str_impl!(u8) +to_option_impl!(u8) +to_str_impl!(u16) +to_option_impl!(u16) +to_str_impl!(u32) +to_option_impl!(u32) +to_str_impl!(u64) +to_option_impl!(u64) +to_str_impl!(float) +to_option_impl!(float) +to_str_impl!(f32) +to_option_impl!(f32) +to_str_impl!(f64) +to_option_impl!(f64) + +impl<'self> ToSql for &'self str { + fn to_sql(&self) -> Option<~[u8]> { + Some(self.as_bytes().to_owned()) + } +} + +impl ToSql for Option<~str> { + fn to_sql(&self) -> Option<~[u8]> { + do self.chain_ref |val| { + val.to_sql() + } + } +} + +impl<'self> ToSql for Option<&'self str> { + fn to_sql(&self) -> Option<~[u8]> { + do self.chain |val| { + val.to_sql() + } } } diff --git a/src/message.rs b/src/message.rs index 7a3c634f..4617fd0f 100644 --- a/src/message.rs +++ b/src/message.rs @@ -41,7 +41,7 @@ pub struct RowDescriptionEntry { pub enum FrontendMessage<'self> { /// portal, stmt, formats, values, result formats - Bind(&'self str, &'self str, &'self [i16], &'self [Option<&'self [u8]>], + Bind(&'self str, &'self str, &'self [i16], &'self [Option<~[u8]>], &'self [i16]), Close(u8, &'self str), Describe(u8, &'self str), @@ -90,11 +90,11 @@ impl WriteMessage for W { for value in values.iter() { match *value { None => { - buf.write_be_i16_(-1); + buf.write_be_i32_(-1); } - Some(value) => { + Some(ref value) => { buf.write_be_i32_(value.len() as i32); - buf.write(value); + buf.write(*value); } } } diff --git a/src/test.rs b/src/test.rs index 698f351c..62a73630 100644 --- a/src/test.rs +++ b/src/test.rs @@ -1,13 +1,13 @@ extern mod postgres; -use postgres::PostgresConnection; +use postgres::{PostgresConnection, ToSql}; #[test] fn test_basic() { let conn = PostgresConnection::connect("postgres://postgres@127.0.0.1:5432"); do conn.in_transaction |conn| { - conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)").update(); + conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)").update([]); Err::<(), ()>(()) }; @@ -18,13 +18,36 @@ fn test_query() { let conn = PostgresConnection::connect("postgres://postgres@127.0.0.1:5432"); do conn.in_transaction |conn| { - conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)").update(); - conn.prepare("INSERT INTO foo (id) VALUES (1), (2)").update(); + conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)").update([]); + conn.prepare("INSERT INTO foo (id) VALUES ($1), ($2)") + .update([&1 as &ToSql, &2 as &ToSql]); let stmt = conn.prepare("SELECT * from foo ORDER BY id"); - let result = stmt.query(); + let result = stmt.query([]); assert_eq!(~[1, 2], result.iter().map(|row| { row[0] }).collect()); Err::<(), ()>(()) }; } + +#[test] +fn test_nulls() { + let conn = PostgresConnection::connect("postgres://postgres@127.0.0.1:5432"); + + do conn.in_transaction |conn| { + conn.prepare("CREATE TABLE foo ( + id BIGINT PRIMARY KEY, + val VARCHAR + )").update([]); + conn.prepare("INSERT INTO foo (id, val) VALUES ($1, $2), ($3, $4)") + .update([&1 as &ToSql, & &"foobar" as &ToSql, + &2 as &ToSql, &None::<~str> as &ToSql]); + let stmt = conn.prepare("SELECT id, val FROM foo ORDER BY id"); + let result = stmt.query([]); + + assert_eq!(~[Some(~"foobar"), None], + result.iter().map(|row| { row[1] }).collect()); + + Err::<(), ()>(()) + }; +}