1use crate::{
2 BoxFuture, Container, Controller, EventHandler, EventHandlerDef, Injectable,
3 RegisterableEventHandler, WebSocketGateway,
4};
5use std::{
6 any::{Any, TypeId},
7 collections::{HashMap, HashSet},
8 fmt::Debug,
9 future::Future,
10 sync::Arc,
11 time::Instant,
12};
13
14type ProviderValue = Arc<dyn Any + Send + Sync>;
15type BuildProviderFn =
16 Box<dyn for<'a> Fn(&'a Container) -> BoxFuture<'a, crate::Result<ProviderValue>> + Send + Sync>;
17type LifecycleFn =
18 Box<dyn for<'a> Fn(&'a ProviderValue) -> BoxFuture<'a, crate::Result<()>> + Send + Sync>;
19
20#[derive(Clone, Copy)]
22pub struct ProviderDependency {
23 type_id: TypeId,
24 type_name: &'static str,
25}
26
27impl ProviderDependency {
28 pub fn of<T: Send + Sync + 'static>() -> Self {
29 Self {
30 type_id: TypeId::of::<T>(),
31 type_name: std::any::type_name::<T>(),
32 }
33 }
34
35 pub(crate) fn type_id(&self) -> TypeId {
36 self.type_id
37 }
38}
39
40#[macro_export]
42macro_rules! provider_dependencies {
43 ($($dependency:ty),* $(,)?) => {
44 vec![$($crate::ProviderDependency::of::<$dependency>()),*]
45 };
46}
47
48pub struct ProviderDef {
49 type_id: TypeId,
50 type_name: &'static str,
51 dependencies: Vec<ProviderDependency>,
52 build: BuildProviderFn,
53 init_fn: LifecycleFn,
54 bootstrap_fn: LifecycleFn,
55 shutdown_fn: LifecycleFn,
56}
57
58impl ProviderDef {
59 pub fn of<T: Injectable>() -> Self {
60 Self {
61 type_id: TypeId::of::<T>(),
62 type_name: std::any::type_name::<T>(),
63 dependencies: T::dependencies(),
64 build: Box::new(|container| {
65 Box::pin(async move { Ok(Arc::new(T::create(container).await?) as ProviderValue) })
66 }),
67 init_fn: Box::new(|value| {
68 let value = downcast_provider::<T>(value);
69 Box::pin(async move { value?.on_module_init().await })
70 }),
71 bootstrap_fn: Box::new(|value| {
72 let value = downcast_provider::<T>(value);
73 Box::pin(async move { value?.on_bootstrap().await })
74 }),
75 shutdown_fn: Box::new(|value| {
76 let value = downcast_provider::<T>(value);
77 Box::pin(async move { value?.on_shutdown().await })
78 }),
79 }
80 }
81
82 pub fn instance<T: Send + Sync + 'static>(value: T) -> Self {
85 let value = Arc::new(value) as ProviderValue;
86 Self {
87 type_id: TypeId::of::<T>(),
88 type_name: std::any::type_name::<T>(),
89 dependencies: vec![],
90 build: Box::new(move |_| {
91 let value = value.clone();
92 Box::pin(async move { Ok(value) })
93 }),
94 init_fn: noop_lifecycle(),
95 bootstrap_fn: noop_lifecycle(),
96 shutdown_fn: noop_lifecycle(),
97 }
98 }
99
100 pub fn async_factory<T, Fut, E>(
101 dependencies: Vec<ProviderDependency>,
102 factory: impl Fn(Arc<Container>) -> Fut + Send + Sync + 'static,
103 ) -> Self
104 where
105 T: Send + Sync + 'static,
106 Fut: Future<Output = std::result::Result<T, E>> + Send + 'static,
107 E: Debug + Send + 'static,
108 {
109 Self {
110 type_id: TypeId::of::<T>(),
111 type_name: std::any::type_name::<T>(),
112 dependencies,
113 build: Box::new(move |container| {
114 let future = factory(Arc::new(container.clone()));
115 Box::pin(async move {
116 let value = future.await.map_err(|err| {
117 crate::exception::startup_error(format!(
118 "async factory failed for {}: {:?}",
119 std::any::type_name::<T>(),
120 err
121 ))
122 })?;
123 Ok(Arc::new(value) as ProviderValue)
124 })
125 }),
126 init_fn: noop_lifecycle(),
127 bootstrap_fn: noop_lifecycle(),
128 shutdown_fn: noop_lifecycle(),
129 }
130 }
131
132 pub fn type_id(&self) -> TypeId {
133 self.type_id
134 }
135 pub fn type_name(&self) -> &'static str {
136 self.type_name
137 }
138
139 fn assert_registered(&self, container: &Container) -> crate::Result<()> {
140 if container.contains_type_id(self.type_id) {
141 return Ok(());
142 }
143 Err(crate::exception::startup_error(format!(
144 "missing provider at startup: {} was declared by module metadata but was not registered",
145 self.type_name
146 )))
147 }
148
149 async fn run_lifecycle(
150 &self,
151 value: &ProviderValue,
152 hook: &'static str,
153 callback: &LifecycleFn,
154 ) -> crate::Result<()> {
155 callback(value).await.map_err(|err| {
156 crate::exception::startup_error(format!(
157 "{hook} failed for {}: {}: {}",
158 self.type_name, err.error, err.message
159 ))
160 })
161 }
162
163 async fn run_lifecycle_from_container(
164 &self,
165 container: &Container,
166 hook: &'static str,
167 callback: &LifecycleFn,
168 ) -> crate::Result<()> {
169 let value = container.resolve_erased(self.type_id).ok_or_else(|| crate::exception::startup_error(format!(
170 "missing provider during {hook}: {} was declared by module metadata but was not registered", self.type_name
171 )))?;
172 self.run_lifecycle(&value, hook, callback).await
173 }
174}
175
176pub struct ProviderOverrides {
177 defs: HashMap<TypeId, ProviderDef>,
178}
179impl ProviderOverrides {
180 pub fn new() -> Self {
181 Self {
182 defs: HashMap::new(),
183 }
184 }
185 pub fn insert_instance<T: Send + Sync + 'static>(mut self, value: T) -> Self {
186 self.defs
187 .insert(TypeId::of::<T>(), ProviderDef::instance(value));
188 self
189 }
190 pub fn insert_factory<T, Fut, E>(
191 mut self,
192 dependencies: Vec<ProviderDependency>,
193 factory: impl Fn(Arc<Container>) -> Fut + Send + Sync + 'static,
194 ) -> Self
195 where
196 T: Send + Sync + 'static,
197 Fut: Future<Output = std::result::Result<T, E>> + Send + 'static,
198 E: Debug + Send + 'static,
199 {
200 self.defs.insert(
201 TypeId::of::<T>(),
202 ProviderDef::async_factory::<T, Fut, E>(dependencies, factory),
203 );
204 self
205 }
206 pub fn insert(mut self, def: ProviderDef) -> Self {
207 self.defs.insert(def.type_id, def);
208 self
209 }
210 pub(crate) fn into_inner(self) -> HashMap<TypeId, ProviderDef> {
211 self.defs
212 }
213}
214impl Default for ProviderOverrides {
215 fn default() -> Self {
216 Self::new()
217 }
218}
219
220fn downcast_provider<T: Send + Sync + 'static>(value: &ProviderValue) -> crate::Result<Arc<T>> {
221 value.clone().downcast::<T>().map_err(|_| {
222 crate::exception::startup_error(format!(
223 "type mismatch running lifecycle hook for {}",
224 std::any::type_name::<T>()
225 ))
226 })
227}
228fn noop_lifecycle() -> LifecycleFn {
229 Box::new(|_| Box::pin(async { Ok(()) }))
230}
231
232pub struct ControllerDef {
233 pub register_fn: fn(&mut dyn Any),
234 pub route_log_fn: fn(),
235 #[cfg(feature = "openapi")]
236 pub(crate) openapi_routes_fn: fn() -> &'static [crate::openapi::OpenApiRouteDef],
237 provider: ProviderDef,
238}
239impl ControllerDef {
240 pub fn of<C: Controller + Injectable + 'static>() -> Self {
241 Self {
242 register_fn: |any| C::register_routes(any),
243 route_log_fn: || crate::log_controller_routes::<C>(),
244 #[cfg(feature = "openapi")]
245 openapi_routes_fn: || C::openapi_routes(),
246 provider: ProviderDef::of::<C>(),
247 }
248 }
249}
250
251pub struct GatewayDef {
252 pub path: &'static str,
253 pub type_id: TypeId,
254 provider: ProviderDef,
255 kind: GatewayKind,
256}
257enum GatewayKind {
258 WebSocket {
259 resolve_fn: fn(&Container) -> crate::Result<Arc<dyn WebSocketGateway>>,
260 },
261 SocketIo {
262 register_fn: fn(&Container, &dyn Any) -> crate::Result<()>,
263 },
264}
265impl GatewayDef {
266 pub fn websocket<G: WebSocketGateway>(path: &'static str) -> Self {
267 Self {
268 path,
269 type_id: TypeId::of::<G>(),
270 provider: ProviderDef::of::<G>(),
271 kind: GatewayKind::WebSocket {
272 resolve_fn: |c| Ok(c.resolve::<G>()? as Arc<dyn WebSocketGateway>),
273 },
274 }
275 }
276 #[doc(hidden)]
277 pub fn socket_io<G: Injectable>(
278 path: &'static str,
279 register_fn: fn(&Container, &dyn Any) -> crate::Result<()>,
280 ) -> Self {
281 Self {
282 path,
283 type_id: TypeId::of::<G>(),
284 provider: ProviderDef::of::<G>(),
285 kind: GatewayKind::SocketIo { register_fn },
286 }
287 }
288 pub fn resolve(&self, container: &Container) -> crate::Result<Arc<dyn WebSocketGateway>> {
289 match self.kind {
290 GatewayKind::WebSocket { resolve_fn } => resolve_fn(container),
291 GatewayKind::SocketIo { .. } => Err(crate::exception::startup_error(format!(
292 "Socket.IO gateway {} cannot be mounted by an RFC 6455 application",
293 self.path
294 ))),
295 }
296 }
297 #[doc(hidden)]
298 pub fn is_websocket(&self) -> bool {
299 matches!(self.kind, GatewayKind::WebSocket { .. })
300 }
301 #[doc(hidden)]
302 pub fn register_socket_io(&self, container: &Container, handle: &dyn Any) -> crate::Result<()> {
303 match self.kind {
304 GatewayKind::WebSocket { .. } => Ok(()),
305 GatewayKind::SocketIo { register_fn } => register_fn(container, handle),
306 }
307 }
308}
309pub trait Gateway: Injectable {
310 #[doc(hidden)]
311 fn definition() -> GatewayDef;
312}
313
314pub struct ModuleDef {
315 type_id: TypeId,
316 type_name: &'static str,
317 metadata_fn: fn() -> ModuleMetadata,
318 pub(crate) route_log_fn: fn(),
319}
320impl ModuleDef {
321 pub fn of<M: Module + 'static>() -> Self {
322 Self {
323 type_id: TypeId::of::<M>(),
324 type_name: std::any::type_name::<M>(),
325 metadata_fn: M::register,
326 route_log_fn: || crate::log_module_routes::<M>(),
327 }
328 }
329}
330pub trait Module {
331 fn register() -> ModuleMetadata;
332}
333
334pub struct ModuleMetadata {
335 pub imports: Vec<ModuleDef>,
336 pub providers: Vec<ProviderDef>,
337 pub controllers: Vec<ControllerDef>,
338 pub event_handlers: Vec<EventHandlerDef>,
339 pub gateways: Vec<GatewayDef>,
340 exports: Vec<ProviderDependency>,
341 global: bool,
342}
343impl ModuleMetadata {
344 pub fn new() -> Self {
345 Self {
346 imports: vec![],
347 providers: vec![],
348 controllers: vec![],
349 event_handlers: vec![],
350 gateways: vec![],
351 exports: vec![],
352 global: false,
353 }
354 }
355 pub fn global() -> Self {
357 let mut metadata = Self::new();
358 metadata.global = true;
359 metadata
360 }
361 pub fn import<M: Module + 'static>(mut self) -> Self {
362 self.imports.push(ModuleDef::of::<M>());
363 self
364 }
365 pub fn provider<T: Injectable>(mut self) -> Self {
366 self.providers.push(ProviderDef::of::<T>());
367 self
368 }
369 pub fn provider_async_factory<T, Fut, E>(
370 mut self,
371 dependencies: Vec<ProviderDependency>,
372 factory: impl Fn(Arc<Container>) -> Fut + Send + Sync + 'static,
373 ) -> Self
374 where
375 T: Send + Sync + 'static,
376 Fut: Future<Output = std::result::Result<T, E>> + Send + 'static,
377 E: Debug + Send + 'static,
378 {
379 self.providers.push(ProviderDef::async_factory::<T, Fut, E>(
380 dependencies,
381 factory,
382 ));
383 self
384 }
385 pub fn controller<C: Controller + Injectable + 'static>(mut self) -> Self {
386 self.controllers.push(ControllerDef::of::<C>());
387 self
388 }
389 pub fn gateway<G: Gateway>(mut self) -> Self {
390 self.gateways.push(G::definition());
391 self
392 }
393 pub fn event_handler<H>(mut self) -> Self
394 where
395 H: RegisterableEventHandler + EventHandler<H::Event>,
396 {
397 self.event_handlers.push(EventHandlerDef::of::<H>());
398 self
399 }
400 pub fn event_handler_for<E, H>(mut self) -> Self
401 where
402 E: Clone + Send + Sync + 'static,
403 H: Injectable + EventHandler<E>,
404 {
405 self.event_handlers
406 .push(EventHandlerDef::for_event::<E, H>());
407 self
408 }
409 pub fn export<T: Send + Sync + 'static>(mut self) -> Self {
411 self.exports.push(ProviderDependency::of::<T>());
412 self
413 }
414}
415impl Default for ModuleMetadata {
416 fn default() -> Self {
417 Self::new()
418 }
419}
420
421struct ModuleNode {
422 type_id: TypeId,
423 type_name: &'static str,
424 metadata: ModuleMetadata,
425 imports: Vec<usize>,
426}
427struct ModuleGraph {
428 nodes: Vec<ModuleNode>,
429}
430#[derive(Clone, Copy)]
431enum ProviderSlot {
432 Provider(usize),
433 Controller(usize),
434 Gateway(usize),
435}
436#[derive(Clone, Copy)]
437struct ProviderRegistration {
438 module: usize,
439 slot: ProviderSlot,
440}
441
442impl ModuleGraph {
443 fn discover<M: Module + 'static>() -> crate::Result<Self> {
444 Self::discover_from(ModuleDef::of::<M>())
445 }
446 fn discover_from(root: ModuleDef) -> crate::Result<Self> {
447 fn visit(
448 def: ModuleDef,
449 graph: &mut Vec<ModuleNode>,
450 states: &mut HashMap<TypeId, u8>,
451 path: &mut Vec<&'static str>,
452 ) -> crate::Result<usize> {
453 match states.get(&def.type_id).copied() {
454 Some(2) => {
455 return Ok(graph
456 .iter()
457 .position(|node| node.type_id == def.type_id)
458 .expect("discovered module missing"));
459 }
460 Some(1) => {
461 let mut cycle = path.clone();
462 cycle.push(def.type_name);
463 return Err(crate::exception::startup_error(format!(
464 "circular module import: {}",
465 cycle.join(" -> ")
466 )));
467 }
468 _ => {}
469 }
470 states.insert(def.type_id, 1);
471 path.push(def.type_name);
472 let metadata = (def.metadata_fn)();
473 let index = graph.len();
474 graph.push(ModuleNode {
475 type_id: def.type_id,
476 type_name: def.type_name,
477 metadata,
478 imports: vec![],
479 });
480 let imports = std::mem::take(&mut graph[index].metadata.imports);
481 for import in imports {
482 let child = visit(import, graph, states, path)?;
483 graph[index].imports.push(child);
484 }
485 path.pop();
486 states.insert(def.type_id, 2);
487 Ok(index)
488 }
489 let mut nodes = vec![];
490 let mut states = HashMap::new();
491 let root_index = visit(root, &mut nodes, &mut states, &mut vec![])?;
492 let mut ordered = Vec::new();
494 let mut seen = HashSet::new();
495 fn order(
496 index: usize,
497 nodes: &Vec<ModuleNode>,
498 seen: &mut HashSet<usize>,
499 ordered: &mut Vec<usize>,
500 ) {
501 if !seen.insert(index) {
502 return;
503 }
504 for &child in &nodes[index].imports {
505 order(child, nodes, seen, ordered);
506 }
507 ordered.push(index);
508 }
509 order(root_index, &nodes, &mut seen, &mut ordered);
510 let mut remap = HashMap::new();
511 for (new, old) in ordered.iter().enumerate() {
512 remap.insert(*old, new);
513 }
514 let mut slots: Vec<Option<ModuleNode>> = nodes.into_iter().map(Some).collect();
515 let mut result = Vec::with_capacity(slots.len());
516 for old in ordered {
517 let mut node = slots[old]
518 .take()
519 .expect("module graph ordering repeated a node");
520 node.imports = node.imports.into_iter().map(|i| remap[&i]).collect();
521 result.push(node);
522 }
523 let _ = remap[&root_index];
524 Ok(Self { nodes: result })
525 }
526
527 fn definitions(&self) -> Vec<ProviderRegistration> {
528 let mut values = vec![];
529 for (module, node) in self.nodes.iter().enumerate() {
530 values.extend(
531 (0..node.metadata.providers.len()).map(|slot| ProviderRegistration {
532 module,
533 slot: ProviderSlot::Provider(slot),
534 }),
535 );
536 values.extend(
537 (0..node.metadata.controllers.len()).map(|slot| ProviderRegistration {
538 module,
539 slot: ProviderSlot::Controller(slot),
540 }),
541 );
542 for slot in 0..node.metadata.gateways.len() {
543 values.push(ProviderRegistration {
544 module,
545 slot: ProviderSlot::Gateway(slot),
546 });
547 }
548 }
549 values
550 }
551 fn def(&self, registration: ProviderRegistration) -> &ProviderDef {
552 match registration.slot {
553 ProviderSlot::Provider(index) => {
554 &self.nodes[registration.module].metadata.providers[index]
555 }
556 ProviderSlot::Controller(index) => {
557 &self.nodes[registration.module].metadata.controllers[index].provider
558 }
559 ProviderSlot::Gateway(index) => {
560 &self.nodes[registration.module].metadata.gateways[index].provider
561 }
562 }
563 }
564 fn local_types(&self, module: usize) -> HashSet<TypeId> {
565 self.definitions()
566 .into_iter()
567 .filter(|r| r.module == module)
568 .map(|r| self.def(r).type_id)
569 .collect()
570 }
571 fn preflight(&self, container: &Container) -> crate::Result<Vec<ProviderRegistration>> {
572 let definitions = self.definitions();
573 let mut by_type = HashMap::new();
574 for registration in &definitions {
575 let def = self.def(*registration);
576 if let Some(existing) = by_type.insert(def.type_id, *registration) {
577 if !container.has_pending_override(def.type_id)
578 && !container.was_overridden(def.type_id)
579 {
580 return Err(crate::exception::startup_error(format!(
581 "duplicate provider registration for {} in {} and {}",
582 def.type_name,
583 self.nodes[existing.module].type_name,
584 self.nodes[registration.module].type_name
585 )));
586 }
587 }
588 }
589 let locals: Vec<HashSet<TypeId>> =
590 (0..self.nodes.len()).map(|i| self.local_types(i)).collect();
591 let mut exports = vec![HashSet::new(); self.nodes.len()];
592 for index in 0..self.nodes.len() {
593 let node = &self.nodes[index];
594 for export in &node.metadata.exports {
595 let imported = node
596 .imports
597 .iter()
598 .any(|&child| exports[child].contains(&export.type_id));
599 if !locals[index].contains(&export.type_id) && !imported {
600 return Err(crate::exception::startup_error(format!(
601 "module {} cannot export {}: it is neither declared locally nor exported by a direct import",
602 node.type_name, export.type_name
603 )));
604 }
605 exports[index].insert(export.type_id);
606 }
607 }
608 let global_exports: HashSet<TypeId> = self
609 .nodes
610 .iter()
611 .enumerate()
612 .filter(|(_, node)| node.metadata.global)
613 .flat_map(|(index, _)| exports[index].iter().copied())
614 .collect();
615 let visible = |module: usize, type_id: TypeId| {
616 locals[module].contains(&type_id)
617 || global_exports.contains(&type_id)
618 || self.nodes[module]
619 .imports
620 .iter()
621 .any(|&child| exports[child].contains(&type_id))
622 };
623 for registration in &definitions {
624 let production = self.def(*registration);
625 let effective = container
626 .pending_override(production.type_id)
627 .unwrap_or(production);
628 for dependency in &effective.dependencies {
629 if dependency.type_id == TypeId::of::<crate::Logger>() {
630 continue;
631 }
632 if by_type.contains_key(&dependency.type_id) {
633 if !visible(registration.module, dependency.type_id) {
634 return Err(crate::exception::startup_error(format!(
635 "{} depends on {} but it is not visible in module {}; import and export its module",
636 effective.type_name,
637 dependency.type_name,
638 self.nodes[registration.module].type_name
639 )));
640 }
641 } else if !container.contains_type_id(dependency.type_id) {
642 return Err(crate::exception::startup_error(format!(
643 "missing provider at startup: {} depends on {} but no provider is registered",
644 effective.type_name, dependency.type_name
645 )));
646 }
647 }
648 }
649 for (module, node) in self.nodes.iter().enumerate() {
650 for handler in &node.metadata.event_handlers {
651 handler.assert_registered_or_declared(&locals[module])?;
652 if !visible(module, TypeId::of::<crate::EventBus>()) {
653 return Err(crate::exception::startup_error(format!(
654 "no provider registered for EventBus in module {}; import EventModule",
655 node.type_name
656 )));
657 }
658 }
659 }
660 let mut sorted = vec![];
661 let mut states = HashMap::new();
662 let mut stack = vec![];
663 fn schedule(
664 reg: ProviderRegistration,
665 graph: &ModuleGraph,
666 by_type: &HashMap<TypeId, ProviderRegistration>,
667 container: &Container,
668 states: &mut HashMap<TypeId, u8>,
669 stack: &mut Vec<&'static str>,
670 sorted: &mut Vec<ProviderRegistration>,
671 ) -> crate::Result<()> {
672 let def = graph.def(reg);
673 match states.get(&def.type_id).copied() {
674 Some(2) => return Ok(()),
675 Some(1) => {
676 let mut cycle = stack.clone();
677 cycle.push(def.type_name);
678 return Err(crate::exception::startup_error(format!(
679 "provider dependency cycle: {}",
680 cycle.join(" -> ")
681 )));
682 }
683 _ => {}
684 }
685 states.insert(def.type_id, 1);
686 stack.push(def.type_name);
687 let effective = container.pending_override(def.type_id).unwrap_or(def);
688 for dep in &effective.dependencies {
689 if let Some(next) = by_type.get(&dep.type_id) {
690 schedule(*next, graph, by_type, container, states, stack, sorted)?;
691 }
692 }
693 stack.pop();
694 states.insert(def.type_id, 2);
695 sorted.push(reg);
696 Ok(())
697 }
698 for registration in definitions {
699 schedule(
700 registration,
701 self,
702 &by_type,
703 container,
704 &mut states,
705 &mut stack,
706 &mut sorted,
707 )?;
708 }
709 Ok(sorted)
710 }
711 fn def_by_type(&self, type_id: TypeId) -> Option<&ProviderDef> {
712 self.definitions().into_iter().find_map(|r| {
713 let def = self.def(r);
714 (def.type_id == type_id).then_some(def)
715 })
716 }
717}
718
719async fn initialize_graph(graph: &ModuleGraph, container: &mut Container) -> crate::Result<()> {
720 let order = graph.preflight(container)?;
721 for registration in order {
722 let declared = graph.def(registration);
723 container.mark_provider_declared(declared.type_id);
724 if container.contains_type_id(declared.type_id)
725 || container.was_overridden(declared.type_id)
726 {
727 continue;
728 }
729 let start = Instant::now();
730 let override_def = container.take_pending_override(declared.type_id);
731 if override_def.is_some() {
732 container.mark_override_applied(declared.type_id);
733 }
734 let effective = override_def.as_ref().unwrap_or(declared);
735 let scoped_container =
736 container.scoped_for_provider(effective.type_name, &effective.dependencies);
737 let value = match (effective.build)(&scoped_container).await {
738 Ok(value) => value,
739 Err(error) => {
740 rollback_graph(graph, container).await;
741 return Err(error);
742 }
743 };
744 container.register_erased(effective.type_id, value.clone());
745 if let Err(error) = effective
746 .run_lifecycle(&value, "on_module_init", &effective.init_fn)
747 .await
748 {
749 rollback_graph(graph, container).await;
750 return Err(error);
751 }
752 if override_def.is_none() {
753 container.record_initialized_provider(declared.type_id);
754 }
755 crate::log_provider_initialized(effective.type_name, start.elapsed());
756 }
757 for node in &graph.nodes {
758 for handler in &node.metadata.event_handlers {
759 if let Err(error) = handler.register(container) {
760 rollback_graph(graph, container).await;
761 return Err(error);
762 }
763 }
764 crate::log_module_initialized(node.type_name, std::time::Duration::ZERO);
765 }
766 Ok(())
767}
768
769async fn bootstrap_graph(graph: &ModuleGraph, container: &Container) -> crate::Result<()> {
770 let order = graph.preflight(container)?;
771 let initialized: HashSet<TypeId> = container.initialized_provider_types().into_iter().collect();
772 for registration in order {
773 let def = graph.def(registration);
774 if initialized.contains(&def.type_id)
775 && !container.was_overridden(def.type_id)
776 && container.begin_provider_bootstrap(def.type_id)
777 {
778 if let Err(error) = def
779 .run_lifecycle_from_container(container, "on_bootstrap", &def.bootstrap_fn)
780 .await
781 {
782 rollback_graph(graph, container).await;
783 return Err(error);
784 }
785 }
786 }
787 Ok(())
788}
789
790async fn rollback_graph(graph: &ModuleGraph, container: &Container) {
791 for type_id in container.take_initialized_providers().into_iter().rev() {
792 if let Some(def) = graph.def_by_type(type_id) {
793 let _ = def
794 .run_lifecycle_from_container(container, "on_shutdown", &def.shutdown_fn)
795 .await;
796 }
797 }
798}
799
800pub async fn register_module<M: Module + 'static>(container: &mut Container) -> crate::Result<()> {
801 let graph = ModuleGraph::discover::<M>()?;
802 initialize_graph(&graph, container).await
803}
804pub async fn bootstrap_module<M: Module + 'static>(container: &Container) -> crate::Result<()> {
805 let graph = ModuleGraph::discover::<M>()?;
806 bootstrap_graph(&graph, container).await
807}
808pub async fn shutdown_module<M: Module + 'static>(container: &Container) -> crate::Result<()> {
809 let graph = ModuleGraph::discover::<M>()?;
810 let mut first = None;
811 for type_id in container.take_initialized_providers().into_iter().rev() {
812 if let Some(def) = graph.def_by_type(type_id) {
813 if let Err(error) = def
814 .run_lifecycle_from_container(container, "on_shutdown", &def.shutdown_fn)
815 .await
816 {
817 if first.is_none() {
818 first = Some(error);
819 }
820 }
821 }
822 }
823 first.map_or(Ok(()), Err)
824}
825
826pub async fn build_container<M: Module + 'static>() -> crate::Result<Container> {
827 build_container_with_overrides::<M>(ProviderOverrides::new()).await
828}
829#[doc(hidden)]
830pub async fn build_container_with_setup<M: Module + 'static>(
831 setup: impl FnOnce(&mut Container),
832) -> crate::Result<Container> {
833 crate::log_application_starting();
834 let mut container = Container::new();
835 setup(&mut container);
836 if let Err(error) = register_module::<M>(&mut container).await {
837 return Err(error);
838 }
839 if let Err(error) =
840 validate_module_providers::<M>(&container).and_then(|_| validate_gateway_paths::<M>())
841 {
842 let graph = ModuleGraph::discover::<M>()?;
843 rollback_graph(&graph, &container).await;
844 return Err(error);
845 }
846 if let Err(error) = bootstrap_module::<M>(&container).await {
847 return Err(error);
848 }
849 Ok(container)
850}
851pub async fn build_container_with_overrides<M: Module + 'static>(
852 overrides: ProviderOverrides,
853) -> crate::Result<Container> {
854 crate::log_application_starting();
855 let mut container = Container::new();
856 container.seed_overrides(overrides);
857 if let Err(error) = register_module::<M>(&mut container).await {
858 return Err(error);
859 }
860 if let Err(error) = container
861 .assert_no_unused_overrides()
862 .and_then(|_| validate_module_providers::<M>(&container))
863 .and_then(|_| validate_gateway_paths::<M>())
864 {
865 let graph = ModuleGraph::discover::<M>()?;
866 rollback_graph(&graph, &container).await;
867 return Err(error);
868 }
869 if let Err(error) = bootstrap_module::<M>(&container).await {
870 return Err(error);
871 }
872 Ok(container)
873}
874
875fn validate_gateway_paths<M: Module + 'static>() -> crate::Result<()> {
876 let graph = ModuleGraph::discover::<M>()?;
877 let mut paths = HashSet::new();
878 for node in &graph.nodes {
879 for gateway in &node.metadata.gateways {
880 if !gateway.path.starts_with('/') {
881 return Err(crate::exception::startup_error(format!(
882 "websocket gateway path must start with '/': {}",
883 gateway.path
884 )));
885 }
886 if !paths.insert((gateway.path, gateway.is_websocket())) {
887 return Err(crate::exception::startup_error(format!(
888 "duplicate gateway path: {}",
889 gateway.path
890 )));
891 }
892 }
893 }
894 Ok(())
895}
896pub fn validate_module_providers<M: Module + 'static>(container: &Container) -> crate::Result<()> {
897 let graph = ModuleGraph::discover::<M>()?;
898 graph.preflight(container)?;
899 for registration in graph.definitions() {
900 graph.def(registration).assert_registered(container)?;
901 }
902 Ok(())
903}
904
905pub fn register_module_controllers<M: Module + 'static>(any: &mut dyn Any) {
906 if let Ok(graph) = ModuleGraph::discover::<M>() {
907 for node in graph.nodes {
908 for controller in &node.metadata.controllers {
909 (controller.register_fn)(any);
910 }
911 }
912 }
913}
914#[cfg(feature = "openapi")]
915#[doc(hidden)]
916pub fn visit_module_openapi_routes<M: Module + 'static>(
917 visitor: &mut impl FnMut(&crate::openapi::OpenApiRouteDef),
918) {
919 if let Ok(graph) = ModuleGraph::discover::<M>() {
920 for node in graph.nodes {
921 for controller in &node.metadata.controllers {
922 for route in (controller.openapi_routes_fn)() {
923 visitor(route);
924 }
925 }
926 }
927 }
928}
929pub fn visit_module_gateways<M: Module + 'static>(visitor: &mut impl FnMut(&GatewayDef)) {
930 if let Ok(graph) = ModuleGraph::discover::<M>() {
931 let mut seen = HashSet::new();
932 for node in graph.nodes {
933 for gateway in &node.metadata.gateways {
934 if seen.insert(gateway.type_id) {
935 visitor(gateway);
936 }
937 }
938 }
939 }
940}