diff --git a/src/lib.rs b/src/lib.rs index 3e17993e..1d65fdfb 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -38,18 +38,13 @@ impl PostgresConnection { next_stmt_id: Cell::new(0) }; - do conn.stream.with_mut_ref |s| { - let mut args = HashMap::new(); - args.insert(~"user", username.to_owned()); - s.write_message(&StartupMessage(args)); - } + let mut args = HashMap::new(); + args.insert(&"user", username); + conn.write_message(&StartupMessage(args)); - let resp = do conn.stream.with_mut_ref |s| { - s.read_message() - }; - match resp { + match conn.read_message() { AuthenticationOk => (), - _ => fail!("Bad response: %?", resp) + resp => fail!("Bad response: %?", resp) } conn.finish_connect(); @@ -59,47 +54,79 @@ impl PostgresConnection { fn finish_connect(&self) { loop { - match self.stream.with_mut_ref(|s| { s.read_message() }) { + match self.read_message() { ParameterStatus(param, value) => printfln!("Param %s = %s", param, value), - BackendKeyData(*) => loop, - ReadyForQuery(_) => break, + BackendKeyData(*) => (), + ReadyForQuery(*) => break, resp => fail!("Bad response: %?", resp) } } } + fn write_message(&self, message: &FrontendMessage) { + do self.stream.with_mut_ref |s| { + s.write_message(message); + } + } + + fn read_message(&self) -> BackendMessage { + do self.stream.with_mut_ref |s| { + s.read_message() + } + } + pub fn prepare<'a>(&'a self, query: &str) -> PostgresStatement<'a> { let id = self.next_stmt_id.take(); - let query_name = ifmt!("statement_{}", id); + let stmt_name = ifmt!("statement_{}", id); self.next_stmt_id.put_back(id + 1); - do self.stream.with_mut_ref |s| { - s.write_message(&Parse(query_name.clone(), query.to_owned(), ~[])); - } + let types = []; + self.write_message(&Parse(stmt_name, query, types)); + self.write_message(&Sync); - do self.stream.with_mut_ref |s| { - s.write_message(&Sync); - } - - match self.stream.with_mut_ref(|s| { s.read_message() }) { + match self.read_message() { ParseComplete => (), ErrorResponse(ref data) => fail!("Error: %?", data), resp => fail!("Bad response: %?", resp) } - match self.stream.with_mut_ref(|s| { s.read_message() }) { - ReadyForQuery(*) => (), + self.wait_for_ready(); + + self.write_message(&Describe('S' as u8, stmt_name)); + self.write_message(&Sync); + + let num_params = match self.read_message() { + ParameterDescription(ref types) => types.len(), + resp => fail!("Bad response: %?", resp) + }; + + match self.read_message() { + RowDescription(*) | NoData => (), resp => fail!("Bad response: %?", resp) } - PostgresStatement { conn: self, name: query_name } + self.wait_for_ready(); + + PostgresStatement { + conn: self, + name: stmt_name, + num_params: num_params + } + } + + fn wait_for_ready(&self) { + match self.read_message() { + ReadyForQuery(*) => (), + resp => fail!("Bad response: %?", resp) + } } } pub struct PostgresStatement<'self> { priv conn: &'self PostgresConnection, - priv name: ~str + priv name: ~str, + priv num_params: uint } impl<'self> PostgresStatement<'self> { diff --git a/src/message.rs b/src/message.rs index 76b3863d..32f8f060 100644 --- a/src/message.rs +++ b/src/message.rs @@ -5,6 +5,7 @@ use std::rt::io::extensions::{ReaderUtil, ReaderByteConversions, use std::rt::io::mem::{MemWriter, MemReader}; use std::hashmap::HashMap; use std::sys; +use std::vec; pub static PROTOCOL_VERSION: i32 = 0x0003_0000; @@ -12,16 +13,30 @@ pub enum BackendMessage { AuthenticationOk, BackendKeyData(i32, i32), ErrorResponse(HashMap), + NoData, + ParameterDescription(~[i32]), ParameterStatus(~str, ~str), ParseComplete, - ReadyForQuery(u8) + ReadyForQuery(u8), + RowDescription(~[RowDescriptionEntry]) } -pub enum FrontendMessage { +pub struct RowDescriptionEntry { + name: ~str, + table_oid: i32, + column_id: i16, + type_oid: i32, + type_size: i16, + type_modifier: i32, + format: i16 +} + +pub enum FrontendMessage<'self> { + Describe(u8, &'self str), /// name, query, parameter types - Parse(~str, ~str, ~[i32]), - Query(~str), - StartupMessage(HashMap<~str, ~str>), + Parse(&'self str, &'self str, &'self [i32]), + Query(&'self str), + StartupMessage(HashMap<&'self str, &'self str>), Sync, Terminate } @@ -48,6 +63,11 @@ impl WriteMessage for W { let mut ident = None; match *message { + Describe(variant, ref name) => { + ident = Some('D'); + buf.write_u8_(variant); + buf.write_string(*name); + } Parse(ref name, ref query, ref param_types) => { ident = Some('P'); buf.write_string(*name); @@ -118,7 +138,6 @@ impl ReadMessage for R { debug!("Reading message"); let ident = self.read_u8_(); - debug!("Ident %?", ident); // subtract size of length value let len = self.read_be_i32_() as uint - sys::size_of::(); let mut buf = MemReader::new(self.read_bytes(len)); @@ -127,8 +146,11 @@ impl ReadMessage for R { '1' => ParseComplete, 'E' => read_error_message(&mut buf), 'K' => BackendKeyData(buf.read_be_i32_(), buf.read_be_i32_()), + 'n' => NoData, 'R' => read_auth_message(&mut buf), 'S' => ParameterStatus(buf.read_string(), buf.read_string()), + 't' => read_parameter_description(&mut buf), + 'T' => read_row_description(&mut buf), 'Z' => ReadyForQuery(buf.read_u8_()), ident => fail!("Unknown message identifier `%c`", ident) }; @@ -158,3 +180,33 @@ fn read_auth_message(buf: &mut MemReader) -> BackendMessage { val => fail!("Unknown Authentication identifier `%?`", val) } } + +fn read_parameter_description(buf: &mut MemReader) -> BackendMessage { + let len = buf.read_be_i16_() as uint; + let mut types = vec::with_capacity(len); + + do len.times() { + types.push(buf.read_be_i32_()); + } + + ParameterDescription(types) +} + +fn read_row_description(buf: &mut MemReader) -> BackendMessage { + let len = buf.read_be_i16_() as uint; + let mut types = vec::with_capacity(len); + + do len.times() { + types.push(RowDescriptionEntry { + name: buf.read_string(), + table_oid: buf.read_be_i32_(), + column_id: buf.read_be_i16_(), + type_oid: buf.read_be_i32_(), + type_size: buf.read_be_i16_(), + type_modifier: buf.read_be_i32_(), + format: buf.read_be_i16_() + }); + } + + RowDescription(types) +} diff --git a/src/test.rs b/src/test.rs index a6cf2eff..29649a90 100644 --- a/src/test.rs +++ b/src/test.rs @@ -7,5 +7,5 @@ fn test_connect() { let conn = PostgresConnection::connect("postgres://127.0.0.1:54322", "postgres"); - conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)"); + let stmt = conn.prepare("CREATE TABLE foo (id BIGINT PRIMARY KEY)"); }