1use super::{EnumVariant, Field, FieldId, FieldPrimitive, FieldTy, Model, ModelId};
2
3use crate::{Result, stmt};
4use indexmap::IndexMap;
5use std::collections::HashSet;
6
7#[derive(Debug)]
26pub enum Resolved<'a> {
27 Field(&'a Field),
29 Variant(&'a EnumVariant),
31}
32
33#[derive(Debug, Default)]
50pub struct Schema {
51 pub models: IndexMap<ModelId, Model>,
53}
54
55#[derive(Default)]
56struct Builder {
57 models: IndexMap<ModelId, Model>,
58}
59
60impl Schema {
61 pub fn from_macro(models: impl IntoIterator<Item = Model>) -> Result<Self> {
66 Builder::from_macro(models)
67 }
68
69 pub fn field(&self, id: FieldId) -> &Field {
75 self.model(id.model)
76 .fields()
77 .get(id.index)
78 .expect("invalid field ID")
79 }
80
81 pub fn models(&self) -> impl Iterator<Item = &Model> {
83 self.models.values()
84 }
85
86 pub fn get_model(&self, id: impl Into<ModelId>) -> Option<&Model> {
88 self.models.get(&id.into())
89 }
90
91 pub fn model(&self, id: impl Into<ModelId>) -> &Model {
97 self.models.get(&id.into()).expect("invalid model ID")
98 }
99
100 pub fn fields(&self, id: impl Into<ModelId>) -> &[Field] {
111 self.model(id).fields()
112 }
113
114 pub fn project_fields<'s>(
131 &'s self,
132 model: ModelId,
133 projection: &[usize],
134 ) -> impl Iterator<Item = &'s Field> {
135 let mut current = Some(model);
136 let mut steps = projection.iter();
137
138 std::iter::from_fn(move || {
139 let &index = steps.next()?;
140 let field = self.get_model(current?)?.fields().get(index)?;
141 current = match field.expr_ty() {
142 stmt::Type::Model(nested) => Some(*nested),
143 _ => None,
144 };
145 Some(field)
146 })
147 }
148
149 pub fn resolve<'a>(
161 &'a self,
162 root: &'a Model,
163 projection: &stmt::Projection,
164 ) -> Option<Resolved<'a>> {
165 let [first, rest @ ..] = projection.as_slice() else {
166 return None;
167 };
168
169 let mut current_field = root.as_root_unwrap().fields.get(*first)?;
171
172 let mut steps = rest.iter();
175 while let Some(step) = steps.next() {
176 match ¤t_field.ty {
177 FieldTy::Primitive(FieldPrimitive {
184 ty: stmt::Type::Model(embed_id),
185 ..
186 }) => {
187 let consumed = rest.len() - steps.as_slice().len();
192 let doc_path = &rest[consumed - 1..];
193 return (self.project_fields(*embed_id, doc_path).count() == doc_path.len())
194 .then_some(Resolved::Field(current_field));
195 }
196 FieldTy::Primitive(..) => {
197 return None;
199 }
200 FieldTy::Embedded(embedded) => {
201 let target = self.model(embedded.target);
202 match target {
203 Model::EmbeddedStruct(s) => {
204 current_field = s.fields.get(*step)?;
205 }
206 Model::EmbeddedEnum(e) => {
207 let variant_index = *step;
208 let variant = e.variants.get(variant_index)?;
209
210 if let Some(field_step) = steps.next() {
212 current_field = e.variant_fields(variant_index).get(*field_step)?;
215 } else {
216 return Some(Resolved::Variant(variant));
218 }
219 }
220 _ => return None,
221 }
222 }
223 FieldTy::BelongsTo(belongs_to) => {
224 current_field = belongs_to.target(self).as_root_unwrap().fields.get(*step)?;
225 }
226 FieldTy::Has(has) => {
227 current_field = has.target(self).as_root_unwrap().fields.get(*step)?;
228 }
229 FieldTy::Via(via) => {
230 current_field = via.target(self).as_root_unwrap().fields.get(*step)?;
231 }
232 };
233 }
234
235 Some(Resolved::Field(current_field))
236 }
237
238 pub fn resolve_field<'a>(
243 &'a self,
244 root: &'a Model,
245 projection: &stmt::Projection,
246 ) -> Option<&'a Field> {
247 match self.resolve(root, projection) {
248 Some(Resolved::Field(field)) => Some(field),
249 _ => None,
250 }
251 }
252
253 pub fn resolve_field_path<'a>(&'a self, path: &stmt::Path) -> Option<&'a Field> {
256 let model = self.model(path.root.as_model_unwrap());
257 self.resolve_field(model, &path.projection)
258 }
259}
260
261impl Builder {
262 pub(crate) fn from_macro(models: impl IntoIterator<Item = Model>) -> Result<Schema> {
263 let mut builder = Self { ..Self::default() };
264
265 for model in models {
266 builder.models.insert(model.id(), model);
267 }
268
269 builder.process_models()?;
270 builder.into_schema()
271 }
272
273 fn into_schema(self) -> Result<Schema> {
274 Ok(Schema {
275 models: self.models,
276 })
277 }
278
279 fn process_models(&mut self) -> Result<()> {
280 self.link_relations()?;
283 self.resolve_via_targets()?;
284 self.verify_no_eager_load_cycles()?;
285
286 Ok(())
287 }
288
289 fn resolve_via_targets(&mut self) -> crate::Result<()> {
298 let mut updates = Vec::new();
300
301 for curr in 0..self.models.len() {
302 if self.models[curr].is_embedded() {
303 continue;
304 }
305 let src = self.models[curr].id();
306 for index in 0..self.models[curr].as_root_unwrap().fields.len() {
307 let field = &self.models[curr].as_root_unwrap().fields[index];
308 let FieldTy::Via(via) = &field.ty else {
309 continue;
310 };
311 let Some(terminal) = via.terminal else {
312 continue;
313 };
314
315 let projection = via.path.projection.as_slice();
317 let relation_steps = &projection[..projection.len() - 1];
318 let field_name = field.name.app_unwrap().to_string();
319 let target = self.walk_via_relation_chain(src, relation_steps, &field_name)?;
320
321 let terminal_field = &self.models[&target].as_root_unwrap().fields[terminal];
323 if !matches!(terminal_field.ty, FieldTy::Primitive(_)) {
324 return Err(crate::Error::invalid_schema(format!(
325 "the `via` terminal `{}::{}` is not a scalar field",
326 self.models[&target].name().upper_camel_case(),
327 terminal_field.name.app_unwrap(),
328 )));
329 }
330
331 updates.push((curr, index, target));
332 }
333 }
334
335 for (curr, index, target) in updates {
336 if let FieldTy::Via(via) = &mut self.models[curr].as_root_mut_unwrap().fields[index].ty
337 {
338 via.target = target;
339 }
340 }
341
342 Ok(())
343 }
344
345 fn walk_via_relation_chain(
348 &self,
349 declaring: ModelId,
350 steps: &[usize],
351 field_name: &str,
352 ) -> crate::Result<ModelId> {
353 let mut current = declaring;
354 let mut queue: Vec<usize> = steps.iter().rev().copied().collect();
355
356 while let Some(idx) = queue.pop() {
357 let field = &self.models[¤t].as_root_unwrap().fields[idx];
358 match &field.ty {
359 FieldTy::Has(has) => current = has.target,
360 FieldTy::BelongsTo(belongs_to) => current = belongs_to.target,
361 FieldTy::Via(inner) => {
364 let inner_projection = inner.path.projection.as_slice();
365 let inner_steps = match inner.terminal {
366 Some(_) => &inner_projection[..inner_projection.len() - 1],
367 None => inner_projection,
368 };
369 for step in inner_steps.iter().rev() {
370 queue.push(*step);
371 }
372 }
373 _ => {
374 return Err(crate::Error::invalid_schema(format!(
375 "the `via` path for `{}::{}` traverses `{}`, which is not a relation",
376 self.models[&declaring].name().upper_camel_case(),
377 field_name,
378 field.name.app_unwrap(),
379 )));
380 }
381 }
382 }
383
384 Ok(current)
385 }
386
387 fn verify_no_eager_load_cycles(&self) -> crate::Result<()> {
388 let mut visited = HashSet::new();
389 let mut model_stack = Vec::new();
390 let mut field_stack = Vec::new();
391
392 for model in self.models.values() {
393 if model.is_embedded() {
394 continue;
395 }
396 self.visit_eager_load_graph(
397 model.id(),
398 &mut visited,
399 &mut model_stack,
400 &mut field_stack,
401 )?;
402 }
403
404 Ok(())
405 }
406
407 fn visit_eager_load_graph(
408 &self,
409 model_id: ModelId,
410 visited: &mut HashSet<ModelId>,
411 model_stack: &mut Vec<ModelId>,
412 field_stack: &mut Vec<FieldId>,
413 ) -> crate::Result<()> {
414 if model_stack.contains(&model_id) {
415 return Ok(());
416 }
417
418 if !visited.insert(model_id) {
419 return Ok(());
420 }
421
422 model_stack.push(model_id);
423
424 let model = self.models[&model_id].as_root_unwrap();
425 for field in &model.fields {
426 let Some(target) = eager_relation_target(field) else {
427 continue;
428 };
429
430 if let Some(pos) = model_stack.iter().position(|id| *id == target) {
431 let mut cycle = field_stack[pos..].to_vec();
432 cycle.push(field.id);
433 return Err(crate::Error::invalid_schema(format!(
434 "eager relation cycle detected: {}",
435 self.format_eager_load_cycle(&cycle, target)
436 )));
437 }
438
439 field_stack.push(field.id);
440 self.visit_eager_load_graph(target, visited, model_stack, field_stack)?;
441 field_stack.pop();
442 }
443
444 model_stack.pop();
445 Ok(())
446 }
447
448 fn format_eager_load_cycle(&self, fields: &[FieldId], target: ModelId) -> String {
449 let mut parts = Vec::new();
450 for field_id in fields {
451 let model = &self.models[&field_id.model];
452 let field = &model.as_root_unwrap().fields[field_id.index];
453 parts.push(format!(
454 "{}::{}",
455 model.name().upper_camel_case(),
456 field.name.app_unwrap()
457 ));
458 }
459 parts.push(self.models[&target].name().upper_camel_case());
460 parts.join(" -> ")
461 }
462
463 fn link_relations(&mut self) -> crate::Result<()> {
465 for curr in 0..self.models.len() {
473 if self.models[curr].is_embedded() {
474 continue;
475 }
476 for index in 0..self.models[curr].as_root_unwrap().fields.len() {
477 let model = &self.models[curr];
478 let src = model.id();
479 let field = &model.as_root_unwrap().fields[index];
480
481 if let FieldTy::Has(has) = &field.ty
482 && has.is_many()
483 {
484 let target = has.target;
485 let field_name = field.name.app_unwrap().to_string();
486 let pair = if has.pair_id.is_placeholder() {
487 self.find_has_many_pair(src, target, &field_name)?
488 } else {
489 self.validate_pair(src, target, &field_name, has.pair_id)?;
490 has.pair_id
491 };
492 self.models[curr].as_root_mut_unwrap().fields[index]
493 .ty
494 .as_has_mut_unwrap()
495 .pair_id = pair;
496 }
497 }
498 }
499
500 for curr in 0..self.models.len() {
502 if self.models[curr].is_embedded() {
503 continue;
504 }
505 for index in 0..self.models[curr].as_root_unwrap().fields.len() {
506 let model = &self.models[curr];
507 let src = model.id();
508 let field = &model.as_root_unwrap().fields[index];
509
510 match &field.ty {
511 FieldTy::Has(has) if has.is_one() => {
512 let target = has.target;
513 let field_name = field.name.app_unwrap().to_string();
514 let pair = if has.pair_id.is_placeholder() {
515 match self.find_belongs_to_pair(src, target, &field_name)? {
516 Some(pair) => pair,
517 None => {
518 return Err(crate::Error::invalid_schema(format!(
519 "field `{}::{}` has no matching `BelongsTo` relation on the target model",
520 self.models[curr].name().upper_camel_case(),
521 field_name,
522 )));
523 }
524 }
525 } else {
526 self.validate_pair(src, target, &field_name, has.pair_id)?;
527 has.pair_id
528 };
529
530 self.models[curr].as_root_mut_unwrap().fields[index]
531 .ty
532 .as_has_mut_unwrap()
533 .pair_id = pair;
534 }
535 FieldTy::BelongsTo(belongs_to) => {
536 assert!(!belongs_to.foreign_key.is_placeholder());
537 continue;
538 }
539 _ => {}
540 }
541 }
542 }
543
544 for curr in 0..self.models.len() {
546 if self.models[curr].is_embedded() {
547 continue;
548 }
549 for index in 0..self.models[curr].as_root_unwrap().fields.len() {
550 let model = &self.models[curr];
551 let field_id = model.as_root_unwrap().fields[index].id;
552
553 let pair = match &self.models[curr].as_root_unwrap().fields[index].ty {
554 FieldTy::BelongsTo(belongs_to) => {
555 let mut pair = None;
556 let target = match self.models.get_index_of(&belongs_to.target) {
557 Some(target) => target,
558 None => {
559 let model = &self.models[curr];
560 return Err(crate::Error::invalid_schema(format!(
561 "field `{}::{}` references a model that was not registered \
562 with the schema; did you forget to register it with `Db::builder()`?",
563 model.name().upper_camel_case(),
564 model.as_root_unwrap().fields[index].name(),
565 )));
566 }
567 };
568
569 for target_index in 0..self.models[target].as_root_unwrap().fields.len() {
570 pair = match &self.models[target].as_root_unwrap().fields[target_index]
571 .ty
572 {
573 FieldTy::Has(has) if has.pair_id == field_id => {
574 assert!(pair.is_none());
575 Some(
576 self.models[target].as_root_unwrap().fields[target_index]
577 .id,
578 )
579 }
580 _ => continue,
581 }
582 }
583
584 if pair.is_none() {
585 continue;
586 }
587
588 pair
589 }
590 _ => continue,
591 };
592
593 self.models[curr].as_root_mut_unwrap().fields[index]
594 .ty
595 .as_belongs_to_mut_unwrap()
596 .pair = pair;
597 }
598 }
599
600 Ok(())
601 }
602
603 fn find_belongs_to_pair(
604 &self,
605 src: ModelId,
606 target: ModelId,
607 field_name: &str,
608 ) -> crate::Result<Option<FieldId>> {
609 let src_model = &self.models[&src];
610
611 let target = match self.models.get(&target) {
612 Some(target) => target,
613 None => {
614 return Err(crate::Error::invalid_schema(format!(
615 "field `{}::{}` references a model that was not registered with the schema; \
616 did you forget to register it with `Db::builder()`?",
617 src_model.name().upper_camel_case(),
618 field_name,
619 )));
620 }
621 };
622
623 let belongs_to: Vec<_> = target
625 .as_root_unwrap()
626 .fields
627 .iter()
628 .filter(|field| match &field.ty {
629 FieldTy::BelongsTo(rel) => rel.target == src,
630 _ => false,
631 })
632 .collect();
633
634 match &belongs_to[..] {
635 [field] => Ok(Some(field.id)),
636 [] => Ok(None),
637 _ => Err(crate::Error::invalid_schema(format!(
638 "model `{}` has more than one `BelongsTo` relation targeting `{}`; \
639 disambiguate by adding `pair = <field>` on the paired `has_many`/`has_one` \
640 field",
641 target.name().upper_camel_case(),
642 src_model.name().upper_camel_case(),
643 ))),
644 }
645 }
646
647 fn find_has_many_pair(
648 &mut self,
649 src: ModelId,
650 target: ModelId,
651 field_name: &str,
652 ) -> crate::Result<FieldId> {
653 if let Some(field_id) = self.find_belongs_to_pair(src, target, field_name)? {
654 return Ok(field_id);
655 }
656
657 Err(crate::Error::invalid_schema(format!(
658 "field `{}::{}` has no matching `BelongsTo` relation on the target model",
659 self.models[&src].name().upper_camel_case(),
660 field_name,
661 )))
662 }
663
664 fn validate_pair(
668 &self,
669 src: ModelId,
670 target: ModelId,
671 field_name: &str,
672 pair: FieldId,
673 ) -> crate::Result<()> {
674 let src_model = &self.models[&src];
675
676 let target_model = match self.models.get(&target) {
677 Some(target) => target,
678 None => {
679 return Err(crate::Error::invalid_schema(format!(
680 "field `{}::{}` references a model that was not registered with the schema; \
681 did you forget to register it with `Db::builder()`?",
682 src_model.name().upper_camel_case(),
683 field_name,
684 )));
685 }
686 };
687
688 if pair.model != target {
689 return Err(crate::Error::invalid_schema(format!(
690 "field `{}::{}` specifies a `pair` on a model other than its target `{}`",
691 src_model.name().upper_camel_case(),
692 field_name,
693 target_model.name().upper_camel_case(),
694 )));
695 }
696
697 let paired = &target_model.as_root_unwrap().fields[pair.index];
698 match &paired.ty {
699 FieldTy::BelongsTo(rel) if rel.target == src => Ok(()),
700 _ => Err(crate::Error::invalid_schema(format!(
701 "field `{}::{}` specifies `pair = {}`, but `{}::{}` is not a `BelongsTo` \
702 targeting `{}`",
703 src_model.name().upper_camel_case(),
704 field_name,
705 paired.name.app_unwrap(),
706 target_model.name().upper_camel_case(),
707 paired.name.app_unwrap(),
708 src_model.name().upper_camel_case(),
709 ))),
710 }
711 }
712}
713
714fn eager_relation_target(field: &Field) -> Option<ModelId> {
715 if field.deferred {
716 return None;
717 }
718
719 field.relation_target_id()
720}