use alloc::format;
use alloc::string::ToString;
use alloc::vec::Vec;
use spg_sql::ast::BinOp;
use spg_storage::Value;
use super::EvalError;
fn inet_arg_text(v: &Value<'_>) -> Option<alloc::string::String> {
match v {
Value::Text(s) => Some(s.to_string()),
Value::Inet { family, bits, addr } | Value::Cidr { family, bits, addr } => {
Some(crate::conversions::format_inet(*family, *bits, addr))
}
_ => None,
}
}
pub(super) fn inet_host(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("host() takes one TEXT arg, got {} args", args.len()),
});
}
};
let host = s.split('/').next().unwrap_or("").to_string();
Ok(Value::text(host))
}
pub(super) fn inet_network(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("network() takes one TEXT arg, got {} args", args.len()),
});
}
};
let mut split = s.splitn(2, '/');
let host = split.next().unwrap_or("").to_string();
let mask: u32 = split.next().and_then(|m| m.parse().ok()).unwrap_or(32);
if !host.contains('.') {
return Ok(Value::text(s));
}
let Some(addr) = ipv4_to_u32(&host) else {
return Ok(Value::text(s));
};
let mask = mask.min(32);
let masked = if mask == 0 {
0
} else {
addr & (u32::MAX << (32 - mask))
};
Ok(Value::text(alloc::format!(
"{}.{}.{}.{}/{mask}",
(masked >> 24) & 0xFF,
(masked >> 16) & 0xFF,
(masked >> 8) & 0xFF,
masked & 0xFF
)))
}
fn ipv4_to_u32(host: &str) -> Option<u32> {
let octets: Vec<&str> = host.split('.').collect();
if octets.len() != 4 {
return None;
}
let mut addr: u32 = 0;
for oct in &octets {
let b: u32 = oct.parse().ok()?;
if b > 255 {
return None;
}
addr = (addr << 8) | b;
}
Some(addr)
}
pub(super) fn inet_family(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("family() takes one TEXT arg, got {} args", args.len()),
});
}
};
let host = s.split('/').next().unwrap_or("");
if host.contains(':') {
Ok(Value::Int(6))
} else {
Ok(Value::Int(4))
}
}
pub(super) fn inet_netmask(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("netmask() takes one TEXT arg, got {} args", args.len()),
});
}
};
let mut split = s.splitn(2, '/');
let host = split.next().unwrap_or("");
let mask: u32 = split.next().and_then(|m| m.parse().ok()).unwrap_or(32);
if host.contains(':') {
return Ok(Value::text(s));
}
let mask = mask.min(32);
let bits: u32 = if mask == 0 {
0
} else {
u32::MAX << (32 - mask)
};
Ok(Value::text(alloc::format!(
"{}.{}.{}.{}",
(bits >> 24) & 0xFF,
(bits >> 16) & 0xFF,
(bits >> 8) & 0xFF,
bits & 0xFF
)))
}
pub(super) fn inet_hostmask(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("hostmask() takes one TEXT arg, got {} args", args.len()),
});
}
};
let mut split = s.splitn(2, '/');
let host = split.next().unwrap_or("");
let mask: u32 = split.next().and_then(|m| m.parse().ok()).unwrap_or(32);
if host.contains(':') {
return Ok(Value::text(s));
}
let mask = mask.min(32);
let bits: u32 = if mask == 0 {
u32::MAX
} else {
!(u32::MAX << (32 - mask))
};
Ok(Value::text(alloc::format!(
"{}.{}.{}.{}",
(bits >> 24) & 0xFF,
(bits >> 16) & 0xFF,
(bits >> 8) & 0xFF,
bits & 0xFF
)))
}
pub(super) fn inet_broadcast(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("broadcast() takes one TEXT arg, got {} args", args.len()),
});
}
};
let mut split = s.splitn(2, '/');
let host = split.next().unwrap_or("");
let mask: u32 = split.next().and_then(|m| m.parse().ok()).unwrap_or(32);
if host.contains(':') {
return Ok(Value::text(s));
}
let octets: Vec<u32> = host.split('.').filter_map(|o| o.parse().ok()).collect();
if octets.len() != 4 {
return Ok(Value::text(s));
}
let addr = (octets[0] << 24) | (octets[1] << 16) | (octets[2] << 8) | octets[3];
let mask = mask.min(32);
let host_bits: u32 = if mask == 0 {
u32::MAX
} else {
!(u32::MAX << (32 - mask))
};
let bcast = addr | host_bits;
Ok(Value::text(alloc::format!(
"{}.{}.{}.{}/{mask}",
(bcast >> 24) & 0xFF,
(bcast >> 16) & 0xFF,
(bcast >> 8) & 0xFF,
bcast & 0xFF
)))
}
pub(super) fn inet_merge(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let unpack = |v: &Value<'_>| -> Option<(u8, u8, [u8; 16])> {
match v {
Value::Inet { family, bits, addr } | Value::Cidr { family, bits, addr } => {
Some((*family, *bits, *addr))
}
Value::Text(s) => crate::conversions::parse_inet_text(s),
_ => None,
}
};
if args.iter().any(|v| matches!(v, Value::Null)) {
return Ok(Value::Null);
}
let (Some((fa, ba, aa)), Some((fb, bb, ab))) =
(args.first().and_then(unpack), args.get(1).and_then(unpack))
else {
return Err(EvalError::TypeMismatch {
detail: "inet_merge() takes 2 inet/cidr args".into(),
});
};
if fa != fb {
return Err(EvalError::TypeMismatch {
detail: "cannot merge addresses from different families".into(),
});
}
let nbits = u16::from(ba.min(bb));
let mut common: u16 = 0;
'outer: for byte in 0..16usize {
for bit in 0..8u16 {
let pos = (byte as u16) * 8 + bit;
if pos >= nbits {
break 'outer;
}
let mask = 0x80u8 >> bit;
if (aa[byte] & mask) != (ab[byte] & mask) {
break 'outer;
}
common = pos + 1;
}
}
let mut addr = [0u8; 16];
for byte in 0..16usize {
let bit_base = (byte as u16) * 8;
let keep = common.saturating_sub(bit_base).min(8) as u8;
let mask: u8 = if keep == 0 { 0 } else { 0xffu8 << (8 - keep) };
addr[byte] = aa[byte] & mask;
}
#[allow(clippy::cast_possible_truncation)]
Ok(Value::Cidr {
family: fa,
bits: common as u8,
addr,
})
}
pub(super) fn macaddr8_set7bit(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let mut out: [u8; 8] = match args {
[Value::Macaddr8(b)] => *b,
[Value::Null] => return Ok(Value::Null),
[Value::Text(s)] => {
let bytes: Vec<u8> = s
.split([':', '-'])
.filter_map(|part| u8::from_str_radix(part, 16).ok())
.collect();
<[u8; 8]>::try_from(bytes.as_slice()).map_err(|_| EvalError::TypeMismatch {
detail: alloc::format!("macaddr8_set7bit(): invalid macaddr8 '{s}'"),
})?
}
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!(
"macaddr8_set7bit() takes one macaddr8 arg, got {} args",
args.len()
),
});
}
};
out[0] |= 0x02;
Ok(Value::Macaddr8(out))
}
pub(super) fn inet_same_family(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let fam = |v: &Value<'_>| -> Option<u8> {
match v {
Value::Inet { family, .. } | Value::Cidr { family, .. } => Some(*family),
Value::Text(s) => Some(if s.split('/').next().unwrap_or("").contains(':') {
6
} else {
4
}),
_ => None,
}
};
if args.iter().any(|v| matches!(v, Value::Null)) {
return Ok(Value::Null);
}
match (args.first().and_then(fam), args.get(1).and_then(fam)) {
(Some(a), Some(b)) => Ok(Value::Bool(a == b)),
_ => Err(EvalError::TypeMismatch {
detail: "inet_same_family() takes 2 inet args".into(),
}),
}
}
pub(super) fn inet_masklen(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Text(s)] => s.clone(),
[Value::Inet { family, bits, addr }] | [Value::Cidr { family, bits, addr }] => {
alloc::borrow::Cow::Owned(crate::conversions::format_inet(*family, *bits, addr))
}
[Value::Null] => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("masklen() takes one TEXT arg, got {} args", args.len()),
});
}
};
let mask: i32 = s
.split_once('/')
.and_then(|(_, m)| m.parse().ok())
.unwrap_or(32);
Ok(Value::Int(mask))
}
pub(super) fn inet_set_masklen(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
if args.iter().any(|v| matches!(v, Value::Null)) {
return Ok(Value::Null);
}
let (family, addr, is_cidr) = match args.first() {
Some(Value::Inet { family, addr, .. }) => (*family, *addr, false),
Some(Value::Cidr { family, addr, .. }) => (*family, *addr, true),
_ => {
return Err(EvalError::TypeMismatch {
detail: "set_masklen() first arg must be inet/cidr".into(),
});
}
};
let n = match args.get(1) {
Some(Value::SmallInt(v)) => i64::from(*v),
Some(Value::Int(v)) => i64::from(*v),
Some(Value::BigInt(v)) => *v,
_ => {
return Err(EvalError::TypeMismatch {
detail: "set_masklen() second arg must be an integer".into(),
});
}
};
let max = if family == 4 { 32 } else { 128 };
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let bits = n.clamp(0, max) as u8;
if is_cidr {
let mut addr = addr;
let nbytes = if family == 4 { 4 } else { 16 };
for byte in 0..nbytes {
let bit_base = (byte as u16) * 8;
let keep = u16::from(bits).saturating_sub(bit_base).min(8) as u8;
let mask: u8 = if keep == 0 { 0 } else { 0xffu8 << (8 - keep) };
addr[byte] &= mask;
}
Ok(Value::Cidr { family, bits, addr })
} else {
Ok(Value::Inet { family, bits, addr })
}
}
pub(super) fn inet_abbrev(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let (family, bits, addr, is_cidr) = match args.first() {
Some(Value::Cidr { family, bits, addr }) => (*family, *bits, *addr, true),
Some(Value::Inet { family, bits, addr }) => (*family, *bits, *addr, false),
Some(Value::Null) => return Ok(Value::Null),
_ => {
return Err(EvalError::TypeMismatch {
detail: "abbrev() arg must be inet/cidr".into(),
});
}
};
if !is_cidr || family != 4 {
return Ok(Value::text(crate::conversions::format_inet(
family, bits, &addr,
)));
}
let sig = ((usize::from(bits) + 7) / 8).max(1);
let parts: alloc::vec::Vec<alloc::string::String> = addr[0..sig]
.iter()
.map(alloc::string::ToString::to_string)
.collect();
Ok(Value::text(alloc::format!("{}/{}", parts.join("."), bits)))
}
struct InetNet {
bytes: [u8; 16],
family_bytes: u8,
prefix_bits: u8,
}
fn parse_inet_text(s: &str) -> Option<InetNet> {
let mut split = s.splitn(2, '/');
let host = split.next()?;
let mask_str = split.next();
if host.contains(':') {
let bytes = parse_ipv6(host)?;
let prefix_bits = match mask_str {
Some(m) => m.parse::<u8>().ok().filter(|&n| n <= 128)?,
None => 128,
};
let mut out = [0u8; 16];
out.copy_from_slice(&bytes);
Some(InetNet {
bytes: out,
family_bytes: 16,
prefix_bits,
})
} else {
let bytes = parse_ipv4(host)?;
let prefix_bits = match mask_str {
Some(m) => m.parse::<u8>().ok().filter(|&n| n <= 32)?,
None => 32,
};
let mut out = [0u8; 16];
out[..4].copy_from_slice(&bytes);
Some(InetNet {
bytes: out,
family_bytes: 4,
prefix_bits,
})
}
}
fn parse_ipv4(s: &str) -> Option<[u8; 4]> {
let parts: Vec<&str> = s.split('.').collect();
if parts.len() != 4 {
return None;
}
let mut out = [0u8; 4];
for (i, p) in parts.iter().enumerate() {
out[i] = p.parse::<u8>().ok()?;
}
Some(out)
}
fn parse_ipv6(s: &str) -> Option<[u8; 16]> {
let (head, tail) = match s.find("::") {
Some(idx) => (&s[..idx], Some(&s[idx + 2..])),
None => (s, None),
};
let head_groups: Vec<&str> = if head.is_empty() {
Vec::new()
} else {
head.split(':').collect()
};
let tail_groups: Vec<&str> = match tail {
Some(t) if !t.is_empty() => t.split(':').collect(),
_ => Vec::new(),
};
let head_len = head_groups.len();
let tail_len = tail_groups.len();
if tail.is_none() {
if head_len != 8 {
return None;
}
} else if head_len + tail_len > 7 {
return None;
}
let mut words = [0u16; 8];
for (i, g) in head_groups.iter().enumerate() {
words[i] = u16::from_str_radix(g, 16).ok()?;
}
let tail_start = 8 - tail_len;
for (i, g) in tail_groups.iter().enumerate() {
words[tail_start + i] = u16::from_str_radix(g, 16).ok()?;
}
let mut out = [0u8; 16];
for (i, w) in words.iter().enumerate() {
out[i * 2] = (w >> 8) as u8;
out[i * 2 + 1] = (w & 0xff) as u8;
}
Some(out)
}
fn network_prefix_eq(a: &InetNet, b: &InetNet, prefix_bits: u8) -> bool {
let full_bytes = (prefix_bits / 8) as usize;
if a.bytes[..full_bytes] != b.bytes[..full_bytes] {
return false;
}
let extra = prefix_bits % 8;
if extra == 0 {
return true;
}
let mask: u8 = 0xff << (8 - extra);
(a.bytes[full_bytes] & mask) == (b.bytes[full_bytes] & mask)
}
fn inet_contained_eq(a: &InetNet, b: &InetNet) -> bool {
if a.family_bytes != b.family_bytes {
return false;
}
if a.prefix_bits < b.prefix_bits {
return false;
}
network_prefix_eq(a, b, b.prefix_bits)
}
fn inet_networks_equal(a: &InetNet, b: &InetNet) -> bool {
if a.family_bytes != b.family_bytes {
return false;
}
if a.prefix_bits != b.prefix_bits {
return false;
}
network_prefix_eq(a, b, a.prefix_bits)
}
pub(super) fn inet_op_bool_result(
op: BinOp,
l: &Value,
r: &Value,
) -> Result<Value<'static>, EvalError> {
if matches!(l, Value::Null) || matches!(r, Value::Null) {
return Ok(Value::Null);
}
let to_inet = |v: &Value| -> Result<InetNet, EvalError> {
match v {
Value::Inet { family, bits, addr } | Value::Cidr { family, bits, addr } => {
Ok(InetNet {
bytes: *addr,
family_bytes: *family,
prefix_bits: *bits,
})
}
Value::Text(s) => parse_inet_text(s).ok_or_else(|| EvalError::TypeMismatch {
detail: format!("invalid inet text: {s:?}"),
}),
_ => Err(EvalError::TypeMismatch {
detail: format!(
"inet operator requires INET/CIDR/TEXT operands, got {}",
crate::conversions::pg_type_name_for_error_opt(v.data_type())
),
}),
}
};
let a = to_inet(l)?;
let b = to_inet(r)?;
let result = match op {
BinOp::InetContainedByEq => inet_contained_eq(&a, &b),
BinOp::InetContainedBy => inet_contained_eq(&a, &b) && !inet_networks_equal(&a, &b),
BinOp::InetContainsEq => inet_contained_eq(&b, &a),
BinOp::InetContains => inet_contained_eq(&b, &a) && !inet_networks_equal(&a, &b),
BinOp::InetOverlap => inet_contained_eq(&a, &b) || inet_contained_eq(&b, &a),
_ => unreachable!("inet_op_bool_result called with non-inet op"),
};
Ok(Value::Bool(result))
}
pub(super) fn mysql_inet_aton(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Null] => return Ok(Value::Null),
[Value::Text(s)] => s,
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("inet_aton() takes one TEXT arg, got {} args", args.len()),
});
}
};
match parse_ipv4(s) {
Some(b) => Ok(Value::BigInt(
(i64::from(b[0]) << 24)
| (i64::from(b[1]) << 16)
| (i64::from(b[2]) << 8)
| i64::from(b[3]),
)),
None => Ok(Value::Null),
}
}
pub(super) fn mysql_inet_ntoa(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let n = match args {
[Value::Null] => return Ok(Value::Null),
[Value::Int(n)] => i64::from(*n),
[Value::SmallInt(n)] => i64::from(*n),
[Value::BigInt(n)] => *n,
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!(
"inet_ntoa() takes one integer arg, got {}",
args.first().map_or_else(
|| alloc::string::String::from("no arguments"),
|a| crate::conversions::pg_type_name_for_error_opt(a.data_type()),
)
),
});
}
};
if !(0..=0xFFFF_FFFF).contains(&n) {
return Ok(Value::Null);
}
let n = n as u32;
Ok(Value::text(alloc::format!(
"{}.{}.{}.{}",
(n >> 24) & 0xFF,
(n >> 16) & 0xFF,
(n >> 8) & 0xFF,
n & 0xFF
)))
}
fn format_ipv6(bytes: &[u8; 16]) -> alloc::string::String {
let mut groups = [0u16; 8];
for (i, g) in groups.iter_mut().enumerate() {
*g = (u16::from(bytes[i * 2]) << 8) | u16::from(bytes[i * 2 + 1]);
}
let (mut best_start, mut best_len) = (usize::MAX, 0usize);
let mut i = 0;
while i < 8 {
if groups[i] == 0 {
let start = i;
while i < 8 && groups[i] == 0 {
i += 1;
}
let len = i - start;
if len > best_len {
best_start = start;
best_len = len;
}
} else {
i += 1;
}
}
let mut out = alloc::string::String::new();
if best_len >= 2 {
for (idx, g) in groups.iter().enumerate().take(best_start) {
if idx > 0 {
out.push(':');
}
out.push_str(&alloc::format!("{g:x}"));
}
out.push_str("::");
for (idx, g) in groups.iter().enumerate().skip(best_start + best_len) {
if idx > best_start + best_len {
out.push(':');
}
out.push_str(&alloc::format!("{g:x}"));
}
} else {
for (idx, g) in groups.iter().enumerate() {
if idx > 0 {
out.push(':');
}
out.push_str(&alloc::format!("{g:x}"));
}
}
out
}
pub(super) fn mysql_inet6_aton(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let s = match args {
[Value::Null] => return Ok(Value::Null),
[Value::Text(s)] => s,
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!("inet6_aton() takes one TEXT arg, got {} args", args.len()),
});
}
};
if let Some(b) = parse_ipv4(s) {
return Ok(Value::Bytes(alloc::borrow::Cow::Owned(b.to_vec())));
}
match parse_ipv6(s) {
Some(b) => Ok(Value::Bytes(alloc::borrow::Cow::Owned(b.to_vec()))),
None => Ok(Value::Null),
}
}
pub(super) fn mysql_inet6_ntoa(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
let bytes = match args {
[Value::Null] => return Ok(Value::Null),
[Value::Bytes(b)] => b.as_ref(),
_ => {
return Err(EvalError::TypeMismatch {
detail: alloc::format!(
"inet6_ntoa() takes one binary arg, got {}",
args.first().map_or_else(
|| alloc::string::String::from("no arguments"),
|a| crate::conversions::pg_type_name_for_error_opt(a.data_type()),
)
),
});
}
};
match bytes.len() {
4 => Ok(Value::text(alloc::format!(
"{}.{}.{}.{}",
bytes[0],
bytes[1],
bytes[2],
bytes[3]
))),
16 => {
let mut b = [0u8; 16];
b.copy_from_slice(bytes);
Ok(Value::text(format_ipv6(&b)))
}
_ => Ok(Value::Null),
}
}
pub(super) fn mysql_is_ipv4(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
match args {
[Value::Null] => Ok(Value::Null),
[Value::Text(s)] => Ok(Value::Bool(parse_ipv4(s).is_some())),
_ => Err(EvalError::TypeMismatch {
detail: alloc::format!("is_ipv4() takes one TEXT arg, got {} args", args.len()),
}),
}
}
pub(super) fn mysql_is_ipv6(args: &[Value<'_>]) -> Result<Value<'static>, EvalError> {
match args {
[Value::Null] => Ok(Value::Null),
[Value::Text(s)] => Ok(Value::Bool(!s.is_empty() && parse_ipv6(s).is_some())),
_ => Err(EvalError::TypeMismatch {
detail: alloc::format!("is_ipv6() takes one TEXT arg, got {} args", args.len()),
}),
}
}