use super::CqlType;
use crate::types::UdtTypeDef;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
pub fn split_qualified_udt<'a>(
reference: &'a str,
default_keyspace: &'a str,
) -> (&'a str, &'a str) {
match reference.split_once('.') {
Some((keyspace, bare_name)) => (keyspace, bare_name),
None => (default_keyspace, reference),
}
}
pub fn udt_registry_from_cql(cql: &str, default_keyspace: &str) -> UdtRegistry {
use super::cql_parser::{parse_create_type, split_cql_statements};
let mut registry = UdtRegistry::new();
for stmt in split_cql_statements(cql) {
if let Ok((_, (name, keyspace, fields))) = parse_create_type(&stmt) {
let keyspace = keyspace.unwrap_or_else(|| default_keyspace.to_string());
let mut def = UdtTypeDef::new(keyspace, name);
for (field_name, field_type) in fields {
let cql = CqlType::parse(&field_type).unwrap_or(CqlType::Blob);
def = def.with_field(field_name, cql, true);
}
registry.register_udt(def);
}
}
registry
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct UdtRegistry {
udts: HashMap<String, HashMap<String, UdtTypeDef>>,
}
impl UdtRegistry {
pub fn new() -> Self {
Self {
udts: HashMap::new(),
}
}
pub fn with_cassandra5_defaults() -> Self {
let mut registry = Self::new();
registry.load_cassandra5_system_udts();
registry
}
pub fn register_udt(&mut self, udt_def: UdtTypeDef) {
let keyspace_udts = self.udts.entry(udt_def.keyspace.clone()).or_default();
keyspace_udts.insert(udt_def.name.clone(), udt_def);
}
pub fn get_udt(&self, keyspace: &str, name: &str) -> Option<&UdtTypeDef> {
self.udts.get(keyspace)?.get(name)
}
pub fn get_udt_qualified(
&self,
default_keyspace: &str,
reference: &str,
) -> Option<&UdtTypeDef> {
let reference = reference.strip_prefix("udt:").unwrap_or(reference);
let (keyspace, bare_name) = split_qualified_udt(reference, default_keyspace);
self.get_udt(keyspace, bare_name)
.or_else(|| self.get_udt(default_keyspace, reference))
}
pub fn get_keyspace_udts(&self, keyspace: &str) -> Option<&HashMap<String, UdtTypeDef>> {
self.udts.get(keyspace)
}
pub fn list_udt_names(&self, keyspace: &str) -> Vec<&str> {
self.udts
.get(keyspace)
.map(|udts| udts.keys().map(|s| s.as_str()).collect())
.unwrap_or_default()
}
pub fn contains_udt(&self, keyspace: &str, name: &str) -> bool {
self.udts
.get(keyspace)
.map(|udts| udts.contains_key(name))
.unwrap_or(false)
}
pub fn remove_udt(&mut self, keyspace: &str, name: &str) -> Option<UdtTypeDef> {
self.udts.get_mut(keyspace)?.remove(name)
}
pub fn clear_keyspace(&mut self, keyspace: &str) {
self.udts.remove(keyspace);
}
pub fn total_udts(&self) -> usize {
self.udts.values().map(|udts| udts.len()).sum()
}
fn load_cassandra5_system_udts(&mut self) {
let address_udt = UdtTypeDef::new("system".to_string(), "address".to_string())
.with_field("street".to_string(), CqlType::Text, true)
.with_field("street2".to_string(), CqlType::Text, true)
.with_field("city".to_string(), CqlType::Text, true)
.with_field("state".to_string(), CqlType::Text, true)
.with_field("zip_code".to_string(), CqlType::Text, true)
.with_field("country".to_string(), CqlType::Text, true)
.with_field(
"coordinates".to_string(),
CqlType::Tuple(vec![CqlType::Double, CqlType::Double]),
true,
);
self.register_udt(address_udt);
let person_udt = UdtTypeDef::new("system".to_string(), "person".to_string())
.with_field("id".to_string(), CqlType::Uuid, false)
.with_field("first_name".to_string(), CqlType::Text, false)
.with_field("last_name".to_string(), CqlType::Text, false)
.with_field("middle_name".to_string(), CqlType::Text, true)
.with_field("age".to_string(), CqlType::Int, true)
.with_field("email".to_string(), CqlType::Text, true)
.with_field(
"phone_numbers".to_string(),
CqlType::Set(Box::new(CqlType::Text)),
true,
)
.with_field(
"addresses".to_string(),
CqlType::List(Box::new(CqlType::Udt("address".to_string(), vec![]))),
true,
)
.with_field(
"metadata".to_string(),
CqlType::Map(Box::new(CqlType::Text), Box::new(CqlType::Text)),
true,
);
self.register_udt(person_udt);
let contact_info_udt = UdtTypeDef::new("system".to_string(), "contact_info".to_string())
.with_field(
"person".to_string(),
CqlType::Udt("person".to_string(), vec![]),
false,
)
.with_field(
"primary_address".to_string(),
CqlType::Udt("address".to_string(), vec![]),
true,
)
.with_field(
"emergency_contacts".to_string(),
CqlType::List(Box::new(CqlType::Udt("person".to_string(), vec![]))),
true,
)
.with_field("last_updated".to_string(), CqlType::Timestamp, true);
self.register_udt(contact_info_udt);
}
pub fn resolve_type(&self, ty: &CqlType, keyspace: &str) -> CqlType {
self.resolve_type_depth(ty, keyspace, 0)
}
fn resolve_type_depth(&self, ty: &CqlType, keyspace: &str, depth: usize) -> CqlType {
const MAX_DEPTH: usize = 32;
if depth >= MAX_DEPTH {
return ty.clone();
}
match ty {
CqlType::List(inner) => CqlType::List(Box::new(self.resolve_type_depth(
inner,
keyspace,
depth + 1,
))),
CqlType::Set(inner) => CqlType::Set(Box::new(self.resolve_type_depth(
inner,
keyspace,
depth + 1,
))),
CqlType::Frozen(inner) => CqlType::Frozen(Box::new(self.resolve_type_depth(
inner,
keyspace,
depth + 1,
))),
CqlType::Map(k, v) => CqlType::Map(
Box::new(self.resolve_type_depth(k, keyspace, depth + 1)),
Box::new(self.resolve_type_depth(v, keyspace, depth + 1)),
),
CqlType::Tuple(types) => CqlType::Tuple(
types
.iter()
.map(|t| self.resolve_type_depth(t, keyspace, depth + 1))
.collect(),
),
CqlType::Udt(name, fields) if fields.is_empty() => self
.resolve_udt_reference(name, keyspace, depth)
.unwrap_or_else(|| ty.clone()),
CqlType::Udt(name, fields) => CqlType::Udt(
name.clone(),
fields
.iter()
.map(|(fname, ftype)| {
(
fname.clone(),
self.resolve_type_depth(ftype, keyspace, depth + 1),
)
})
.collect(),
),
CqlType::Custom(name) => {
let udt_name = name.strip_prefix("udt:").unwrap_or(name);
if super::is_udt_identifier(udt_name) {
self.resolve_udt_reference(udt_name, keyspace, depth)
.unwrap_or_else(|| ty.clone())
} else {
ty.clone()
}
}
other => other.clone(),
}
}
fn resolve_udt_reference(
&self,
udt_name: &str,
keyspace: &str,
depth: usize,
) -> Option<CqlType> {
let (lookup_keyspace, bare_name) = split_qualified_udt(udt_name, keyspace);
let def = self
.get_udt(lookup_keyspace, bare_name)
.or_else(|| self.get_udt("system", bare_name))?;
let fields = def
.fields
.iter()
.map(|f| {
(
f.name.clone(),
self.resolve_type_depth(&f.field_type, &def.keyspace, depth + 1),
)
})
.collect();
Some(CqlType::Udt(bare_name.to_string(), fields))
}
pub fn resolve_udt_with_dependencies(
&self,
keyspace: &str,
name: &str,
) -> crate::Result<&UdtTypeDef> {
let udt = self.get_udt(keyspace, name).ok_or_else(|| {
crate::Error::schema(format!(
"UDT '{}' not found in keyspace '{}'",
name, keyspace
))
})?;
for field in &udt.fields {
self.validate_field_type_dependencies(&field.field_type, keyspace)?;
}
Ok(udt)
}
fn validate_field_type_dependencies(
&self,
field_type: &CqlType,
keyspace: &str,
) -> crate::Result<()> {
match field_type {
CqlType::Udt(udt_name, _) => {
if !self.contains_udt(keyspace, udt_name) {
return Err(crate::Error::schema(format!(
"UDT dependency '{}' not found in keyspace '{}'",
udt_name, keyspace
)));
}
}
CqlType::List(inner) | CqlType::Set(inner) | CqlType::Frozen(inner) => {
self.validate_field_type_dependencies(inner, keyspace)?;
}
CqlType::Map(key_type, value_type) => {
self.validate_field_type_dependencies(key_type, keyspace)?;
self.validate_field_type_dependencies(value_type, keyspace)?;
}
CqlType::Tuple(field_types) => {
for tuple_field_type in field_types {
self.validate_field_type_dependencies(tuple_field_type, keyspace)?;
}
}
_ => {} }
Ok(())
}
pub fn get_dependent_udts(&self, keyspace: &str, udt_name: &str) -> Vec<&UdtTypeDef> {
let mut dependents = Vec::new();
if let Some(keyspace_udts) = self.udts.get(keyspace) {
for udt in keyspace_udts.values() {
if udt.name == udt_name {
continue; }
if self.udt_depends_on(udt, udt_name) {
dependents.push(udt);
}
}
}
dependents
}
fn udt_depends_on(&self, udt: &UdtTypeDef, target_udt: &str) -> bool {
for field in &udt.fields {
if self.field_type_depends_on(&field.field_type, target_udt) {
return true;
}
}
false
}
#[allow(clippy::only_used_in_recursion)]
fn field_type_depends_on(&self, field_type: &CqlType, target_udt: &str) -> bool {
match field_type {
CqlType::Udt(udt_name, _) => udt_name == target_udt,
CqlType::List(inner) | CqlType::Set(inner) | CqlType::Frozen(inner) => {
self.field_type_depends_on(inner, target_udt)
}
CqlType::Map(key_type, value_type) => {
self.field_type_depends_on(key_type, target_udt)
|| self.field_type_depends_on(value_type, target_udt)
}
CqlType::Tuple(field_types) => field_types
.iter()
.any(|ft| self.field_type_depends_on(ft, target_udt)),
_ => false,
}
}
pub fn register_udt_with_validation(&mut self, udt_def: UdtTypeDef) -> crate::Result<()> {
for field in &udt_def.fields {
self.validate_field_type_dependencies(&field.field_type, &udt_def.keyspace)?;
}
if self.would_create_circular_dependency(&udt_def) {
return Err(crate::Error::schema(format!(
"Registering UDT '{}' would create circular dependency",
udt_def.name
)));
}
self.register_udt(udt_def);
Ok(())
}
fn would_create_circular_dependency(&self, udt_def: &UdtTypeDef) -> bool {
for field in &udt_def.fields {
if self.field_type_depends_on(&field.field_type, &udt_def.name) {
return true;
}
}
false
}
pub fn export_definitions(&self, keyspace: &str) -> Vec<String> {
let mut definitions = Vec::new();
if let Some(keyspace_udts) = self.udts.get(keyspace) {
for udt in keyspace_udts.values() {
let mut def = format!("CREATE TYPE {}.{} (\n", keyspace, udt.name);
for (i, field) in udt.fields.iter().enumerate() {
if i > 0 {
def.push_str(",\n");
}
def.push_str(&format!(
" {} {}",
field.name,
self.format_cql_type(&field.field_type)
));
}
def.push_str("\n);");
definitions.push(def);
}
}
definitions
}
#[allow(clippy::only_used_in_recursion)]
fn format_cql_type(&self, cql_type: &CqlType) -> String {
match cql_type {
CqlType::Boolean => "boolean".to_string(),
CqlType::TinyInt => "tinyint".to_string(),
CqlType::SmallInt => "smallint".to_string(),
CqlType::Int => "int".to_string(),
CqlType::BigInt => "bigint".to_string(),
CqlType::Counter => "counter".to_string(),
CqlType::Float => "float".to_string(),
CqlType::Double => "double".to_string(),
CqlType::Text | CqlType::Varchar => "text".to_string(),
CqlType::Ascii => "ascii".to_string(),
CqlType::Blob => "blob".to_string(),
CqlType::Timestamp => "timestamp".to_string(),
CqlType::Date => "date".to_string(),
CqlType::Time => "time".to_string(),
CqlType::Uuid => "uuid".to_string(),
CqlType::TimeUuid => "timeuuid".to_string(),
CqlType::Inet => "inet".to_string(),
CqlType::Duration => "duration".to_string(),
CqlType::Varint => "varint".to_string(),
CqlType::Decimal => "decimal".to_string(),
CqlType::List(inner) => format!("list<{}>", self.format_cql_type(inner)),
CqlType::Set(inner) => format!("set<{}>", self.format_cql_type(inner)),
CqlType::Map(key, value) => format!(
"map<{}, {}>",
self.format_cql_type(key),
self.format_cql_type(value)
),
CqlType::Udt(name, _) => name.clone(),
CqlType::Tuple(types) => {
let type_strs: Vec<String> =
types.iter().map(|t| self.format_cql_type(t)).collect();
format!("tuple<{}>", type_strs.join(", "))
}
CqlType::Frozen(inner) => format!("frozen<{}>", self.format_cql_type(inner)),
CqlType::Custom(name) => name.clone(),
}
}
}
#[cfg(test)]
mod resolve_tests {
use super::*;
const DDL: &str = "\
CREATE TYPE ks.address_type (street text, city text); \
CREATE TYPE ks.contact_info (email text, address frozen<address_type>);";
#[test]
fn get_udt_qualified_splits_keyspace_and_strips_udt_prefix() {
let reg = udt_registry_from_cql(DDL, "ks");
assert!(reg.get_udt_qualified("other", "ks.address_type").is_some());
assert!(reg.get_udt_qualified("ks", "udt:address_type").is_some());
assert!(reg.get_udt_qualified("ks", "address_type").is_some());
assert!(reg.get_udt_qualified("ks", "nope.address_type").is_none());
}
#[test]
fn get_udt_qualified_non_regressive_for_dotted_bare_name() {
let mut reg = UdtRegistry::new();
reg.register_udt(
UdtTypeDef::new("ks".to_string(), "my.type".to_string()).with_field(
"v".to_string(),
CqlType::Text,
true,
),
);
let got = reg
.get_udt_qualified("ks", "my.type")
.expect("dotted bare name resolves");
assert_eq!(got.name, "my.type");
}
#[test]
fn from_cql_registers_every_create_type() {
let reg = udt_registry_from_cql(DDL, "ks");
assert!(reg.contains_udt("ks", "address_type"));
assert!(reg.contains_udt("ks", "contact_info"));
assert_eq!(reg.total_udts(), 2);
}
#[test]
fn from_cql_no_create_type_is_empty() {
let reg = udt_registry_from_cql("CREATE TABLE ks.t (id int PRIMARY KEY, v text)", "ks");
assert_eq!(reg.total_udts(), 0);
}
#[test]
fn resolve_type_rewrites_custom_udt_in_list_to_struct() {
let reg = udt_registry_from_cql(DDL, "ks");
let parsed = CqlType::parse("list<frozen<address_type>>").unwrap();
let resolved = reg.resolve_type(&parsed, "ks");
match &resolved {
CqlType::List(inner) => match inner.as_ref() {
CqlType::Frozen(udt) => match udt.as_ref() {
CqlType::Udt(name, fields) => {
assert_eq!(name, "address_type");
assert_eq!(fields.len(), 2, "street + city resolved from the registry");
assert_eq!(fields[0].0, "street");
}
other => panic!("inner must resolve to Udt, got {other:?}"),
},
other => panic!("expected Frozen wrapper, got {other:?}"),
},
other => panic!("expected List, got {other:?}"),
}
}
#[test]
fn resolve_type_recurses_into_nested_udt_fields() {
let reg = udt_registry_from_cql(DDL, "ks");
let parsed = CqlType::parse("frozen<contact_info>").unwrap();
let resolved = reg.resolve_type(&parsed, "ks");
let inner = match &resolved {
CqlType::Frozen(inner) => inner.as_ref(),
other => panic!("expected Frozen, got {other:?}"),
};
let fields = match inner {
CqlType::Udt(_, fields) => fields,
other => panic!("expected Udt, got {other:?}"),
};
let (_, addr_type) = fields
.iter()
.find(|(n, _)| n == "address")
.expect("address field");
match addr_type {
CqlType::Frozen(a) => assert!(
matches!(a.as_ref(), CqlType::Udt(n, f) if n == "address_type" && f.len() == 2),
"nested address field must resolve to the full address_type Struct"
),
CqlType::Udt(n, f) => {
assert_eq!(n, "address_type");
assert_eq!(f.len(), 2);
}
other => panic!("nested address must resolve to a Udt, got {other:?}"),
}
}
#[test]
fn resolve_type_unknown_reference_is_left_unchanged() {
let reg = UdtRegistry::new();
let parsed = CqlType::parse("list<frozen<missing_type>>").unwrap();
assert_eq!(reg.resolve_type(&parsed, "ks"), parsed);
}
#[test]
fn resolve_type_leaves_primitives_untouched() {
let reg = udt_registry_from_cql(DDL, "ks");
let parsed = CqlType::parse("map<text, int>").unwrap();
assert_eq!(reg.resolve_type(&parsed, "ks"), parsed);
}
#[test]
fn resolve_type_recurses_into_a_partially_resolved_udt_node() {
let reg = udt_registry_from_cql(DDL, "ks");
let partial = CqlType::Udt(
"contact_info".to_string(),
vec![
("email".to_string(), CqlType::Text),
(
"address".to_string(),
CqlType::Frozen(Box::new(CqlType::Custom("udt:address_type".to_string()))),
),
],
);
let resolved = reg.resolve_type(&partial, "ks");
let fields = match &resolved {
CqlType::Udt(_, fields) => fields,
other => panic!("expected Udt, got {other:?}"),
};
let (_, addr) = fields
.iter()
.find(|(n, _)| n == "address")
.expect("address");
let inner = match addr {
CqlType::Frozen(inner) => inner.as_ref(),
other => panic!("expected Frozen, got {other:?}"),
};
assert!(
matches!(inner, CqlType::Udt(n, f) if n == "address_type" && f.len() == 2),
"the inner Custom(\"udt:address_type\") must resolve to the full Struct, got {inner:?}"
);
}
#[test]
fn resolve_type_qualified_reference_resolves_to_bare_node_name() {
let reg = udt_registry_from_cql(DDL, "ks");
let parsed = CqlType::parse("frozen<ks.address_type>").unwrap();
let resolved = reg.resolve_type(&parsed, "other_ks");
match &resolved {
CqlType::Frozen(inner) => match inner.as_ref() {
CqlType::Udt(name, fields) => {
assert_eq!(name, "address_type", "node name stays BARE, not qualified");
assert_eq!(fields.len(), 2, "resolved against the ks keyspace");
}
other => panic!("expected Udt, got {other:?}"),
},
other => panic!("expected Frozen, got {other:?}"),
}
}
}