use super::{Column, ColumnType, HostType, SqlError, Value, write};
use crate::lir::SqlNames;
use crate::unit::ADDRESS_BASE;
use crate::vocab::{SignClause, SignPosition};
use zarch::ebcdic::CodePage;
const HEADER: usize = 16;
const SQLVAR: usize = 44;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct Invalid(pub u8);
pub(super) const INVALID: SqlError = SqlError { code: -804, state: "07002" };
fn halfword(mem: &[u8], at: usize) -> i16 {
i16::from_be_bytes([mem[at], mem[at + 1]])
}
fn fullword(mem: &[u8], at: usize) -> u32 {
u32::from_be_bytes([mem[at], mem[at + 1], mem[at + 2], mem[at + 3]])
}
fn sqltype(ty: &ColumnType) -> Option<(i16, [u8; 2])> {
let length = |n: u16| n.to_be_bytes();
Some(match *ty {
ColumnType::Char(n) => (452, length(n)),
ColumnType::VarChar(n) => (448, length(n)),
ColumnType::Graphic(n) => (468, length(n)),
ColumnType::VarGraphic(n) => (464, length(n)),
ColumnType::SmallInt => (500, length(2)),
ColumnType::Integer => (496, length(4)),
ColumnType::BigInt => (492, length(8)),
ColumnType::Decimal { precision, scale } => (484, [precision, scale]),
ColumnType::Real => (480, length(4)),
ColumnType::Double => (480, length(8)),
ColumnType::Date => (384, length(10)),
ColumnType::Time => (388, length(8)),
ColumnType::Timestamp(0) => (392, length(19)),
ColumnType::Timestamp(p) => (392, length(20 + u16::from(p))),
ColumnType::Binary(n) => (912, length(n)),
ColumnType::VarBinary(n) => (908, length(n)),
ColumnType::Other(_) => return None,
})
}
pub(super) fn describe(mem: &mut [u8], at: usize, columns: Option<&[Column]>, names: SqlNames, page: &'static CodePage) -> Result<usize, Result<Invalid, String>> {
if at + HEADER > mem.len() {
return Err(Ok(Invalid(7)));
}
let sqln = halfword(mem, at + 12);
let columns = columns.unwrap_or_default();
let sqld = i16::try_from(columns.len()).map_err(|_| Ok(Invalid(7)))?;
if sqln < 0 || at + HEADER + SQLVAR * sqln as usize > mem.len() {
return Err(Ok(Invalid(7)));
}
let put = |mem: &mut [u8], from: usize, len: usize, value: Value, ty: HostType| {
let _ = write(&value, &mut mem[from..from + len], &ty, page);
};
put(mem, at, 8, Value::Char("SQLDA".into()), HostType::Char(8));
put(mem, at + 8, 4, Value::Int(i64::from(sqln) * SQLVAR as i64 + HEADER as i64), HostType::Integer { signed: true });
put(mem, at + 14, 2, Value::Int(sqld.into()), HostType::SmallInt { signed: true });
if sqld == 0 || sqld > sqln {
return Ok(HEADER);
}
for (k, column) in columns.iter().enumerate() {
let Some((code, length)) = sqltype(&column.ty) else {
let ColumnType::Other(name) = &column.ty else { unreachable!("sqltype describes every other type") };
return Err(Err(format!("column {} is {name}, which has no Db2 type", column.name)));
};
let var = at + HEADER + SQLVAR * k;
mem[var..var + 2].copy_from_slice(&(code + i16::from(column.nullable)).to_be_bytes());
mem[var + 2..var + 4].copy_from_slice(&length);
let ccsid = match column.ty {
ColumnType::Char(_) | ColumnType::VarChar(_) | ColumnType::Date | ColumnType::Time | ColumnType::Timestamp(_) => Some(page.ccsid),
ColumnType::Graphic(_) | ColumnType::VarGraphic(_) => page.dbcs_ccsid(),
_ => None,
};
mem[var + 4..var + 8].copy_from_slice(&u32::from(ccsid.unwrap_or(0)).to_be_bytes());
mem[var + 8..var + 12].fill(0);
let name = match names {
SqlNames::Labels => String::new(),
SqlNames::Names | SqlNames::Any => column.name.chars().take(30).collect(),
};
put(mem, var + 12, 2, Value::Int(name.chars().count() as i64), HostType::SmallInt { signed: true });
put(mem, var + 14, 30, Value::Char(name), HostType::Char(30));
}
Ok(HEADER + SQLVAR * columns.len())
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct Var {
pub offset: usize,
pub len: usize,
pub ty: HostType,
pub indicator: Option<usize>,
}
fn host_type(code: i16, sqllen: [u8; 2]) -> Option<(HostType, usize)> {
let n = usize::from(u16::from_be_bytes(sqllen));
let (precision, scale) = (u32::from(sqllen[0]), u32::from(sqllen[1]));
let n32 = n as u32;
Some(match code {
452 | 384 | 388 | 392 => (HostType::Char(n32), n),
448 | 456 => (HostType::VarChar(n32), 2 + n),
468 => (HostType::Graphic(n32), 2 * n),
464 | 472 => (HostType::VarGraphic(n32), 2 + 2 * n),
500 => (HostType::SmallInt { signed: true }, 2),
496 => (HostType::Integer { signed: true }, 4),
492 => (HostType::BigInt { signed: true }, 8),
480 if n == 4 => (HostType::Real, 4),
480 if n == 8 => (HostType::Double, 8),
484 if (1..=31).contains(&precision) && scale <= precision => (HostType::Decimal { digits: precision, scale, signed: true }, precision as usize / 2 + 1),
504 if (1..=31).contains(&precision) && scale <= precision => {
let sign = Some(SignClause { position: SignPosition::Leading, separate: true });
(HostType::Zoned { digits: precision, scale, signed: true, sign }, precision as usize + 1)
}
_ => return None,
})
}
fn addressed(address: u32, len: usize, mem_len: usize) -> Option<usize> {
let offset = address.checked_sub(ADDRESS_BASE)? as usize;
(address != 0 && offset + len <= mem_len).then_some(offset)
}
pub(super) fn vars(mem: &[u8], at: usize, input: bool) -> Result<Vec<Var>, Invalid> {
if at + HEADER > mem.len() {
return Err(Invalid(7));
}
let (sqldabc, sqln, sqld) = (fullword(mem, at + 8) as i32, halfword(mem, at + 12), halfword(mem, at + 14));
if sqln < 0 || sqld < 0 {
return Err(Invalid(7));
}
if i64::from(sqldabc) < i64::from(sqln) * SQLVAR as i64 + HEADER as i64 {
return Err(Invalid(14));
}
if sqld > sqln {
return Err(Invalid(11));
}
if at + HEADER + SQLVAR * sqld as usize > mem.len() {
return Err(Invalid(7));
}
let (unknown, bad_pointer) = if input { (8, 12) } else { (16, 13) };
(0..sqld as usize)
.map(|k| {
let var = at + HEADER + SQLVAR * k;
let code = halfword(mem, var);
let (ty, len) = host_type(code & !1, [mem[var + 2], mem[var + 3]]).ok_or(Invalid(unknown))?;
let offset = addressed(fullword(mem, var + 4), len, mem.len()).ok_or(Invalid(bad_pointer))?;
let indicator = match code & 1 {
0 => None,
_ => Some(addressed(fullword(mem, var + 8), 2, mem.len()).ok_or(Invalid(bad_pointer))?),
};
Ok(Var { offset, len, ty, indicator })
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn page() -> &'static CodePage {
CodePage::by_ccsid(1140).expect("1140")
}
fn sqlda(sqln: i16) -> Vec<u8> {
let mut mem = vec![0u8; 16 + 44 * sqln.max(0) as usize + 64];
mem[12..14].copy_from_slice(&sqln.to_be_bytes());
mem
}
#[test]
fn describe_fills_the_header_and_each_sqlvar() {
let columns = [
Column { name: "NAME".into(), ty: ColumnType::Char(10), nullable: false },
Column { name: "AMT".into(), ty: ColumnType::Decimal { precision: 7, scale: 2 }, nullable: true },
];
let mut mem = sqlda(3);
describe(&mut mem, 0, Some(&columns), SqlNames::Names, page()).unwrap();
assert_eq!(&mem[0..8], &[0xE2, 0xD8, 0xD3, 0xC4, 0xC1, 0x40, 0x40, 0x40]);
assert_eq!((fullword(&mem, 8), halfword(&mem, 12), halfword(&mem, 14)), (16 + 44 * 3, 3, 2));
assert_eq!((halfword(&mem, 16), &mem[18..20], fullword(&mem, 20)), (452, &[0, 10][..], 1140));
assert_eq!((halfword(&mem, 28), &mem[30..34]), (4, &[0xD5, 0xC1, 0xD4, 0xC5][..]));
assert_eq!((halfword(&mem, 60), &mem[62..64], fullword(&mem, 64)), (485, &[7, 2][..], 0));
}
#[test]
fn too_few_sqlvars_or_no_query_set_sqld_alone() {
let columns = [Column { name: "A".into(), ty: ColumnType::Integer, nullable: false }];
let mut mem = sqlda(0);
describe(&mut mem, 0, Some(&columns), SqlNames::Names, page()).unwrap();
assert_eq!((halfword(&mem, 14), fullword(&mem, 8)), (1, 16));
let mut mem = sqlda(2);
describe(&mut mem, 0, None, SqlNames::Names, page()).unwrap();
assert_eq!((halfword(&mem, 14), halfword(&mem, 16)), (0, 0));
let other = [Column { name: "B".into(), ty: ColumnType::Other("PostgreSQL type OID 16".into()), nullable: true }];
assert!(matches!(describe(&mut sqlda(1), 0, Some(&other), SqlNames::Names, page()), Err(Err(why)) if why.contains("OID 16")));
}
#[test]
fn using_descriptor_reads_each_sqlvar_s_type_and_addresses() {
let mut mem = sqlda(2);
mem[8..12].copy_from_slice(&(16 + 88u32).to_be_bytes());
mem[14..16].copy_from_slice(&2i16.to_be_bytes());
mem[16..18].copy_from_slice(&453i16.to_be_bytes());
mem[18..20].copy_from_slice(&8u16.to_be_bytes());
mem[20..24].copy_from_slice(&(ADDRESS_BASE + 110).to_be_bytes());
mem[24..28].copy_from_slice(&(ADDRESS_BASE + 118).to_be_bytes());
mem[60..62].copy_from_slice(&484i16.to_be_bytes());
mem[62..64].copy_from_slice(&[5, 2]);
mem[64..68].copy_from_slice(&(ADDRESS_BASE + 120).to_be_bytes());
let got = vars(&mem, 0, true).unwrap();
assert_eq!(got[0], Var { offset: 110, len: 8, ty: HostType::Char(8), indicator: Some(118) });
assert_eq!(got[1], Var { offset: 120, len: 3, ty: HostType::Decimal { digits: 5, scale: 2, signed: true }, indicator: None });
mem[64..68].fill(0);
assert_eq!(vars(&mem, 0, true), Err(Invalid(12)));
assert_eq!(vars(&mem, 0, false), Err(Invalid(13)));
mem[60..62].copy_from_slice(&999i16.to_be_bytes());
assert_eq!(vars(&mem, 0, true), Err(Invalid(8)));
mem[14..16].copy_from_slice(&3i16.to_be_bytes());
assert_eq!(vars(&mem, 0, true), Err(Invalid(11)));
mem[8..12].copy_from_slice(&16u32.to_be_bytes());
assert_eq!(vars(&mem, 0, true), Err(Invalid(14)));
}
}