use std::collections::HashSet;
use crate::ast::{EnumBacking, PrimitiveType, SemanticType, TypeExpr};
use crate::diagnostic::{Diagnostic, ErrorClass};
use crate::ir::{
CompiledSchema, Encoding, FieldEncoding, ResolvedType, TypeDef, TypeId, WireSize,
POISON_TYPE_ID,
};
pub fn check(compiled: &mut CompiledSchema) -> Vec<Diagnostic> {
let mut diags = Vec::new();
check_recursion(compiled, &mut diags);
let decl_ids: Vec<TypeId> = compiled.declarations.clone();
for &id in &decl_ids {
match compiled.registry.get(id) {
Some(TypeDef::Enum(en)) => {
let wire_bits = compute_enum_wire_bits(en);
if let Some(TypeDef::Enum(en)) = compiled.registry.get_mut(id) {
en.wire_bits = wire_bits;
}
}
Some(TypeDef::Flags(fl)) => {
let wire_bytes = compute_flags_wire_bytes(fl);
if let Some(TypeDef::Flags(fl)) = compiled.registry.get_mut(id) {
fl.wire_bytes = wire_bytes;
}
}
_ => {}
}
}
for &id in &decl_ids {
if let Some(def) = compiled.registry.get(id) {
match def {
TypeDef::Message(_) => {
let mut computing = HashSet::new();
let ws = compute_message_wire_size(id, compiled, &mut computing);
if let Some(TypeDef::Message(msg)) = compiled.registry.get_mut(id) {
msg.wire_size = Some(ws);
}
}
TypeDef::Union(_) => {
let mut computing = HashSet::new();
let ws = compute_union_wire_size(id, compiled, &mut computing);
if let Some(TypeDef::Union(un)) = compiled.registry.get_mut(id) {
un.wire_size = Some(ws);
}
}
_ => {}
}
}
}
check_impl_conformance(compiled, &mut diags);
diags
}
fn cycle_wire_size() -> WireSize {
WireSize::Variable {
min_bits: 0,
max_bits: None,
}
}
fn compute_type_wire_size(
ty: &ResolvedType,
enc: &FieldEncoding,
compiled: &CompiledSchema,
computing: &mut HashSet<TypeId>,
) -> WireSize {
match &enc.encoding {
Encoding::Varint => varint_wire_size(ty),
Encoding::ZigZag => zigzag_wire_size(ty),
Encoding::Delta(inner) => {
let inner_enc = FieldEncoding {
encoding: *inner.clone(),
limit: enc.limit,
};
compute_type_wire_size(ty, &inner_enc, compiled, computing)
}
Encoding::Default => compute_resolved_type_wire_size(ty, compiled, computing),
}
}
fn compute_resolved_type_wire_size(
ty: &ResolvedType,
compiled: &CompiledSchema,
computing: &mut HashSet<TypeId>,
) -> WireSize {
match ty {
ResolvedType::Primitive(p) => primitive_wire_size(p),
ResolvedType::SubByte(s) => WireSize::Fixed(s.bits as u64),
ResolvedType::Semantic(s) => semantic_wire_size(s),
ResolvedType::Named(id) => named_type_wire_size(*id, compiled, computing),
ResolvedType::Optional(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
match inner_ws {
WireSize::Fixed(bits) => WireSize::Variable {
min_bits: 1,
max_bits: Some(1 + bits),
},
WireSize::Variable { max_bits, .. } => WireSize::Variable {
min_bits: 1,
max_bits: max_bits.map(|m| 1 + m),
},
}
}
ResolvedType::Array(_) => WireSize::Variable {
min_bits: 8,
max_bits: None,
},
ResolvedType::FixedArray(inner, size) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
match inner_ws {
WireSize::Fixed(bits) => WireSize::Fixed(bits * size),
WireSize::Variable { min_bits, .. } => WireSize::Variable {
min_bits: min_bits * size,
max_bits: None,
},
}
}
ResolvedType::Map(_, _) | ResolvedType::Set(_) => WireSize::Variable {
min_bits: 8,
max_bits: None,
},
ResolvedType::Result(ok, err) => {
let ok_ws = compute_resolved_type_wire_size(ok, compiled, computing);
let err_ws = compute_resolved_type_wire_size(err, compiled, computing);
let min_ok = wire_size_min_bits(&ok_ws);
let min_err = wire_size_min_bits(&err_ws);
let min = 1 + std::cmp::min(min_ok, min_err);
let max = match (wire_size_max_bits(&ok_ws), wire_size_max_bits(&err_ws)) {
(Some(a), Some(b)) => Some(1 + std::cmp::max(a, b)),
_ => None,
};
WireSize::Variable {
min_bits: min,
max_bits: max,
}
}
ResolvedType::Vec2(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 2)
}
ResolvedType::Vec3(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 3)
}
ResolvedType::Vec4(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 4)
}
ResolvedType::Quat(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 4)
}
ResolvedType::Mat3(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 9)
}
ResolvedType::Mat4(inner) => {
let inner_ws = compute_resolved_type_wire_size(inner, compiled, computing);
multiply_wire_size(&inner_ws, 16)
}
ResolvedType::BitsInline(names) => WireSize::Fixed(names.len() as u64),
}
}
fn primitive_wire_size(p: &PrimitiveType) -> WireSize {
let bits = match p {
PrimitiveType::Bool => 1,
PrimitiveType::U8 | PrimitiveType::I8 => 8,
PrimitiveType::U16 | PrimitiveType::I16 => 16,
PrimitiveType::U32 | PrimitiveType::I32 | PrimitiveType::F32 | PrimitiveType::Fixed32 => 32,
PrimitiveType::U64 | PrimitiveType::I64 | PrimitiveType::F64 | PrimitiveType::Fixed64 => 64,
PrimitiveType::Void => 0,
};
WireSize::Fixed(bits)
}
fn semantic_wire_size(s: &SemanticType) -> WireSize {
match s {
SemanticType::String | SemanticType::Bytes => WireSize::Variable {
min_bits: 0,
max_bits: None,
},
SemanticType::Rgb => WireSize::Fixed(24),
SemanticType::Uuid => WireSize::Fixed(128),
SemanticType::Timestamp => WireSize::Fixed(64),
SemanticType::Hash => WireSize::Fixed(256),
}
}
fn varint_wire_size(ty: &ResolvedType) -> WireSize {
let max_bits = match ty {
ResolvedType::Primitive(PrimitiveType::U16) => 24,
ResolvedType::Primitive(PrimitiveType::U32) => 40,
ResolvedType::Primitive(PrimitiveType::U64) => 80,
_ => 80,
};
WireSize::Variable {
min_bits: 8,
max_bits: Some(max_bits),
}
}
fn zigzag_wire_size(ty: &ResolvedType) -> WireSize {
let max_bits = match ty {
ResolvedType::Primitive(PrimitiveType::I16) => 24,
ResolvedType::Primitive(PrimitiveType::I32) => 40,
ResolvedType::Primitive(PrimitiveType::I64) => 80,
_ => 80,
};
WireSize::Variable {
min_bits: 8,
max_bits: Some(max_bits),
}
}
fn named_type_wire_size(
id: TypeId,
compiled: &CompiledSchema,
computing: &mut HashSet<TypeId>,
) -> WireSize {
if computing.contains(&id) {
return cycle_wire_size();
}
match compiled.registry.get(id) {
Some(TypeDef::Enum(en)) => {
WireSize::Fixed(u64::from(en.wire_bits))
}
Some(TypeDef::Flags(fl)) => {
WireSize::Fixed(u64::from(fl.wire_bytes) * 8)
}
Some(TypeDef::Newtype(nt)) => {
let terminal = nt.terminal_type.clone();
compute_resolved_type_wire_size(&terminal, compiled, computing)
}
Some(TypeDef::Message(msg)) => {
if let Some(ws) = msg.wire_size.clone() {
return ws;
}
compute_message_wire_size(id, compiled, computing)
}
Some(TypeDef::Union(un)) => {
if let Some(ws) = un.wire_size.clone() {
return ws;
}
compute_union_wire_size(id, compiled, computing)
}
Some(TypeDef::Config(_))
| Some(TypeDef::GenericAlias(_))
| Some(TypeDef::Trait(_))
| Some(TypeDef::Impl(_))
| None => WireSize::Variable {
min_bits: 0,
max_bits: None,
},
}
}
fn compute_message_wire_size(
id: TypeId,
compiled: &CompiledSchema,
computing: &mut HashSet<TypeId>,
) -> WireSize {
let msg = match compiled.registry.get(id) {
Some(TypeDef::Message(m)) => m,
_ => return WireSize::Fixed(0),
};
if msg.fields.is_empty() {
return WireSize::Fixed(0);
}
let fields: Vec<(ResolvedType, FieldEncoding)> = msg
.fields
.iter()
.map(|f| (f.resolved_type.clone(), f.encoding.clone()))
.collect();
computing.insert(id);
let mut total_min: u64 = 0;
let mut total_max: Option<u64> = Some(0);
let mut is_variable = false;
for (resolved_type, encoding) in &fields {
let ws = compute_type_wire_size(resolved_type, encoding, compiled, computing);
match ws {
WireSize::Fixed(bits) => {
total_min += bits;
if let Some(ref mut max) = total_max {
*max += bits;
}
}
WireSize::Variable { min_bits, max_bits } => {
is_variable = true;
total_min += min_bits;
match (total_max, max_bits) {
(Some(cur), Some(field_max)) => total_max = Some(cur + field_max),
_ => total_max = None,
}
}
}
}
computing.remove(&id);
if is_variable {
WireSize::Variable {
min_bits: total_min,
max_bits: total_max,
}
} else {
WireSize::Fixed(total_min)
}
}
fn compute_union_wire_size(
id: TypeId,
compiled: &CompiledSchema,
computing: &mut HashSet<TypeId>,
) -> WireSize {
let un = match compiled.registry.get(id) {
Some(TypeDef::Union(u)) => u,
_ => return WireSize::Fixed(0),
};
if un.variants.is_empty() {
return WireSize::Variable {
min_bits: 8,
max_bits: Some(8),
};
}
let tag_min: u64 = 8;
let mut max_variant_bits: Option<u64> = Some(0);
let mut min_variant_bits: u64 = u64::MAX;
let variants: Vec<Vec<(ResolvedType, FieldEncoding)>> = un
.variants
.iter()
.map(|v| {
v.fields
.iter()
.map(|f| (f.resolved_type.clone(), f.encoding.clone()))
.collect()
})
.collect();
computing.insert(id);
for variant_fields in &variants {
let mut var_min: u64 = 0;
let mut var_max: Option<u64> = Some(0);
for (resolved_type, encoding) in variant_fields {
let ws = compute_type_wire_size(resolved_type, encoding, compiled, computing);
match ws {
WireSize::Fixed(bits) => {
var_min += bits;
if let Some(ref mut max) = var_max {
*max += bits;
}
}
WireSize::Variable { min_bits, max_bits } => {
var_min += min_bits;
match (var_max, max_bits) {
(Some(cur), Some(field_max)) => var_max = Some(cur + field_max),
_ => var_max = None,
}
}
}
}
min_variant_bits = std::cmp::min(min_variant_bits, var_min);
match (max_variant_bits, var_max) {
(Some(cur), Some(v)) => max_variant_bits = Some(std::cmp::max(cur, v)),
_ => max_variant_bits = None,
}
}
computing.remove(&id);
if min_variant_bits == u64::MAX {
min_variant_bits = 0;
}
WireSize::Variable {
min_bits: tag_min + min_variant_bits,
max_bits: max_variant_bits.map(|m| tag_min + m),
}
}
fn wire_size_min_bits(ws: &WireSize) -> u64 {
match ws {
WireSize::Fixed(bits) => *bits,
WireSize::Variable { min_bits, .. } => *min_bits,
}
}
fn wire_size_max_bits(ws: &WireSize) -> Option<u64> {
match ws {
WireSize::Fixed(bits) => Some(*bits),
WireSize::Variable { max_bits, .. } => *max_bits,
}
}
fn multiply_wire_size(ws: &WireSize, multiplier: u64) -> WireSize {
match ws {
WireSize::Fixed(bits) => WireSize::Fixed(bits * multiplier),
WireSize::Variable { min_bits, max_bits } => WireSize::Variable {
min_bits: min_bits * multiplier,
max_bits: max_bits.map(|m| m * multiplier),
},
}
}
struct RecursionState<'a> {
direct_path: HashSet<TypeId>,
visited: HashSet<TypeId>,
compiled: &'a CompiledSchema,
origin_span: crate::span::Span,
diags: &'a mut Vec<Diagnostic>,
}
fn check_recursion(compiled: &CompiledSchema, diags: &mut Vec<Diagnostic>) {
for &id in &compiled.declarations {
if let Some(TypeDef::Message(msg)) = compiled.registry.get(id) {
let mut state = RecursionState {
direct_path: {
let mut s = HashSet::new();
s.insert(id);
s
},
visited: HashSet::new(),
compiled,
origin_span: msg.span,
diags,
};
let fields: Vec<(ResolvedType, FieldEncoding)> = msg
.fields
.iter()
.map(|f| (f.resolved_type.clone(), f.encoding.clone()))
.collect();
for (ty, _) in &fields {
walk_type_for_recursion(ty, true, &mut state);
}
}
}
}
fn compute_enum_wire_bits(en: &crate::ir::EnumDef) -> u8 {
if let Some(backing) = &en.backing {
return match backing {
EnumBacking::U8 => 8,
EnumBacking::U16 => 16,
EnumBacking::U32 => 32,
EnumBacking::U64 => 64,
};
}
let max_ordinal = en.variants.iter().map(|v| v.ordinal).max().unwrap_or(0);
let min_bits: u8 = if max_ordinal == 0 {
1
} else {
let n = u64::from(max_ordinal) + 1;
let leading = (n - 1).leading_zeros();
let bits = 64u8.saturating_sub(leading as u8);
std::cmp::max(bits, 1)
};
if en.annotations.non_exhaustive {
std::cmp::max(min_bits, 8)
} else {
std::cmp::max(min_bits, 1)
}
}
fn compute_flags_wire_bytes(fl: &crate::ir::FlagsDef) -> u8 {
let max_bit = fl.bits.iter().map(|b| b.bit).max().unwrap_or(0);
match max_bit {
0..=7 => 1,
8..=15 => 2,
16..=31 => 4,
_ => 8,
}
}
fn walk_type_for_recursion(ty: &ResolvedType, direct: bool, state: &mut RecursionState<'_>) {
match ty {
ResolvedType::Named(id) => {
if direct && state.direct_path.contains(id) {
state.diags.push(Diagnostic::error(
state.origin_span,
ErrorClass::RecursiveTypeInfinite,
"type contains infinite direct recursion",
));
return;
}
if !direct && state.direct_path.contains(id) {
return;
}
if state.visited.contains(id) {
return;
}
state.visited.insert(*id);
match state.compiled.registry.get(*id) {
Some(TypeDef::Message(msg)) => {
let was_new = if direct {
state.direct_path.insert(*id)
} else {
false
};
let fields: Vec<(ResolvedType, FieldEncoding)> = msg
.fields
.iter()
.map(|f| (f.resolved_type.clone(), f.encoding.clone()))
.collect();
for (field_ty, _) in &fields {
walk_type_for_recursion(field_ty, direct, state);
}
if was_new {
state.direct_path.remove(id);
}
}
Some(TypeDef::Union(un)) => {
let variant_fields: Vec<Vec<ResolvedType>> = un
.variants
.iter()
.map(|v| v.fields.iter().map(|f| f.resolved_type.clone()).collect())
.collect();
for fields in &variant_fields {
for field_ty in fields {
walk_type_for_recursion(field_ty, false, state);
}
}
}
Some(TypeDef::Newtype(nt)) => {
let inner = nt.inner_type.clone();
walk_type_for_recursion(&inner, direct, state);
}
_ => {} }
}
ResolvedType::Optional(inner) | ResolvedType::Array(inner) => {
walk_type_for_recursion(inner, false, state);
}
ResolvedType::FixedArray(inner, _) => {
walk_type_for_recursion(inner, false, state);
}
ResolvedType::Map(k, v) => {
walk_type_for_recursion(k, false, state);
walk_type_for_recursion(v, false, state);
}
ResolvedType::Set(inner) => {
walk_type_for_recursion(inner, false, state);
}
ResolvedType::Result(ok, err) => {
walk_type_for_recursion(ok, false, state);
walk_type_for_recursion(err, false, state);
}
ResolvedType::Vec2(inner)
| ResolvedType::Vec3(inner)
| ResolvedType::Vec4(inner)
| ResolvedType::Quat(inner)
| ResolvedType::Mat3(inner)
| ResolvedType::Mat4(inner) => {
walk_type_for_recursion(inner, false, state);
}
_ => {} }
}
use crate::ir::{ImplDef, TraitDef};
use smol_str::SmolStr;
fn check_impl_conformance(compiled: &CompiledSchema, diags: &mut Vec<Diagnostic>) {
for (impl_id, impl_def) in compiled.impls() {
let Some(trait_id) = compiled.registry.impl_trait_id(impl_id) else {
continue;
};
let Some(TypeDef::Trait(trait_def)) = compiled.registry.get(trait_id) else {
continue;
};
check_single_impl_conformance(
impl_def,
trait_def,
compiled.registry.impl_trait_span(impl_id),
compiled,
diags,
);
}
}
fn check_single_impl_conformance(
impl_def: &ImplDef,
trait_def: &TraitDef,
trait_span: Option<crate::span::Span>,
compiled: &CompiledSchema,
diags: &mut Vec<Diagnostic>,
) {
if impl_def.type_args.len() != trait_def.type_params.len() {
diags.push(Diagnostic::error(
trait_span.unwrap_or(impl_def.span),
ErrorClass::UnresolvedType,
format!(
"trait '{}' has {} type parameters but impl provides {}",
impl_def.trait_name,
trait_def.type_params.len(),
impl_def.type_args.len()
),
));
return;
}
check_trait_fields(impl_def, trait_def, compiled, diags);
check_trait_functions(impl_def, compiled, diags);
}
fn check_trait_fields(
impl_def: &ImplDef,
trait_def: &TraitDef,
compiled: &CompiledSchema,
diags: &mut Vec<Diagnostic>,
) {
let target_type_def = match &impl_def.target_type {
ResolvedType::Named(id) => compiled.registry.get(*id),
_ => None,
};
let Some(target_def) = target_type_def else {
return;
};
let target_fields = match target_def {
TypeDef::Message(m) => &m.fields,
_ => {
diags.push(Diagnostic::error(
impl_def.span,
ErrorClass::UnresolvedType,
format!(
"impl target '{:?}' is not a message type",
impl_def.target_type
),
));
return;
}
};
let type_param_names: Vec<&str> = trait_def
.type_params
.iter()
.map(|p| p.name.node.as_str())
.collect();
for trait_field in &trait_def.fields {
let substituted_ty = substitute_into_type_expr(
&trait_field.unresolved_ty,
&type_param_names,
&impl_def.type_args,
compiled,
);
let found = target_fields.iter().any(|f| {
f.name == trait_field.name && types_compatible(&f.resolved_type, &substituted_ty)
});
if !found {
diags.push(Diagnostic::error(
impl_def.span,
ErrorClass::UnresolvedType,
format!(
"impl for '{:?}' missing required trait field '{}' of type '{:?}'",
impl_def.target_type, trait_field.name, substituted_ty
),
));
}
}
}
fn check_trait_functions(
impl_def: &ImplDef,
compiled: &CompiledSchema,
diags: &mut Vec<Diagnostic>,
) {
use crate::codegen::portable::{project_impl, PortableFunctionError};
let error = match project_impl(compiled, impl_def) {
Ok(_) => return,
Err(error) => error,
};
let class = match &error {
PortableFunctionError::MissingFunction { .. } => ErrorClass::ImplFnMissing,
PortableFunctionError::ExtraFunction { .. } => ErrorClass::ImplFnExtra,
PortableFunctionError::DuplicateFunction { .. }
| PortableFunctionError::SignatureMismatch { .. } => ErrorClass::ImplFnSignatureMismatch,
PortableFunctionError::InvalidAssignmentTarget { .. } => {
ErrorClass::ImplFnAssignmentInvalid
}
PortableFunctionError::ReturnMismatch { .. }
| PortableFunctionError::StatementsAfterReturn { .. } => ErrorClass::ImplFnReturnMismatch,
PortableFunctionError::TypeMismatch { context, .. } if context == "return" => {
ErrorClass::ImplFnReturnMismatch
}
PortableFunctionError::UnsupportedCall { .. }
| PortableFunctionError::UnsupportedMethodCall { .. } => {
return;
}
PortableFunctionError::UnknownTrait { .. }
| PortableFunctionError::InvalidTarget
| PortableFunctionError::ExternalFunction { .. }
| PortableFunctionError::UnknownLocal { .. }
| PortableFunctionError::UnknownField { .. }
| PortableFunctionError::DuplicateLocal { .. }
| PortableFunctionError::UnsupportedExpressionStatement { .. }
| PortableFunctionError::TypeMismatch { .. }
| PortableFunctionError::InvalidOperator { .. }
| PortableFunctionError::UnresolvedType { .. } => ErrorClass::ImplFnBodyTypeMismatch,
};
diags.push(Diagnostic::error(impl_def.span, class, error.to_string()));
}
fn types_compatible(a: &ResolvedType, b: &ResolvedType) -> bool {
if a == b {
return true;
}
match (a, b) {
(ResolvedType::Named(id_a), ResolvedType::Named(id_b)) => id_a == id_b,
_ => false,
}
}
fn substitute_into_type_expr(
expr: &TypeExpr,
type_params: &[&str],
type_args: &[ResolvedType],
compiled: &CompiledSchema,
) -> ResolvedType {
match expr {
TypeExpr::Named(name) => {
if let Some(idx) = type_params.iter().position(|&p| p == name.as_str()) {
if idx < type_args.len() {
return type_args[idx].clone();
}
}
resolve_type_name(name, compiled)
}
TypeExpr::Primitive(p) => ResolvedType::Primitive(*p),
TypeExpr::SubByte(s) => ResolvedType::SubByte(*s),
TypeExpr::Semantic(s) => ResolvedType::Semantic(*s),
TypeExpr::Generic(name, arg) => {
let inner = Box::new(substitute_into_type_expr(
&arg.node,
type_params,
type_args,
compiled,
));
ResolvedType::Named(resolve_generic_type(name, *inner, compiled))
}
TypeExpr::Optional(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Optional(inner)
}
TypeExpr::Array(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Array(inner)
}
TypeExpr::FixedArray(inner, size) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::FixedArray(inner, *size)
}
TypeExpr::Set(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Set(inner)
}
TypeExpr::Map(key, value) => {
let key = Box::new(substitute_into_type_expr(
&key.node,
type_params,
type_args,
compiled,
));
let value = Box::new(substitute_into_type_expr(
&value.node,
type_params,
type_args,
compiled,
));
ResolvedType::Map(key, value)
}
TypeExpr::Result(ok, err) => {
let ok = Box::new(substitute_into_type_expr(
&ok.node,
type_params,
type_args,
compiled,
));
let err = Box::new(substitute_into_type_expr(
&err.node,
type_params,
type_args,
compiled,
));
ResolvedType::Result(ok, err)
}
TypeExpr::Vec2(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Vec2(inner)
}
TypeExpr::Vec3(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Vec3(inner)
}
TypeExpr::Vec4(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Vec4(inner)
}
TypeExpr::Quat(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Quat(inner)
}
TypeExpr::Mat3(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Mat3(inner)
}
TypeExpr::Mat4(inner) => {
let inner = Box::new(substitute_into_type_expr(
&inner.node,
type_params,
type_args,
compiled,
));
ResolvedType::Mat4(inner)
}
TypeExpr::BitsInline(names) => ResolvedType::BitsInline(names.clone()),
TypeExpr::Qualified(ns, name) => {
let qualified_name: SmolStr = format!("{ns}.{name}").into();
resolve_type_name(&qualified_name, compiled)
}
}
}
fn resolve_generic_type(
name: &SmolStr,
inner_type: ResolvedType,
compiled: &CompiledSchema,
) -> TypeId {
if let Some((_, crate::ir::TypeDef::GenericAlias(alias_def))) =
compiled.find_type(name.as_str())
{
let type_params: Vec<&str> = alias_def.type_params.iter().map(|p| p.as_str()).collect();
let type_args = vec![inner_type];
let substituted =
substitute_into_type_expr(&alias_def.target_type, &type_params, &type_args, compiled);
if let ResolvedType::Named(id) = substituted {
return id;
}
}
POISON_TYPE_ID
}
fn resolve_type_name(name: &SmolStr, compiled: &CompiledSchema) -> ResolvedType {
if let Some((id, _)) = compiled.find_type(name.as_str()) {
ResolvedType::Named(id)
} else {
ResolvedType::Named(POISON_TYPE_ID)
}
}