diff --git a/README.md b/README.md index 943a615a..6d3b83d3 100644 --- a/README.md +++ b/README.md @@ -256,6 +256,10 @@ types. The driver currently supports the following conversions: types::range::Range<Timespec> TSRANGE, TSTZRANGE + + types::array::ArrayBase<i32> + INT4[], INT4[][], ... + diff --git a/test.rs b/test.rs index e9d2e03b..22453d81 100644 --- a/test.rs +++ b/test.rs @@ -32,6 +32,7 @@ use lib::error::{DbError, QueryCanceled, InvalidCatalogName}; use lib::types::{ToSql, FromSql, PgInt4, PgVarchar}; +use lib::types::array::{ArrayBase}; use lib::types::range::{Range, Inclusive, Exclusive, RangeBound}; use lib::pool::PostgresConnectionPool; @@ -420,6 +421,19 @@ fn test_tstzrange_params() { test_timespec_range_params("TSTZRANGE"); } +#[test] +fn test_int4array_params() { + test_type("INT4[]", + [(Some(ArrayBase::from_vec(~[Some(0i32), Some(1), None], 1)), + "'{0,1,NULL}'"), + (None, "NULL")]); + let mut a = ArrayBase::from_vec(~[Some(0i32), Some(1)], 0); + a.wrap(-1); + a.push_move(ArrayBase::from_vec(~[None, Some(3)], 0)); + test_type("INT4[][]", + [(Some(a), "'[-1:0][0:1]={{0,1},{NULL,3}}'")]); +} + fn test_nan_param(sql_type: &str) { let conn = PostgresConnection::connect("postgres://postgres@localhost", &NoSsl); let stmt = conn.prepare("SELECT 'NaN'::" + sql_type); diff --git a/types/array.rs b/types/array.rs new file mode 100644 index 00000000..5cb5772d --- /dev/null +++ b/types/array.rs @@ -0,0 +1,282 @@ +//! Multi-dimensional arrays with per-dimension specifiable lower bounds + +use std::cast; +use std::vec::VecIterator; + +#[deriving(Eq, Clone)] +pub struct DimensionInfo { + len: uint, + lower_bound: int, +} + +pub trait Array { + fn get_dimension_info<'a>(&'a self) -> &'a [DimensionInfo]; + fn slice<'a>(&'a self, idx: int) -> ArraySlice<'a, T>; + fn get<'a>(&'a self, idx: int) -> &'a T; +} + +pub trait MutableArray : Array { + fn slice_mut<'a>(&'a mut self, idx: int) -> MutArraySlice<'a, T> { + MutArraySlice { slice: self.slice(idx) } + } + + fn get_mut<'a>(&'a mut self, idx: int) -> &'a mut T { + unsafe { cast::transmute_mut(self.get(idx)) } + } +} + +trait InternalArray : Array { + fn shift_idx(&self, idx: int) -> uint { + let shifted_idx = idx - self.get_dimension_info()[0].lower_bound; + assert!(shifted_idx >= 0, "Out of bounds array access"); + shifted_idx as uint + } + + fn raw_get<'a>(&'a self, idx: uint, size: uint) -> &'a T; +} + +#[deriving(Eq, Clone)] +pub struct ArrayBase { + priv info: ~[DimensionInfo], + priv data: ~[T], +} + +impl ArrayBase { + pub fn from_raw(data: ~[T], info: ~[DimensionInfo]) + -> ArrayBase { + assert!(!info.is_empty(), "Cannot create a 0x0 array"); + assert!(data.len() == info.iter().fold(1, |acc, i| acc * i.len), + "Size mismatch"); + ArrayBase { + info: info, + data: data, + } + } + + pub fn from_vec(data: ~[T], lower_bound: int) -> ArrayBase { + ArrayBase { + info: ~[DimensionInfo { + len: data.len(), + lower_bound: lower_bound + }], + data: data + } + } + + pub fn wrap(&mut self, lower_bound: int) { + self.info.unshift(DimensionInfo { + len: 1, + lower_bound: lower_bound + }) + } + + pub fn push_move(&mut self, other: ArrayBase) { + assert!(self.info.len() - 1 == other.info.len(), + "Cannot append differently shaped arrays"); + for (info1, info2) in self.info.iter().skip(1).zip(other.info.iter()) { + assert!(info1 == info2, "Cannot append differently shaped arrays"); + } + self.info[0].len += 1; + self.data.push_all_move(other.data); + } + + pub fn values<'a>(&'a self) -> VecIterator<'a, T> { + self.data.iter() + } +} + +impl Array for ArrayBase { + fn get_dimension_info<'a>(&'a self) -> &'a [DimensionInfo] { + self.info.as_slice() + } + + fn slice<'a>(&'a self, idx: int) -> ArraySlice<'a, T> { + ArraySlice { + parent: BaseParent(self), + idx: self.shift_idx(idx) + } + } + + fn get<'a>(&'a self, idx: int) -> &'a T { + assert!(self.info.len() == 1, + "Attempted to get from a multi-dimensional array"); + self.raw_get(self.shift_idx(idx), 1) + } +} + +impl MutableArray for ArrayBase {} + +impl InternalArray for ArrayBase { + fn raw_get<'a>(&'a self, idx: uint, _size: uint) -> &'a T { + &self.data[idx] + } +} + +enum ArrayParent<'parent, T> { + SliceParent(&'parent ArraySlice<'parent, T>), + BaseParent(&'parent ArrayBase), +} + +pub struct ArraySlice<'parent, T> { + priv parent: ArrayParent<'parent, T>, + priv idx: uint, +} + +impl<'parent, T> Array for ArraySlice<'parent, T> { + fn get_dimension_info<'a>(&'a self) -> &'a [DimensionInfo] { + let info = match self.parent { + SliceParent(p) => p.get_dimension_info(), + BaseParent(p) => p.get_dimension_info() + }; + info.slice_from(1) + } + + fn slice<'a>(&'a self, idx: int) -> ArraySlice<'a, T> { + ArraySlice { + parent: SliceParent(self), + idx: self.shift_idx(idx) + } + } + + fn get<'a>(&'a self, idx: int) -> &'a T { + assert!(self.get_dimension_info().len() == 1, + "Attempted to get from a multi-dimensional array"); + self.raw_get(self.shift_idx(idx), 1) + } +} + +impl<'parent, T> InternalArray for ArraySlice<'parent, T> { + fn raw_get<'a>(&'a self, idx: uint, size: uint) -> &'a T { + let size = size * self.get_dimension_info()[0].len; + let idx = size * self.idx + idx; + match self.parent { + SliceParent(p) => p.raw_get(idx, size), + BaseParent(p) => p.raw_get(idx, size) + } + } +} + +pub struct MutArraySlice<'parent, T> { + priv slice: ArraySlice<'parent, T> +} + +impl<'parent, T> Array for MutArraySlice<'parent, T> { + fn get_dimension_info<'a>(&'a self) -> &'a [DimensionInfo] { + self.slice.get_dimension_info() + } + + fn slice<'a>(&'a self, idx: int) -> ArraySlice<'a, T> { + self.slice.slice(idx) + } + + fn get<'a>(&'a self, idx: int) -> &'a T { + self.slice.get(idx) + } +} + +impl<'parent, T> MutableArray for MutArraySlice<'parent, T> {} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_from_vec() { + let a = ArrayBase::from_vec(~[0, 1, 2], -1); + assert_eq!([DimensionInfo { len: 3, lower_bound: -1 }], + a.get_dimension_info()); + assert_eq!(&0, a.get(-1)); + assert_eq!(&1, a.get(0)); + assert_eq!(&2, a.get(1)); + } + + #[test] + #[should_fail] + fn test_get_2d_fail() { + let mut a = ArrayBase::from_vec(~[0, 1, 2], -1); + a.wrap(1); + a.get(1); + } + + #[test] + #[should_fail] + fn test_2d_slice_range_fail() { + let mut a = ArrayBase::from_vec(~[0, 1, 2], -1); + a.wrap(1); + a.slice(0); + } + + #[test] + fn test_2d_slice_get() { + let mut a = ArrayBase::from_vec(~[0, 1, 2], -1); + a.wrap(1); + let s = a.slice(1); + assert_eq!(&0, s.get(-1)); + assert_eq!(&1, s.get(0)); + assert_eq!(&2, s.get(1)); + } + + #[test] + #[should_fail] + fn test_push_move_wrong_lower_bound() { + let mut a = ArrayBase::from_vec(~[1], -1); + a.push_move(ArrayBase::from_vec(~[2], 0)); + } + + #[test] + #[should_fail] + fn test_push_move_wrong_dims() { + let mut a = ArrayBase::from_vec(~[1], -1); + a.wrap(1); + a.push_move(ArrayBase::from_vec(~[1, 2], -1)); + } + + #[test] + #[should_fail] + fn test_push_move_wrong_dim_count() { + let mut a = ArrayBase::from_vec(~[1], -1); + a.wrap(1); + let mut b = ArrayBase::from_vec(~[2], -1); + b.wrap(1); + a.push_move(b); + } + + #[test] + fn test_push_move_ok() { + let mut a = ArrayBase::from_vec(~[1, 2], 0); + a.wrap(0); + a.push_move(ArrayBase::from_vec(~[3, 4], 0)); + let s = a.slice(0); + assert_eq!(&1, s.get(0)); + assert_eq!(&2, s.get(1)); + let s = a.slice(1); + assert_eq!(&3, s.get(0)); + assert_eq!(&4, s.get(1)); + } + + #[test] + fn test_3d() { + let mut a = ArrayBase::from_vec(~[0, 1], 0); + a.wrap(0); + a.push_move(ArrayBase::from_vec(~[2, 3], 0)); + a.wrap(0); + let mut b = ArrayBase::from_vec(~[4, 5], 0); + b.wrap(0); + b.push_move(ArrayBase::from_vec(~[6, 7], 0)); + a.push_move(b); + let s1 = a.slice(0); + let s2 = s1.slice(0); + assert_eq!(&0, s2.get(0)); + assert_eq!(&1, s2.get(1)); + let s2 = s1.slice(1); + assert_eq!(&2, s2.get(0)); + assert_eq!(&3, s2.get(1)); + let s1 = a.slice(1); + let s2 = s1.slice(0); + assert_eq!(&4, s2.get(0)); + assert_eq!(&5, s2.get(1)); + let s2 = s1.slice(1); + assert_eq!(&6, s2.get(0)); + assert_eq!(&7, s2.get(1)); + } +} diff --git a/types/mod.rs b/types/mod.rs index 650e6ad0..a7686d81 100644 --- a/types/mod.rs +++ b/types/mod.rs @@ -9,9 +9,12 @@ use extra::uuid::Uuid; use std::io::Decorator; use std::io::mem::{MemWriter, BufReader}; use std::str; +use std::vec; +use self::array::{Array, ArrayBase, DimensionInfo}; use self::range::{RangeBound, Inclusive, Exclusive, Range}; +pub mod array; pub mod range; /// A Postgres OID @@ -28,6 +31,7 @@ static TEXTOID: Oid = 25; static JSONOID: Oid = 114; static FLOAT4OID: Oid = 700; static FLOAT8OID: Oid = 701; +static INT4ARRAYOID: Oid = 1007; static BPCHAROID: Oid = 1042; static VARCHAROID: Oid = 1043; static TIMESTAMPOID: Oid = 1114; @@ -73,6 +77,8 @@ pub enum PostgresType { PgFloat4, /// FLOAT8/DOUBLE PRECISION PgFloat8, + /// INT4[] + PgInt4Array, /// TIMESTAMP PgTimestamp, /// TIMESTAMP WITH TIME ZONE @@ -109,6 +115,7 @@ impl PostgresType { JSONOID => PgJson, FLOAT4OID => PgFloat4, FLOAT8OID => PgFloat8, + INT4ARRAYOID => PgInt4Array, TIMESTAMPOID => PgTimestamp, TIMESTAMPZOID => PgTimestampZ, BPCHAROID => PgCharN, @@ -322,6 +329,42 @@ from_option_impl!(Range) from_range_impl!(PgTsRange | PgTstzRange, Timespec) from_option_impl!(Range) +macro_rules! from_array_impl( + ($($oid:ident)|+, $t:ty) => ( + from_map_impl!($($oid)|+, ArrayBase>, |buf| { + let mut rdr = BufReader::new(buf.as_slice()); + + let ndim = rdr.read_be_i32() as uint; + let _has_null = rdr.read_be_i32() == 1; + let _element_type: Oid = rdr.read_be_i32(); + + let mut dim_info = vec::with_capacity(ndim); + for _ in range(0, ndim) { + dim_info.push(DimensionInfo { + len: rdr.read_be_i32() as uint, + lower_bound: rdr.read_be_i32() as int + }); + } + let nele = dim_info.iter().fold(1, |acc, info| acc * info.len); + + let mut elements = vec::with_capacity(nele); + for _ in range(0, nele) { + let len = rdr.read_be_i32(); + if len < 0 { + elements.push(None); + } else { + elements.push(Some(RawFromSql::raw_from_sql(&mut rdr))); + } + } + + ArrayBase::from_raw(elements, dim_info) + }) + ) +) + +from_array_impl!(PgInt4Array, i32) +from_option_impl!(ArrayBase>) + /// A trait for types that can be converted into Postgres values pub trait ToSql { /// Converts the value of `self` into a format appropriate for the Postgres @@ -547,3 +590,38 @@ to_option_impl!(PgInt8Range, Range) to_range_impl!(PgTsRange | PgTstzRange, Timespec, 8) to_option_impl!(PgTsRange | PgTstzRange, Range) + +macro_rules! to_array_impl( + ($($oid:ident)|+, $base_oid:ident, $t:ty, $size:expr) => ( + impl ToSql for ArrayBase> { + fn to_sql(&self, ty: PostgresType) -> (Format, Option<~[u8]>) { + check_types!($($oid)|+, ty) + let mut buf = MemWriter::new(); + + buf.write_be_i32(self.get_dimension_info().len() as i32); + buf.write_be_i32(1); + buf.write_be_i32($base_oid); + + for info in self.get_dimension_info().iter() { + buf.write_be_i32(info.len as i32); + buf.write_be_i32(info.lower_bound as i32); + } + + for v in self.values() { + match *v { + Some(ref val) => { + buf.write_be_i32($size); + val.raw_to_sql(&mut buf); + } + None => buf.write_be_i32(-1) + } + } + + (Binary, Some(buf.inner())) + } + } + ) +) + +to_array_impl!(PgInt4Array, INT4OID, i32, 4) +to_option_impl!(PgInt4Array, ArrayBase>)