diff --git a/Cargo.toml b/Cargo.toml index de96afe..aa71fd2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -49,8 +49,8 @@ tcp-transport = ["dep:tokio"] [dependencies] a3s-acl = { version = "0.2.1", optional = true } a3s-boot-macros = { version = "0.1.2", path = "macros", optional = true } -a3s-event = { version = "0.3.0", path = "../event", default-features = false, optional = true } -a3s-lane = { version = "0.5.1", path = "../lane", default-features = false, optional = true } +a3s-event = { version = "0.3.0", default-features = false, optional = true } +a3s-lane = { version = "0.5.1", default-features = false, optional = true } async-nats = { version = "0.49.1", default-features = false, optional = true } axum = { version = "0.8", features = ["ws"], optional = true } bytes = { version = "1", optional = true } diff --git a/ROADMAP.md b/ROADMAP.md index 5748c5c..27a6c82 100644 --- a/ROADMAP.md +++ b/ROADMAP.md @@ -41,10 +41,18 @@ Implemented today: factories, singleton provider lifecycle hooks, lookup, `FromModuleRef` auto-wired provider factories, named or optional dependency resolution, `ProviderRef` lazy provider handles for forward-reference-style - dependencies, fresh resolution contexts, and dynamic injectable creation. + dependencies, declared eager dependency graphs with automatic request-scope + propagation, public `ContextId` / `ContextIdFactory` identities, fresh or + caller-shared resolution contexts, per-inquirer transient reuse, and dynamic + injectable creation. +- DI-managed application-wide enhancers corresponding to Nest's `APP_GUARD`, + `APP_PIPE`, `APP_INTERCEPTOR`, and `APP_FILTER`, with typed provider markers + for HTTP plus protocol-specific WebSocket and transport variants, automatic + scope bubbling, and per-invocation contextual resolution. - `TestingModule` with module and provider overrides, async provider-aware `compile_async`, and typed HTTP, WebSocket, and transport pipeline overrides - for guards, interceptors, exception filters, and pipes. + for guards, interceptors, exception filters, and pipes, including + provider-backed application enhancers. - `ControllerDefinition` and `RouteDefinition` for HTTP route groups, including specificity-aware path params, catch-all route params, and Nest-style ALL method routes with exact-method precedence. @@ -62,8 +70,9 @@ Implemented today: Nest-style custom route/controller metadata and `#[http_code]` for Nest-style response status metadata, `#[cache_key]` / `#[cache_ttl]` for cache response metadata, `#[header]` for response headers, and `#[redirect]` for redirect - responses. `#[injectable]` implements `FromModuleRef` for unit structs and - named-field structs whose dependencies are `Arc` or `Option>`, with + responses. `#[injectable]` implements `FromModuleRef` and dependency metadata + for unit structs and named-field structs whose dependencies are `Arc`, + `Option>`, `ProviderRef`, or `Option>`, with `#[inject("token")]` for named provider lookup. `#[module]` implements `Module` from Nest-style metadata lists for imports, providers, controllers, routes, gateways, message controllers, exports, @@ -87,6 +96,14 @@ Implemented today: - Nest-style generic pipeline macros: `#[use_guard]`, `#[use_interceptor]`, `#[use_filter]`, and `#[use_pipe]` at HTTP controller/route, WebSocket gateway/subscription, and message controller/pattern scope. +- A shared Nest-style `CallHandler` abstraction for the protocol-specific HTTP, + WebSocket, and message-transport around-interceptor traits. Interceptors can transform or + recover downstream errors, short-circuit without calling `next`, and apply + sequential retries or runtime-specific timeouts around reusable + `next.handle()` calls. Existing + `before` / `short_circuit` / `after` implementations remain compatible + through the default `intercept(...)` methods, and exception filters receive + only errors that remain unrecovered after the interceptor chain. - Nest-style catch-filter targeting with `#[catch]`, `BootErrorKind`, `catch_errors(...)`, `with_catch_filter(...)`, and `use_global_catch_filter(...)`, plus protocol-specific global WebSocket and @@ -158,23 +175,27 @@ Implemented today: module-level forward imports for deliberate circular module relationships, and contextual module import cycle diagnostics. - Provider lifecycle scopes with default singleton providers, request-scoped - providers cached per in-process request context, transient providers built per - resolution, async singleton provider factories awaited during async graph - build, order-independent singleton provider graph initialization, - request-time lookup through `BootRequest`, singleton/transient/request-scoped - provider dependency cycle diagnostics, and singleton provider startup/shutdown - hooks for module init, application bootstrap, module destroy, before - application shutdown, and application shutdown, including OS signal labels - from shutdown hooks. + providers cached per `ContextId`, transient providers reused per inquirer + within a context, async singleton provider factories awaited during async + graph build, automatic request-scope propagation through declared eager + dependency graphs, order-independent singleton provider graph initialization, + request-time lookup and automatic context identity through `BootRequest`, + singleton/transient/request-scoped provider dependency cycle diagnostics, and + singleton provider startup/shutdown hooks for module init, application + bootstrap, module destroy, before application shutdown, and application + shutdown, including OS signal labels from shutdown hooks. - Provider aliases that mirror Nest custom provider `useExisting` semantics and preserve target provider scope. - Lazy `ProviderRef` handles that mirror the useful part of Nest `forwardRef(...)`: explicit delayed provider resolution without weakening - normal cycle diagnostics. + normal cycle diagnostics. Lazy edges intentionally do not propagate request + scope; callers use a captured request context or explicit `resolve(...)`. - Module-level forward imports that mirror Nest `forwardRef(() => Module)` for explicit circular module relationships while preserving normal import-cycle diagnostics. -- Request-scoped route/controller handler factories through `*_scoped` helpers. +- Request-scoped route/controller handler factories through `*_scoped` helpers, + plus provider-backed macro controllers selected automatically for request, + transient, and dependency-bubbled controller graphs. - Middleware with request mutation, short-circuit responses, global/module/ controller/route scopes, `MiddlewareConsumer::apply(...).for_routes(...)` include/exclude rules, filter integration for errors, and adapter validation @@ -189,10 +210,11 @@ Implemented today: WebSocket route registration. - Microservice transports with adapter-neutral `TransportMessage` / `TransportReply`, request-response and event-only message patterns, - provider-backed handlers, validation helpers, transport pipes/guards/ - interceptors, application-wide protocol guards/interceptors/pipes, local and - global protocol exception filters, Nest-style message macros with controller- - and pattern-scoped pipeline decorators, an in-process transport, and an optional + per-dispatch scoped handlers and message controllers, validation helpers, + transport pipes/guards/interceptors, application-wide protocol + guards/interceptors/pipes, local and global protocol exception filters, + Nest-style message macros with controller- and pattern-scoped pipeline + decorators, an in-process transport, and an optional TCP transport for newline-delimited JSON message frames plus an optional Redis Pub/Sub transport and optional NATS request/reply and event subjects plus optional MQTT request/reply and event topics plus optional RabbitMQ @@ -255,7 +277,8 @@ Implemented today: 1. Parameter extraction macros 2. OpenAPI metadata and generator 3. Validation pipeline (implemented) -4. Module encapsulation, dynamic modules, and provider lifecycle scopes (implemented) +4. Module encapsulation, dynamic modules, provider lifecycle scopes, and + application enhancers (implemented) 5. Middleware (implemented) 6. WebSocket gateways (implemented) 7. Microservice transports (implemented) @@ -458,8 +481,8 @@ Tasks: - Add a small `Validate` trait in core or a `validation` feature. (Implemented in core) - Integrate validation after DTO extraction and before handler invocation. - (Implemented with route validation hooks that run after guards, interceptor - `before` hooks, and request pipes for routes carrying validation metadata) + (Implemented with route validation hooks inside the around-interceptor chain, + after request pipes and before the handler) - Support explicit validation pipe composition for projects that prefer a third party crate such as `garde` or `validator`. (Implemented through ordinary `Pipe` composition plus explicit `Validate` implementations) @@ -501,7 +524,7 @@ Acceptance: - Validation does not run for raw handlers unless explicitly configured. (Covered) -## Milestone 4: Module Encapsulation, Dynamic Modules, And Provider Lifecycle Scopes +## Milestone 4: Module Encapsulation, Dynamic Modules, Provider Scopes, And Application Enhancers Nest equivalent: @@ -517,6 +540,8 @@ Nest equivalent: - forward-reference-style provider dependencies - module-level `forwardRef(() => Module)` - lazy module loading +- DI-managed global enhancers through `APP_GUARD`, `APP_PIPE`, + `APP_INTERCEPTOR`, and `APP_FILTER` Current gap: @@ -531,20 +556,50 @@ shutdown, and application shutdown hooks. Managed HTTP and microservice hosts can enable Nest-style shutdown hooks so OS signals close the application through the same signal-aware lifecycle phases. Request-scoped handler factories rebuild route/controller state from the current -request's module context. Provider aliases let one token delegate to an existing +request's module context. Macro controllers automatically use provider-backed +routes when their own scope or an eager dependency graph is contextual. +Singleton dependency trees automatically inherit request context without +changing their declared cache scope. Transient dependencies are cached per +inquirer inside a `ContextId`: one consumer reuses its transient instance, while +different consumers or contexts remain isolated. Direct contextless transient +lookups remain fresh per call for compatibility. `#[injectable]` emits typed, +named, optional, and lazy dependency metadata; hand-written factories can +declare the same graph explicitly without Boot probing factories and repeating +their side effects. Provider aliases let one token delegate to an existing provider token without changing the target provider's lifecycle scope. Explicit module-level forward imports can model deliberate circular module relationships, while ordinary module import cycles still report the active module chain during sync and async application graph builds. Singleton provider factories are initialized after all module provider tokens are registered, so factories can depend on providers declared later in the same module. -`LazyModuleLoader` can load provider-only module graphs on demand, reuse eagerly -registered modules, and resolve async singleton factories through -`load_async(...)`. `ModuleRef` can resolve providers in a fresh temporary -request context and dynamically create unregistered `FromModuleRef` values. +`LazyModuleLoader` can register complete provider-only module graphs on demand, +including forward imports, reuse eagerly registered modules, and resolve async +singleton factories through `load_async(...)`. Late global lazy modules are +rejected before factory execution because mutating global visibility would +invalidate already initialized singleton plans. `ModuleRef` can resolve +providers in a fresh temporary context or reuse a caller-supplied `ContextId`, +bind a context with `context_scope(...)`, and dynamically create unregistered +`FromModuleRef` values. `ContextIdFactory` creates standalone contexts and +discovers the identity attached automatically to each HTTP request. `ProviderRef` can capture a module context and resolve a provider lazily, which gives Rust code an explicit forward-reference-style escape hatch while -keeping ordinary provider cycles diagnostic. +keeping ordinary provider cycles diagnostic. It can also resolve through a +caller-supplied context while preserving the original inquirer. Contextual +factories use non-owning context views so cached lazy refs cannot retain their +own request cache. Provider construction uses panic-safe, branch-local +resolution paths plus single-flight cache slots and cross-thread wait-cycle +diagnostics. Provider definitions can also carry typed application-enhancer +markers that map to Nest's `APP_GUARD`, `APP_PIPE`, `APP_INTERCEPTOR`, and +`APP_FILTER` provider registrations. HTTP, WebSocket, and transport variants +resolve from the declaring module against the current request or message +`ContextId`, so private dependencies remain injectable without widening module +exports. HTTP requests, WebSocket messages, and transport messages each receive +an isolated context; sequential interceptor retries reuse the original context. +Pipelines compose builder globals before provider globals, provider globals +retain module/provider declaration order, and handler-local hooks remain last; +exception filters keep their existing local-to-global handling direction. Lazy +modules reject `APP_*` providers before any factory executes because existing +application pipelines cannot be mutated after build. Tasks: @@ -564,15 +619,32 @@ Tasks: (Implemented; root scopes and global exports are visible to the host) - Add provider lifecycle scopes comparable to Nest singleton, request, and transient providers. (Implemented) -- Make request-scoped providers reuse one instance per request context, +- Make request-scoped providers reuse one instance per resolution context, including dependencies resolved inside request-scoped provider factories. (Implemented) +- Add public `ContextId` / `ContextIdFactory` APIs plus `ModuleRef` and + `ProviderRef` context-aware resolution, including named and optional + variants. (Implemented) +- Reuse transient providers per inquirer within one context while keeping + different consumers and contexts isolated and direct contextless transient + lookup fresh. (Implemented) +- Attach one discoverable `ContextId` to each dispatched HTTP request. + (Implemented through `BootRequest::context_id()` and + `ContextIdFactory::get_by_request(...)`) +- Propagate request context through eager singleton and transient dependency + graphs without rewriting their declared scope. (Implemented through + `ProviderDependency` metadata generated by `#[injectable]` or declared on + `ProviderDefinition`) +- Reject async factories and singleton lifecycle-hook providers whose effective + dependency graph is contextual. (Implemented during graph finalization) - Add singleton provider lifecycle hooks for init, bootstrap, module destroy, before application shutdown, and application shutdown. (Implemented) - Add Nest-style shutdown hook enabling for managed HTTP and microservice hosts. (Implemented with `enable_shutdown_hooks(...)` and default `SIGINT`/`SIGTERM` support) -- Add request-scoped route/controller handler factories. (Implemented) +- Add request-scoped route/controller and transport message-handler factories. + (Implemented, with automatic provider-backed routes and message patterns for + contextual macro controllers) - Add provider aliases comparable to Nest `useExisting`. (Implemented) - Add lazy provider handles comparable to the useful provider side of Nest `forwardRef(...)`. (Implemented with `ProviderRef`) @@ -585,7 +657,25 @@ Tasks: dependency cycles. (Implemented) - Add contextual diagnostics for module import cycles. (Implemented) - Add order-independent singleton provider graph initialization. (Implemented) -- Add provider-only lazy module loading with cached module refs. (Implemented) +- Add provider-only lazy module loading with cached module refs. (Implemented + with full-graph forward-import planning; newly loaded globals are rejected) +- Add typed provider-backed application enhancers corresponding to Nest + `APP_GUARD`, `APP_PIPE`, `APP_INTERCEPTOR`, and `APP_FILTER`, including + WebSocket and transport variants. (Implemented with `app_*::()` and + `with_app_*::()` provider APIs) +- Resolve application enhancers from their declaring module with automatic + scope bubbling and the active HTTP request, WebSocket message, or transport + message `ContextId`; preserve that id across sequential interceptor retries. + (Implemented) +- Preserve deterministic builder-global, provider-global, and local-hook order + without widening module visibility. (Implemented) +- Reject `APP_*` providers in lazy-loaded modules before provider factories + execute. (Implemented) +- Make testing pipeline overrides replace provider-backed components by their + original concrete type and retain original enhancer markers across provider + overrides. (Implemented) +- Add custom context-id selection strategies, durable provider subtrees, and + request-object registration for arbitrary context ids. (Future work) Acceptance: @@ -596,24 +686,69 @@ Acceptance: - Duplicate-provider checks respect module scope. (Covered) - Existing simple module examples continue to work or have a documented migration. (Covered; root module providers remain visible through `BootApplication::get`) -- Transient providers are rebuilt for every resolution. (Covered) -- Request-scoped providers are cached per request and are isolated from other - requests. (Covered) +- Direct contextless transient lookups remain fresh, while repeated transient + injection into one inquirer is reused and different inquirers or context ids + receive distinct instances. (Covered) +- Request-scoped providers are cached per context and are isolated from other + contexts; a caller can reuse that cache through `resolve_with_context(...)` + or `ModuleRef::context_scope(...)`. (Covered) +- Singleton providers with eager request-scoped dependencies are cached once per + context, contextual transients follow per-inquirer reuse, and opaque factories + cannot capture a startup-only request instance. (Covered) - Singleton provider lifecycle hooks run with module lifecycle hooks, reject request/transient provider scopes, and receive explicit shutdown signal labels through signal-aware close helpers. (Covered) - Managed HTTP and microservice hosts can close through signal-aware lifecycle hooks when a configured shutdown signal wins the serve race. (Covered) -- Request-scoped controller handlers are rebuilt for each request and share the - same request-scoped provider cache as `BootRequest::get(...)`. (Covered) +- Request-scoped and dependency-bubbled controller handlers are rebuilt for + each request and share the same request-scoped provider cache as + `BootRequest::get(...)`. (Covered) +- Request-scoped, transient, and dependency-bubbled message controllers resolve + once per transport dispatch, reuse the same controller and `ContextId` across + interceptor retries, and isolate later messages. (Covered) - Provider aliases resolve the same singleton instance, preserve request-scoped resolution, and reject alias cycles with contextual errors. (Covered) - `ProviderRef` resolves lazily, can break an intentional singleton dependency cycle, supports named and optional macro injection, and preserves a - captured request scope. (Covered) -- `ModuleRef::resolve(...)` creates a fresh resolution context for - request-scoped dependency caches, and `ModuleRef::create(...)` can instantiate - `FromModuleRef` values without registering them. (Covered) + captured request scope and inquirer; explicit context resolution reuses a + caller-supplied id. (Covered) +- Context caches release providers that contain lazy refs, construction paths + recover after caught panics, parallel branches do not share sibling frames, + and concurrent single-flight waits distinguish real cycles from converging + dependency graphs. (Covered) +- `ModuleRef::resolve(...)` creates a fresh resolution context, + `resolve_with_context(...)` and its named/optional variants reuse an explicit + context, and `ModuleRef::create(...)` can instantiate `FromModuleRef` values + with distinct synthetic inquirers without registering them. (Covered) +- Every dispatched HTTP request exposes one distinct context identity through + `BootRequest::context_id()` and `ContextIdFactory::get_by_request(...)`. + (Covered) +- Provider-backed HTTP application guards, pipes, interceptors, and filters + share the request context across direct and module routes, and application + filters can handle early routing errors. (Covered) +- Provider-backed HTTP enhancers also apply to framework-provided routes, and + application filters can handle unmatched application routes. (Covered) +- Provider-backed WebSocket and transport application enhancers resolve from a + fresh message context, share one exact context within a dispatch, and isolate + later messages on the same connection or pattern. (Covered) +- Sequential HTTP, WebSocket, and transport interceptor retries reuse the + original request or message context. (Covered) +- Application enhancers resolve through their declaring module, can inject its + private providers, retain deterministic builder/provider/local order, and do + not make those dependencies visible to unrelated handler modules. (Covered) +- Typed custom and named provider definitions can retain application-enhancer + markers. (Covered) +- Alias, value, and async provider definitions can retain the same markers, and + contextual async or lifecycle enhancer graphs are rejected before factory + execution. (Covered) +- Testing pipeline overrides replace provider-backed application enhancers by + original concrete type, provider overrides retain the original marker slot, + and lazy module loading rejects `APP_*` providers before factories run. + (Covered) +- Graph-level enhancer replacement uses `override_provider(...)` with the + original typed or named token, runs before singleton construction, suppresses + the original provider lifecycle, and retains its pipeline marker; pipeline + overrides remain type-safe component-view replacements. (Covered) - Transient and request-scoped provider cycles report the active token chain. (Covered) - Module import cycles report the active module chain during sync and async @@ -624,10 +759,14 @@ Acceptance: - Singleton provider factories can resolve dependencies declared later in the same module, including sync factories that depend on async-built singletons in async builds. (Covered) +- Declared async dependencies are seeded before their consumers across local, + imported, forward-imported, and global module edges; cycles and missing + dependencies fail before factory execution. (Covered) - Lazy module loading returns cached module refs, reuses eagerly imported - modules, resolves imports/exports, supports async singleton providers through - `load_async(...)`, and does not register controllers, routes, gateways, - middleware, message patterns, or lifecycle hooks. (Covered) + modules, resolves imports/exports/forward imports, supports async singleton + providers through `load_async(...)`, rejects new late globals, and does not + register controllers, routes, gateways, middleware, message patterns, or + lifecycle hooks. (Covered) ## Milestone 5: Middleware @@ -639,16 +778,17 @@ Nest equivalent: Tasks: -- Add middleware trait that can inspect/mutate `BootRequest` before guards, - interceptor `before` hooks, and pipes. (Implemented) +- Add middleware trait that can inspect/mutate `BootRequest` before guards, the + around-interceptor chain, and pipes. (Implemented) - Allow middleware to short-circuit with `BootResponse`. (Implemented through `MiddlewareOutcome::Respond`) - Support global, module/controller, and route-scoped registration. (Implemented) - Add Nest-style `MiddlewareConsumer` with `apply`, `exclude`, `for_routes`, and `for_all_routes` for module-scoped route selection. (Implemented) -- Preserve order: middleware, guards, interceptor `before` hooks, pipes, - validation, handler, interceptor `after` hooks, filters. (Covered) +- Preserve order: middleware, guards, the nested around-interceptor chain, + pipes, validation, and the handler; legacy `after` hooks unwind in reverse, + and filters receive only unrecovered errors. (Covered) - Ensure adapter-level request validation remains before middleware. (Covered for Axum) @@ -724,9 +864,12 @@ Acceptance: - Gateway and subscription guards/interceptors/pipes plus application-wide gateway guards/interceptors/pipes run in Nest-style deterministic order. (Covered) +- Gateway interceptors can use reusable `CallHandler` access to short-circuit, + recover or transform errors, and retry the downstream pipeline while + retaining legacy `before` / `after` compatibility. (Covered) - Gateway exception filters can handle matching message dispatch errors and map - them to outbound WebSocket messages, including application-wide WebSocket - filters. (Covered) + unrecovered errors to outbound WebSocket messages, including application-wide + WebSocket filters. (Covered) - Gateways can track active connection ids, join/leave rooms, and deliver direct, room-scoped, or gateway-wide messages to adapter-backed connections. (Covered) @@ -748,6 +891,10 @@ Tasks: `MessagePatternDefinition`, `Module::message_patterns`, `BootApplicationBuilder::message_pattern`, `#[message_controller]`, `#[message_pattern]`, and `#[event_pattern]`) +- Add dispatch-scoped request/event pattern factories and automatically select + them for request-scoped, transient, or dependency-bubbled macro message + controllers. (Implemented with `request_scoped(...)`, `event_scoped(...)`, + and provider-backed macro patterns) - Add field-level payload binding comparable to Nest `@Payload("field")`. (Implemented with `#[payload("field")]` and `TransportMessage::data_field_as(...)` helpers) @@ -767,17 +914,24 @@ Acceptance: - A module can register message handlers independently from HTTP routes. (Covered) - Message handlers can use providers and validation. (Covered) +- Scoped message handlers can inject private declaring-module providers, reuse + one handler/controller across an interceptor retry, and receive a fresh + context on the next request-response or event message. (Covered) - Message handlers can bind individual payload fields, including optional fields, defaults, and parse pipes. (Covered) - Application-wide transport guards/interceptors/pipes run before and around pattern-scoped hooks in Nest-style deterministic order. (Covered) +- Transport interceptors can use reusable `CallHandler` access to short-circuit, + recover or transform errors, and retry the downstream pipeline while + retaining legacy `before` / `after` compatibility. Event-pattern replies are + discarded even when an interceptor short-circuits. (Covered) - Tests cover request-response and event-only patterns. (Covered) - Handler errors preserve `BootError` HTTP exception semantics across TCP, Redis, NATS, MQTT, RabbitMQ, Kafka, and gRPC request-response transports. (Covered) - Transport exception filters can handle matching message dispatch errors and - map them to request-response replies or handled event errors, including - application-wide transport filters. (Covered) + map unrecovered errors to request-response replies or handled event errors, + including application-wide transport filters. (Covered) ## Milestone 8: Technique Modules @@ -905,8 +1059,9 @@ Acceptance: through Nest-style `#[body("field")]` arguments. (Covered) - Testing utilities can compile Nest-style testing modules, override imported modules and providers before controllers are built, override HTTP, - WebSocket, and transport pipeline components, resolve providers, and dispatch - in-process requests. (Covered) + WebSocket, and transport pipeline components including provider-backed + application enhancers, resolve providers, and dispatch in-process requests. + (Covered) - Discovery and reflector utilities can snapshot modules, module graph edges, provider tokens, exports, HTTP route metadata, WebSocket gateways, and message patterns from a built application. (Covered) diff --git a/macros/src/controller/handlers.rs b/macros/src/controller/handlers.rs index 614ee46..d364b19 100644 --- a/macros/src/controller/handlers.rs +++ b/macros/src/controller/handlers.rs @@ -8,12 +8,13 @@ use syn::{Ident, LitStr, Result}; pub(super) fn raw_or_json_request_handler( method_ident: &Ident, input: RouteMethodInput, + receiver: &proc_macro2::TokenStream, ) -> Result { let controller_name = format_ident!("__a3s_boot_{}", method_ident); Ok(match input.into_legacy_arg()? { Some(MethodArg { ident, ty, .. }) => quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |#ident: #ty| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { #controller_name.#method_ident(#ident).await } @@ -22,7 +23,7 @@ pub(super) fn raw_or_json_request_handler( }, None => quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { #controller_name.#method_ident().await } @@ -35,12 +36,13 @@ pub(super) fn raw_or_json_request_handler( pub(super) fn json_body_handler( method_ident: &Ident, input: MethodArg, + receiver: &proc_macro2::TokenStream, ) -> proc_macro2::TokenStream { let controller_name = format_ident!("__a3s_boot_{}", method_ident); let MethodArg { ident, ty, .. } = input; quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |#ident: #ty| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { #controller_name.#method_ident(#ident).await } @@ -52,6 +54,7 @@ pub(super) fn json_body_handler( pub(super) fn extracted_raw_handler( method_ident: &Ident, input: RouteMethodInput, + receiver: &proc_macro2::TokenStream, ) -> Result { let controller_name = format_ident!("__a3s_boot_{}", method_ident); let ExtractedArguments { @@ -64,7 +67,7 @@ pub(super) fn extracted_raw_handler( Ok(quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { @@ -83,6 +86,7 @@ pub(super) fn extracted_json_response_handler( method_ident: &Ident, input: RouteMethodInput, status: proc_macro2::TokenStream, + receiver: &proc_macro2::TokenStream, ) -> Result { let controller_name = format_ident!("__a3s_boot_{}", method_ident); let ExtractedArguments { @@ -95,7 +99,7 @@ pub(super) fn extracted_json_response_handler( Ok(quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { @@ -117,6 +121,7 @@ pub(super) fn rendered_view_handler( input: RouteMethodInput, view: &LitStr, status: proc_macro2::TokenStream, + receiver: &proc_macro2::TokenStream, ) -> Result { let controller_name = format_ident!("__a3s_boot_{}", method_ident); @@ -130,7 +135,7 @@ pub(super) fn rendered_view_handler( let apply_response = response_passthrough_apply(response_passthrough); return Ok(quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { @@ -153,7 +158,7 @@ pub(super) fn rendered_view_handler( Ok(match input.into_legacy_arg()? { Some(MethodArg { ident, ty, .. }) => quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { @@ -170,7 +175,7 @@ pub(super) fn rendered_view_handler( }, None => quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { @@ -190,6 +195,7 @@ pub(super) fn rendered_view_handler( pub(super) fn extracted_sse_handler( method_ident: &Ident, input: RouteMethodInput, + receiver: &proc_macro2::TokenStream, ) -> Result { let controller_name = format_ident!("__a3s_boot_{}", method_ident); let ExtractedArguments { @@ -207,7 +213,7 @@ pub(super) fn extracted_sse_handler( Ok(quote! { { - let #controller_name = ::std::sync::Arc::clone(&self); + let #controller_name = ::std::sync::Arc::clone(#receiver); move |__a3s_boot_request: ::a3s_boot::BootRequest| { let #controller_name = ::std::sync::Arc::clone(&#controller_name); async move { diff --git a/macros/src/controller/mod.rs b/macros/src/controller/mod.rs index c4ef908..5069d97 100644 --- a/macros/src/controller/mod.rs +++ b/macros/src/controller/mod.rs @@ -47,6 +47,7 @@ pub(crate) fn expand_controller( let self_ty = item_impl.self_ty.clone(); let mut routes = Vec::new(); + let mut provider_routes = Vec::new(); let mut errors: Option = None; let (impl_attrs, impl_decorator_errors) = expand_apply_decorators_attrs(&item_impl.attrs); for error in impl_decorator_errors { @@ -283,9 +284,9 @@ pub(crate) fn expand_controller( let validation_options = route_validation.enabled_options(controller_validation); let validation_skipped = route_validation.skip; match route_registration( - route, + route.clone(), method, - input, + input.clone(), validation_options, validation_skipped, &metadata_specs, @@ -298,8 +299,32 @@ pub(crate) fn expand_controller( version_specs.as_ref(), serialization_specs.as_ref(), &openapi_specs, + ControllerRouteKind::Instance, ) { - Ok(registration) => routes.push(registration), + Ok(registration) => { + routes.push(registration); + match route_registration( + route, + method, + input, + validation_options, + validation_skipped, + &metadata_specs, + &cache_specs, + http_code.as_ref(), + &response_specs, + render_spec.as_ref(), + &pipeline_specs, + host_specs.as_ref(), + version_specs.as_ref(), + serialization_specs.as_ref(), + &openapi_specs, + ControllerRouteKind::Provider, + ) { + Ok(registration) => provider_routes.push(registration), + Err(error) => push_error(&mut errors, error), + } + } Err(error) => push_error(&mut errors, error), } } @@ -344,10 +369,49 @@ pub(crate) fn expand_controller( )* Ok(__a3s_boot_controller) } + + pub fn provider_controller() -> ::a3s_boot::Result<::a3s_boot::ControllerDefinition> + where + Self: ::std::marker::Send + ::std::marker::Sync + 'static, + { + let mut __a3s_boot_controller = + ::a3s_boot::ControllerDefinition::new(#prefix)?; + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_openapi; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_metadata?; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_cache; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_pipeline; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_host?; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_version; + )* + #( + __a3s_boot_controller = __a3s_boot_controller.#controller_serialization; + )* + #( + __a3s_boot_controller = #provider_routes; + )* + Ok(__a3s_boot_controller) + } } }) } +#[derive(Clone, Copy, PartialEq, Eq)] +enum ControllerRouteKind { + Instance, + Provider, +} + fn take_route_attrs(attrs: &[Attribute]) -> (Vec, Vec, Vec) { let mut clean_attrs = Vec::new(); let mut routes = Vec::new(); @@ -386,6 +450,7 @@ fn route_registration( version_spec: Option<&VersionSpec>, serialization_spec: Option<&SerializationSpec>, openapi_specs: &[RouteOpenApiSpec], + controller_route_kind: ControllerRouteKind, ) -> Result { if method.sig.asyncness.is_none() { return Err(syn::Error::new_spanned( @@ -395,6 +460,10 @@ fn route_registration( } let method_ident = &method.sig.ident; + let receiver = match controller_route_kind { + ControllerRouteKind::Instance => quote! { &self }, + ControllerRouteKind::Provider => quote! { &__a3s_boot_provider_controller }, + }; let explicit_status = route.args.explicit_status(http_code)?; let status = status_value(explicit_status)?; let path = route.args.path.clone(); @@ -435,7 +504,8 @@ fn route_registration( let route_definition = if let Some(render_spec) = render_spec { let builder = route.kind.raw_builder_ident(); let view = &render_spec.view; - let handler = rendered_view_handler(method_ident, input.clone(), view, status.clone())?; + let handler = + rendered_view_handler(method_ident, input.clone(), view, status.clone(), &receiver)?; quote! { ::a3s_boot::RouteDefinition::#builder(#path, #handler)? .with_response(#status, ::a3s_boot::OpenApiResponse::description("Success")) @@ -457,9 +527,9 @@ fn route_registration( )); } let handler = if input.has_extractors() { - extracted_sse_handler(method_ident, input)? + extracted_sse_handler(method_ident, input, &receiver)? } else { - raw_or_json_request_handler(method_ident, input)? + raw_or_json_request_handler(method_ident, input, &receiver)? }; quote! { ::a3s_boot::RouteDefinition::sse(#path, #handler)? @@ -474,9 +544,9 @@ fn route_registration( } let builder = route.kind.raw_builder_ident(); let handler = if input.has_extractors() { - extracted_raw_handler(method_ident, input)? + extracted_raw_handler(method_ident, input, &receiver)? } else { - raw_or_json_request_handler(method_ident, input)? + raw_or_json_request_handler(method_ident, input, &receiver)? }; quote! { ::a3s_boot::RouteDefinition::#builder(#path, #handler)? @@ -485,8 +555,12 @@ fn route_registration( RouteFlavor::JsonRequest => { if input.has_extractors() { let builder = route.kind.raw_builder_ident(); - let handler = - extracted_json_response_handler(method_ident, input, status.clone())?; + let handler = extracted_json_response_handler( + method_ident, + input, + status.clone(), + &receiver, + )?; json_success_status = Some(status.clone()); quote! { ::a3s_boot::RouteDefinition::#builder(#path, #handler)? @@ -498,7 +572,7 @@ fn route_registration( "this HTTP method does not support JSON route inference", ) })?; - let handler = raw_or_json_request_handler(method_ident, input)?; + let handler = raw_or_json_request_handler(method_ident, input, &receiver)?; quote! { ::a3s_boot::RouteDefinition::#builder(#path, #status, #handler)? } @@ -507,8 +581,12 @@ fn route_registration( RouteFlavor::JsonBody => { if input.has_extractors() { let builder = route.kind.raw_builder_ident(); - let handler = - extracted_json_response_handler(method_ident, input, status.clone())?; + let handler = extracted_json_response_handler( + method_ident, + input, + status.clone(), + &receiver, + )?; json_success_status = Some(status.clone()); quote! { ::a3s_boot::RouteDefinition::#builder(#path, #handler)? @@ -526,7 +604,7 @@ fn route_registration( "this HTTP method does not support JSON route inference", ) })?; - let handler = json_body_handler(method_ident, input); + let handler = json_body_handler(method_ident, input, &receiver); quote! { ::a3s_boot::RouteDefinition::#builder(#path, #status, #handler)? } @@ -535,6 +613,38 @@ fn route_registration( } }; + let mut route_definition = route_definition; + if controller_route_kind == ControllerRouteKind::Provider { + let method = route.kind.http_method_ident(); + let provider_route = route_definition; + route_definition = quote! { + ::a3s_boot::RouteDefinition::new_provider::( + ::a3s_boot::HttpMethod::#method, + #path, + move |__a3s_boot_provider_controller: ::std::sync::Arc| { + let __a3s_boot_provider_route = #provider_route; + Ok(__a3s_boot_provider_route) + }, + )? + }; + + if let Some(render_spec) = render_spec { + let view = &render_spec.view; + route_definition = quote! { + (#route_definition) + .with_response( + #status, + ::a3s_boot::OpenApiResponse::description("Success") + ) + .with_metadata("render:view", #view)? + }; + } else if matches!(flavor, RouteFlavor::JsonRequest | RouteFlavor::JsonBody) + && !metadata_input.has_extractors() + { + json_success_status = Some(status.clone()); + } + } + let route_definition = validation_route_definition( route_definition, &metadata_input, diff --git a/macros/src/controller/routing.rs b/macros/src/controller/routing.rs index b463a47..8e8395e 100644 --- a/macros/src/controller/routing.rs +++ b/macros/src/controller/routing.rs @@ -2,6 +2,7 @@ use quote::{format_ident, quote}; use syn::parse::{Parse, ParseStream}; use syn::{Attribute, Ident, LitInt, LitStr, Result, Token}; +#[derive(Clone)] pub(super) struct RouteArgs { pub(super) path: LitStr, pub(super) status: Option, @@ -72,6 +73,7 @@ impl Parse for RouteArgs { } } +#[derive(Clone)] pub(super) struct RouteSpec { pub(super) kind: RouteKind, pub(super) args: RouteArgs, @@ -96,6 +98,19 @@ pub(super) enum RouteKind { } impl RouteKind { + pub(super) fn http_method_ident(self) -> Ident { + match self { + Self::All => format_ident!("All"), + Self::Get | Self::Sse | Self::GetJson => format_ident!("Get"), + Self::Post | Self::PostJson => format_ident!("Post"), + Self::Put | Self::PutJson => format_ident!("Put"), + Self::Patch | Self::PatchJson => format_ident!("Patch"), + Self::Delete | Self::DeleteJson => format_ident!("Delete"), + Self::Options => format_ident!("Options"), + Self::Head => format_ident!("Head"), + } + } + pub(super) fn from_attribute(attr: &Attribute) -> Option { let ident = attr.path().segments.last()?.ident.to_string(); match ident.as_str() { diff --git a/macros/src/dependency.rs b/macros/src/dependency.rs index 1a7eac4..0b1cedd 100644 --- a/macros/src/dependency.rs +++ b/macros/src/dependency.rs @@ -11,7 +11,7 @@ use crate::{parse_optional_comma, set_once}; pub(crate) fn expand_injectable( mut item_struct: syn::ItemStruct, ) -> Result { - let constructor = injectable_constructor(&mut item_struct)?; + let (constructor, dependencies) = injectable_constructor(&mut item_struct)?; let ident = &item_struct.ident; let mut from_module_ref_generics = item_struct.generics.clone(); from_module_ref_generics @@ -28,6 +28,12 @@ pub(crate) fn expand_injectable( fn from_module_ref(module_ref: &::a3s_boot::ModuleRef) -> ::a3s_boot::Result { #constructor } + + fn provider_dependencies() -> ::std::option::Option< + ::std::vec::Vec<::a3s_boot::ProviderDependency>, + > { + ::std::option::Option::Some(::std::vec![#(#dependencies,)*]) + } } impl #impl_generics #ident #ty_generics #where_clause { @@ -139,7 +145,16 @@ pub(crate) fn expand_module( } else { quote! { Ok(::std::vec![ - #(module_ref.get::<#controllers>()?.controller()?,)* + #({ + if module_ref.provider_is_contextual::<#controllers>()? + || module_ref.provider_scope::<#controllers>()? + != ::a3s_boot::ProviderScope::Singleton + { + <#controllers>::provider_controller()? + } else { + module_ref.get::<#controllers>()?.controller()? + } + },)* ]) } }; @@ -164,9 +179,16 @@ pub(crate) fn expand_module( quote! { let mut __a3s_boot_patterns = ::std::vec::Vec::new(); #( - __a3s_boot_patterns.extend( - module_ref.get::<#message_controllers>()?.message_patterns()? - ); + __a3s_boot_patterns.extend({ + if module_ref.provider_is_contextual::<#message_controllers>()? + || module_ref.provider_scope::<#message_controllers>()? + != ::a3s_boot::ProviderScope::Singleton + { + <#message_controllers>::provider_message_patterns()? + } else { + module_ref.get::<#message_controllers>()?.message_patterns()? + } + }); )* Ok(__a3s_boot_patterns) } @@ -345,50 +367,57 @@ fn export_registration_token(spec: &ModuleExportSpec) -> proc_macro2::TokenStrea } } -fn injectable_constructor(item_struct: &mut syn::ItemStruct) -> Result { +fn injectable_constructor( + item_struct: &mut syn::ItemStruct, +) -> Result<(proc_macro2::TokenStream, Vec)> { match &mut item_struct.fields { - Fields::Unit => Ok(quote! { Ok(Self) }), + Fields::Unit => Ok((quote! { Ok(Self) }, Vec::new())), Fields::Named(fields) => { let mut values = Vec::new(); + let mut dependencies = Vec::new(); for field in fields.named.iter_mut() { let ident = field.ident.clone().ok_or_else(|| { syn::Error::new_spanned(&*field, "#[injectable] requires named fields") })?; let token = take_field_inject_attr(&mut field.attrs)?; - let value = match injectable_field_dependency(&field.ty) { - Some(InjectableFieldDependency::Required(inner)) => match token { + let dependency = injectable_field_dependency(&field.ty).ok_or_else(|| { + syn::Error::new_spanned( + &field.ty, + "#[injectable] fields must be Arc, Option>, ProviderRef, or Option>", + ) + })?; + dependencies.push(provider_dependency_token(dependency, token.as_ref())); + let value = match dependency { + InjectableFieldDependency::Required(inner) => match token { Some(token) => quote! { module_ref.get_named::<#inner>(#token)? }, None => quote! { module_ref.get::<#inner>()? }, }, - Some(InjectableFieldDependency::Optional(inner)) => match token { + InjectableFieldDependency::Optional(inner) => match token { Some(token) => quote! { module_ref.get_optional_named::<#inner>(#token)? }, None => quote! { module_ref.get_optional::<#inner>()? }, }, - Some(InjectableFieldDependency::ProviderRef(inner)) => match token { + InjectableFieldDependency::ProviderRef(inner) => match token { Some(token) => quote! { module_ref.named_provider_ref::<#inner>(#token) }, None => quote! { module_ref.provider_ref::<#inner>() }, }, - Some(InjectableFieldDependency::OptionalProviderRef(inner)) => match token { + InjectableFieldDependency::OptionalProviderRef(inner) => match token { Some(token) => { quote! { module_ref.optional_named_provider_ref::<#inner>(#token)? } } None => quote! { module_ref.optional_provider_ref::<#inner>()? }, }, - None => { - return Err(syn::Error::new_spanned( - &field.ty, - "#[injectable] fields must be Arc, Option>, ProviderRef, or Option>", - )); - } }; values.push(quote! { #ident: #value }); } - Ok(quote! { - Ok(Self { - #(#values,)* - }) - }) + Ok(( + quote! { + Ok(Self { + #(#values,)* + }) + }, + dependencies, + )) } Fields::Unnamed(fields) => Err(syn::Error::new_spanned( fields, @@ -397,6 +426,7 @@ fn injectable_constructor(item_struct: &mut syn::ItemStruct) -> Result { Required(&'a Type), Optional(&'a Type), @@ -404,6 +434,31 @@ enum InjectableFieldDependency<'a> { OptionalProviderRef(&'a Type), } +fn provider_dependency_token( + kind: InjectableFieldDependency<'_>, + token: Option<&LitStr>, +) -> proc_macro2::TokenStream { + let inner = match kind { + InjectableFieldDependency::Required(inner) + | InjectableFieldDependency::Optional(inner) + | InjectableFieldDependency::ProviderRef(inner) + | InjectableFieldDependency::OptionalProviderRef(inner) => inner, + }; + let dependency = match token { + Some(token) => quote! { ::a3s_boot::ProviderDependency::named(#token) }, + None => quote! { ::a3s_boot::ProviderDependency::typed::<#inner>() }, + }; + + match kind { + InjectableFieldDependency::Required(_) => dependency, + InjectableFieldDependency::Optional(_) => quote! { (#dependency).optional() }, + InjectableFieldDependency::ProviderRef(_) => quote! { (#dependency).lazy() }, + InjectableFieldDependency::OptionalProviderRef(_) => { + quote! { (#dependency).optional().lazy() } + } + } +} + fn injectable_field_dependency(field_type: &Type) -> Option> { if let Some(inner) = single_type_argument(field_type, "Arc") { return Some(InjectableFieldDependency::Required(inner)); diff --git a/macros/src/messaging.rs b/macros/src/messaging.rs index fe37887..66288ac 100644 --- a/macros/src/messaging.rs +++ b/macros/src/messaging.rs @@ -27,6 +27,7 @@ pub(crate) fn expand_message_controller( let self_ty = item_impl.self_ty.clone(); let mut patterns = Vec::new(); + let mut provider_patterns = Vec::new(); let mut errors: Option = None; let (impl_attrs, impl_decorator_errors) = expand_apply_decorators_attrs(&item_impl.attrs); for error in impl_decorator_errors { @@ -135,14 +136,31 @@ pub(crate) fn expand_message_controller( match message_pattern_registration( method, input.clone(), - spec, + spec.clone(), validation_options, &controller_metadata, &metadata_specs, &controller_pipeline, &pipeline_specs, + MessageControllerPatternKind::Instance, ) { - Ok(pattern) => patterns.push(pattern), + Ok(pattern) => { + patterns.push(pattern); + match message_pattern_registration( + method, + input.clone(), + spec, + validation_options, + &controller_metadata, + &metadata_specs, + &controller_pipeline, + &pipeline_specs, + MessageControllerPatternKind::Provider, + ) { + Ok(pattern) => provider_patterns.push(pattern), + Err(error) => push_error(&mut errors, error), + } + } Err(error) => push_error(&mut errors, error), } } @@ -165,10 +183,28 @@ pub(crate) fn expand_message_controller( )* Ok(__a3s_boot_patterns) } + + pub fn provider_message_patterns( + ) -> ::a3s_boot::Result<::std::vec::Vec<::a3s_boot::MessagePatternDefinition>> + where + Self: ::std::marker::Send + ::std::marker::Sync + 'static, + { + let mut __a3s_boot_patterns = ::std::vec::Vec::new(); + #( + __a3s_boot_patterns.push(#provider_patterns); + )* + Ok(__a3s_boot_patterns) + } } }) } +#[derive(Clone, Copy, PartialEq, Eq)] +enum MessageControllerPatternKind { + Instance, + Provider, +} + fn take_message_pattern_attrs( attrs: &[Attribute], ) -> (Vec, Vec, Vec) { @@ -213,6 +249,7 @@ fn message_pattern_registration( metadata_specs: &[MetadataSpec], controller_pipeline: &[proc_macro2::TokenStream], pipeline_specs: &[PipelineSpec], + controller_pattern_kind: MessageControllerPatternKind, ) -> Result { let method_ident = &method.sig.ident; let pattern = spec.args.pattern; @@ -220,15 +257,41 @@ fn message_pattern_registration( let args = message_handler_args(input)?; let definition = match spec.kind { MessagePatternAttrKind::Message => { - let handler = message_request_handler(method_ident, raw, args.clone()); - quote! { - ::a3s_boot::MessagePatternDefinition::request(#pattern, #handler)? + let handler = match controller_pattern_kind { + MessageControllerPatternKind::Instance => { + message_request_handler(method_ident, raw, args.clone()) + } + MessageControllerPatternKind::Provider => { + provider_message_request_factory(method_ident, raw, args.clone()) + } + }; + if controller_pattern_kind == MessageControllerPatternKind::Provider { + quote! { + ::a3s_boot::MessagePatternDefinition::request_scoped(#pattern, #handler)? + } + } else { + quote! { + ::a3s_boot::MessagePatternDefinition::request(#pattern, #handler)? + } } } MessagePatternAttrKind::Event => { - let handler = message_event_handler(method_ident, args.clone()); - quote! { - ::a3s_boot::MessagePatternDefinition::event(#pattern, #handler)? + let handler = match controller_pattern_kind { + MessageControllerPatternKind::Instance => { + message_event_handler(method_ident, args.clone()) + } + MessageControllerPatternKind::Provider => { + provider_message_event_factory(method_ident, args.clone()) + } + }; + if controller_pattern_kind == MessageControllerPatternKind::Provider { + quote! { + ::a3s_boot::MessagePatternDefinition::event_scoped(#pattern, #handler)? + } + } else { + quote! { + ::a3s_boot::MessagePatternDefinition::event(#pattern, #handler)? + } } } }; @@ -285,6 +348,31 @@ fn message_request_handler( } } +fn provider_message_request_factory( + method_ident: &Ident, + raw: bool, + args: MessageHandlerArgs, +) -> proc_macro2::TokenStream { + let controller_name = format_ident!("__a3s_boot_message_{}", method_ident); + let bindings = args.bindings(); + let call_args = args.call_args(); + let call = message_request_call(method_ident, &controller_name, raw, quote!(#(#call_args),*)); + quote! { + { + move |__a3s_boot_module_ref: &::a3s_boot::ModuleRef| { + let #controller_name = __a3s_boot_module_ref.get::()?; + Ok(move |__a3s_boot_message: ::a3s_boot::TransportMessage| { + let #controller_name = ::std::sync::Arc::clone(&#controller_name); + async move { + #(#bindings)* + #call + } + }) + } + } + } +} + fn message_request_call( method_ident: &Ident, controller_name: &Ident, @@ -327,6 +415,30 @@ fn message_event_handler( } } +fn provider_message_event_factory( + method_ident: &Ident, + args: MessageHandlerArgs, +) -> proc_macro2::TokenStream { + let controller_name = format_ident!("__a3s_boot_event_{}", method_ident); + let bindings = args.bindings(); + let call_args = args.call_args(); + quote! { + { + move |__a3s_boot_module_ref: &::a3s_boot::ModuleRef| { + let #controller_name = __a3s_boot_module_ref.get::()?; + Ok(move |__a3s_boot_message: ::a3s_boot::TransportMessage| { + let #controller_name = ::std::sync::Arc::clone(&#controller_name); + async move { + #(#bindings)* + let _ = #controller_name.#method_ident(#(#call_args),*).await?; + Ok(()) + } + }) + } + } + } +} + fn message_validation_definition( definition: proc_macro2::TokenStream, args: &MessageHandlerArgs, @@ -540,11 +652,13 @@ impl MessagePatternAttrKind { } } +#[derive(Clone)] struct MessagePatternSpec { kind: MessagePatternAttrKind, args: MessagePatternArgs, } +#[derive(Clone)] struct MessagePatternArgs { pattern: LitStr, raw: Option, diff --git a/src/app/application.rs b/src/app/application.rs index 7132507..e199f8e 100644 --- a/src/app/application.rs +++ b/src/app/application.rs @@ -1,10 +1,12 @@ use super::builder::BootApplicationBuilder; +use crate::pipeline::PipelineComponent; use crate::versioning::ApiVersionCandidate; use crate::{ - ApiVersioning, BootError, BootRequest, BootResponse, DiscoveryService, HttpAdapter, HttpMethod, - LazyModuleLoader, MessagePatternDefinition, MessageTransport, Module, ModuleRef, - OpenApiDocument, OpenApiInfo, ProviderToken, Reflector, Result, RouteDefinition, - TransportMessage, TransportReply, WebSocketGatewayDefinition, + ApiVersioning, BootError, BootRequest, BootResponse, ContextIdFactory, DiscoveryService, + ExceptionFilter, ExecutionContext, HttpAdapter, HttpMethod, LazyModuleLoader, + MessagePatternDefinition, MessageTransport, Module, ModuleRef, OpenApiDocument, OpenApiInfo, + ProviderToken, Reflector, Result, RouteDefinition, SerializationOptions, TransportMessage, + TransportReply, WebSocketGatewayDefinition, }; use std::collections::BTreeMap; use std::fmt; @@ -65,6 +67,7 @@ pub struct BootApplication { pub(crate) module_ref: ModuleRef, pub(crate) module_instances: Vec, pub(crate) api_versioning: Option, + pub(crate) global_filters: Vec>, } impl BootApplication { @@ -252,11 +255,9 @@ impl BootApplication { self.api_versioning.as_ref(), host, ) else { - return Err(BootError::NotFound(format!( - "{} {}", - request.method.as_str(), - request.path - ))); + let error = + BootError::NotFound(format!("{} {}", request.method.as_str(), request.path)); + return self.handle_global_error(request, error).await; }; let candidate = &candidates[candidate_index]; @@ -296,11 +297,35 @@ impl BootApplication { return route.call(request).await; } - Err(BootError::NotFound(format!( - "{} {}", - request.method.as_str(), - request.path - ))) + let error = BootError::NotFound(format!("{} {}", request.method.as_str(), request.path)); + self.handle_global_error(request, error).await + } + + async fn handle_global_error( + &self, + request: BootRequest, + error: BootError, + ) -> Result { + let context_id = ContextIdFactory::create(); + let request = request.with_module_ref(self.module_ref.context_scope(&context_id)); + let context = ExecutionContext::new( + request.clone(), + request.path.clone(), + None, + None, + SerializationOptions::default(), + BTreeMap::new(), + ); + for filter in self.global_filters.iter().rev() { + let filter = filter.resolve(&context_id)?; + if let Some(response) = filter + .catch(context.clone(), error.clone_for_filter()) + .await? + { + return Ok(response); + } + } + Err(error) } /// Dispatch a request and convert unhandled errors into Boot HTTP responses. diff --git a/src/app/builder.rs b/src/app/builder.rs index f376e05..3cbe2cc 100644 --- a/src/app/builder.rs +++ b/src/app/builder.rs @@ -451,14 +451,21 @@ impl BootApplicationBuilder { ModuleRegistry::new(global_ref, self.provider_overrides, self.module_overrides); let mut modules = Vec::new(); let mut module_instances = Vec::new(); - let mut routes = self - .routes - .into_iter() - .map(|route| route.with_pipeline_prefix(&self.global_pipeline)) - .collect::>(); + let mut routes = self.routes; let mut gateways = self.gateways; let mut message_patterns = self.message_patterns; + for module in &self.modules { + let registered = registry.register_module(Arc::clone(module))?; + module_ref.add_visible_scope(registered.module_ref)?; + } + let provider_enhancers = registry.provider_enhancers(); + let mut effective_global_pipeline = self.global_pipeline.clone(); + effective_global_pipeline.append(&provider_enhancers.http); + routes = routes + .into_iter() + .map(|route| route.with_pipeline_prefix(&effective_global_pipeline)) + .collect(); { let mut sink = ModuleRegistrationSink { modules: &mut modules, @@ -467,15 +474,7 @@ impl BootApplicationBuilder { gateways: &mut gateways, message_patterns: &mut message_patterns, }; - - for module in &self.modules { - let registered = registry.register_module( - Arc::clone(module), - &self.global_pipeline, - &mut sink, - )?; - module_ref.add_visible_scope(registered.module_ref)?; - } + registry.finalize(&effective_global_pipeline, &mut sink)?; } for (name, registered) in registry.registered_modules() { lazy_module_loader.seed_module(name, registered.module_ref, registered.exports)?; @@ -490,6 +489,7 @@ impl BootApplicationBuilder { .into_iter() .map(|gateway| { gateway + .with_provider_enhancer_prefix(&provider_enhancers) .with_guard_prefix(&self.global_websocket_guards) .with_interceptor_prefix(&self.global_websocket_interceptors) .with_execution_pipeline_prefix( @@ -502,12 +502,14 @@ impl BootApplicationBuilder { self.global_pipeline.validation_enabled, self.global_pipeline.validation_options, ) + .with_default_module_ref(module_ref.clone()) }) .collect::>(); let message_patterns = message_patterns .into_iter() .map(|pattern| { pattern + .with_provider_enhancer_prefix(&provider_enhancers) .with_guard_prefix(&self.global_transport_guards) .with_interceptor_prefix(&self.global_transport_interceptors) .with_execution_pipeline_prefix( @@ -520,14 +522,15 @@ impl BootApplicationBuilder { self.global_pipeline.validation_enabled, self.global_pipeline.validation_options, ) + .with_default_module_ref(module_ref.clone()) }) .collect::>(); let documented_routes = routes.clone(); for (path, info) in self.openapi_routes { let document = OpenApiDocument::from_routes(info, &documented_routes); - let route = - openapi_json_route(path, document)?.with_pipeline_prefix(&self.global_pipeline); + let route = openapi_json_route(path, document)? + .with_pipeline_prefix(&effective_global_pipeline); let route = apply_global_prefix_to_route( route, self.global_prefix.as_deref(), @@ -542,13 +545,13 @@ impl BootApplicationBuilder { &documented_routes, self.global_prefix.as_deref(), &self.global_prefix_exclusions, - &self.global_pipeline, + &effective_global_pipeline, )?; } #[cfg(feature = "security")] if let Some(cors_preflight) = &self.cors_preflight { - add_cors_preflight_routes(&mut routes, cors_preflight, &self.global_pipeline)?; + add_cors_preflight_routes(&mut routes, cors_preflight, &effective_global_pipeline)?; } routes = routes @@ -569,6 +572,7 @@ impl BootApplicationBuilder { module_ref, module_instances, api_versioning: self.api_versioning, + global_filters: effective_global_pipeline.filters.clone(), }) } @@ -583,14 +587,21 @@ impl BootApplicationBuilder { ModuleRegistry::new(global_ref, self.provider_overrides, self.module_overrides); let mut modules = Vec::new(); let mut module_instances = Vec::new(); - let mut routes = self - .routes - .into_iter() - .map(|route| route.with_pipeline_prefix(&self.global_pipeline)) - .collect::>(); + let mut routes = self.routes; let mut gateways = self.gateways; let mut message_patterns = self.message_patterns; + for module in &self.modules { + let registered = registry.register_module_async(Arc::clone(module)).await?; + module_ref.add_visible_scope(registered.module_ref)?; + } + let provider_enhancers = registry.provider_enhancers(); + let mut effective_global_pipeline = self.global_pipeline.clone(); + effective_global_pipeline.append(&provider_enhancers.http); + routes = routes + .into_iter() + .map(|route| route.with_pipeline_prefix(&effective_global_pipeline)) + .collect(); { let mut sink = ModuleRegistrationSink { modules: &mut modules, @@ -599,13 +610,9 @@ impl BootApplicationBuilder { gateways: &mut gateways, message_patterns: &mut message_patterns, }; - - for module in &self.modules { - let registered = registry - .register_module_async(Arc::clone(module), &self.global_pipeline, &mut sink) - .await?; - module_ref.add_visible_scope(registered.module_ref)?; - } + registry + .finalize_async(&effective_global_pipeline, &mut sink) + .await?; } for (name, registered) in registry.registered_modules() { lazy_module_loader.seed_module(name, registered.module_ref, registered.exports)?; @@ -620,6 +627,7 @@ impl BootApplicationBuilder { .into_iter() .map(|gateway| { gateway + .with_provider_enhancer_prefix(&provider_enhancers) .with_guard_prefix(&self.global_websocket_guards) .with_interceptor_prefix(&self.global_websocket_interceptors) .with_execution_pipeline_prefix( @@ -632,12 +640,14 @@ impl BootApplicationBuilder { self.global_pipeline.validation_enabled, self.global_pipeline.validation_options, ) + .with_default_module_ref(module_ref.clone()) }) .collect::>(); let message_patterns = message_patterns .into_iter() .map(|pattern| { pattern + .with_provider_enhancer_prefix(&provider_enhancers) .with_guard_prefix(&self.global_transport_guards) .with_interceptor_prefix(&self.global_transport_interceptors) .with_execution_pipeline_prefix( @@ -650,14 +660,15 @@ impl BootApplicationBuilder { self.global_pipeline.validation_enabled, self.global_pipeline.validation_options, ) + .with_default_module_ref(module_ref.clone()) }) .collect::>(); let documented_routes = routes.clone(); for (path, info) in self.openapi_routes { let document = OpenApiDocument::from_routes(info, &documented_routes); - let route = - openapi_json_route(path, document)?.with_pipeline_prefix(&self.global_pipeline); + let route = openapi_json_route(path, document)? + .with_pipeline_prefix(&effective_global_pipeline); let route = apply_global_prefix_to_route( route, self.global_prefix.as_deref(), @@ -672,13 +683,13 @@ impl BootApplicationBuilder { &documented_routes, self.global_prefix.as_deref(), &self.global_prefix_exclusions, - &self.global_pipeline, + &effective_global_pipeline, )?; } #[cfg(feature = "security")] if let Some(cors_preflight) = &self.cors_preflight { - add_cors_preflight_routes(&mut routes, cors_preflight, &self.global_pipeline)?; + add_cors_preflight_routes(&mut routes, cors_preflight, &effective_global_pipeline)?; } routes = routes @@ -699,6 +710,7 @@ impl BootApplicationBuilder { module_ref, module_instances, api_versioning: self.api_versioning, + global_filters: effective_global_pipeline.filters.clone(), }) } } diff --git a/src/app/lazy.rs b/src/app/lazy.rs index 1588e85..ba441ee 100644 --- a/src/app/lazy.rs +++ b/src/app/lazy.rs @@ -99,6 +99,9 @@ impl LazyModuleLoader { /// /// Lazy-loaded modules are provider-only: controllers, routes, gateways, /// middleware, message patterns, and lifecycle hooks are not registered. + /// Newly loaded global modules are rejected because changing global + /// visibility after singleton initialization would make the existing + /// provider graph inconsistent. pub fn load(&self, module: M) -> Result where M: Module, @@ -108,8 +111,19 @@ impl LazyModuleLoader { /// Load a shared module on demand and return its provider container. pub fn load_arc(&self, module: Arc) -> Result { - let mut visiting = Vec::new(); - self.load_arc_inner(module, &mut visiting) + let name = validate_lazy_module_name(module.name())?; + if let Some(cached) = self.registry.cached(name)? { + return Ok(cached); + } + + let mut graph = LazyModuleGraph::default(); + let loaded = self.register_arc_inner(module, &mut graph, false)?; + graph.validate()?; + graph.initialize()?; + self.registry.cache_modules(&graph.pending)?; + self.registry + .cached(loaded.name())? + .ok_or_else(|| missing_cached_lazy_module(loaded.name())) } /// Load a module with async singleton provider factories on demand. @@ -122,123 +136,166 @@ impl LazyModuleLoader { /// Load a shared module with async singleton provider factories on demand. pub async fn load_arc_async(&self, module: Arc) -> Result { - let mut visiting = Vec::new(); - self.load_arc_async_inner(module, &mut visiting).await + let name = validate_lazy_module_name(module.name())?; + if let Some(cached) = self.registry.cached(name)? { + return Ok(cached); + } + + let mut graph = LazyModuleGraph::default(); + let loaded = self + .register_arc_async_inner(module, &mut graph, false) + .await?; + graph.validate()?; + graph.initialize_async().await?; + self.registry.cache_modules(&graph.pending)?; + self.registry + .cached(loaded.name())? + .ok_or_else(|| missing_cached_lazy_module(loaded.name())) } - fn load_arc_inner( + fn register_arc_inner( &self, module: Arc, - visiting: &mut Vec, + graph: &mut LazyModuleGraph, + allow_active: bool, ) -> Result { let name = validate_lazy_module_name(module.name())?; if let Some(cached) = self.registry.cached(name)? { return Ok(cached); } - - enter_lazy_module(visiting, name)?; - let result = self.build_lazy_module(module, name, visiting); - visiting.pop(); - result + if let Some(registered) = graph.registered.get(name) { + return Ok(registered.clone()); + } + if let Some(active) = graph.active.get(name) { + if allow_active { + return Ok(active.clone()); + } + return Err(cyclic_lazy_module_error(&graph.visiting, name)); + } + reject_lazy_global(module.as_ref(), name)?; + + enter_lazy_module(&mut graph.visiting, name)?; + let loaded = LazyLoadedModule::new(name.to_string(), ModuleRef::new(), ModuleRef::new()); + graph.active.insert(name.to_string(), loaded.clone()); + let result = self.register_lazy_module(module, &loaded, graph); + graph.active.remove(name); + graph.visiting.pop(); + result?; + + graph.registered.insert(name.to_string(), loaded.clone()); + graph.pending.push(loaded.clone()); + Ok(loaded) } - fn build_lazy_module( + fn register_lazy_module( &self, module: Arc, - name: &str, - visiting: &mut Vec, - ) -> Result { + loaded: &LazyLoadedModule, + graph: &mut LazyModuleGraph, + ) -> Result<()> { let mut imported_modules = Vec::new(); for imported in module.imports() { - imported_modules.push(self.load_arc_inner(imported, visiting)?); + imported_modules.push(self.register_arc_inner(imported, graph, false)?); } - - let module_ref = self.create_module_ref(&imported_modules)?; - for provider in module.providers()? { - module_ref.register(provider)?; + for imported in module.forward_imports() { + imported_modules.push(self.register_arc_inner(imported, graph, true)?); } - module_ref.initialize_local_singletons()?; - let exports = self.create_exports(&module_ref, module.exports()?)?; - if module.is_global() { - self.export_global(&exports)?; + self.prepare_module_ref(&loaded.module_ref, &imported_modules)?; + for provider in module.providers()? { + reject_lazy_provider_enhancers(&provider, loaded.name())?; + loaded.module_ref.register(provider)?; } - - let loaded = LazyLoadedModule::new(name.to_string(), module_ref, exports); - self.registry.cache_module(loaded) + self.populate_exports(&loaded.exports, &loaded.module_ref, module.exports()?) } - fn load_arc_async_inner<'a>( + fn register_arc_async_inner<'a>( &'a self, module: Arc, - visiting: &'a mut Vec, + graph: &'a mut LazyModuleGraph, + allow_active: bool, ) -> BoxFuture<'a, Result> { Box::pin(async move { let name = validate_lazy_module_name(module.name())?; if let Some(cached) = self.registry.cached(name)? { return Ok(cached); } - - enter_lazy_module(visiting, name)?; - let result = self.build_lazy_module_async(module, name, visiting).await; - visiting.pop(); - result + if let Some(registered) = graph.registered.get(name) { + return Ok(registered.clone()); + } + if let Some(active) = graph.active.get(name) { + if allow_active { + return Ok(active.clone()); + } + return Err(cyclic_lazy_module_error(&graph.visiting, name)); + } + reject_lazy_global(module.as_ref(), name)?; + + enter_lazy_module(&mut graph.visiting, name)?; + let loaded = + LazyLoadedModule::new(name.to_string(), ModuleRef::new(), ModuleRef::new()); + graph.active.insert(name.to_string(), loaded.clone()); + let result = self + .register_lazy_module_async(module, &loaded, graph) + .await; + graph.active.remove(name); + graph.visiting.pop(); + result?; + + graph.registered.insert(name.to_string(), loaded.clone()); + graph.pending.push(loaded.clone()); + Ok(loaded) }) } - fn build_lazy_module_async<'a>( + fn register_lazy_module_async<'a>( &'a self, module: Arc, - name: &'a str, - visiting: &'a mut Vec, - ) -> BoxFuture<'a, Result> { + loaded: &'a LazyLoadedModule, + graph: &'a mut LazyModuleGraph, + ) -> BoxFuture<'a, Result<()>> { Box::pin(async move { let mut imported_modules = Vec::new(); for imported in module.imports() { - imported_modules.push(self.load_arc_async_inner(imported, visiting).await?); + imported_modules.push( + self.register_arc_async_inner(imported, graph, false) + .await?, + ); } - - let module_ref = self.create_module_ref(&imported_modules)?; - for provider in module.providers()? { - module_ref.register_async(provider).await?; + for imported in module.forward_imports() { + imported_modules.push(self.register_arc_async_inner(imported, graph, true).await?); } - module_ref.initialize_local_singletons_async().await?; - let exports = self.create_exports(&module_ref, module.exports()?)?; - if module.is_global() { - self.export_global(&exports)?; + self.prepare_module_ref(&loaded.module_ref, &imported_modules)?; + for provider in module.providers()? { + reject_lazy_provider_enhancers(&provider, loaded.name())?; + loaded.module_ref.register_async(provider).await?; } - - let loaded = LazyLoadedModule::new(name.to_string(), module_ref, exports); - self.registry.cache_module(loaded) + self.populate_exports(&loaded.exports, &loaded.module_ref, module.exports()?) }) } - fn create_module_ref(&self, imported_modules: &[LazyLoadedModule]) -> Result { - let module_ref = ModuleRef::new(); + fn prepare_module_ref( + &self, + module_ref: &ModuleRef, + imported_modules: &[LazyLoadedModule], + ) -> Result<()> { module_ref.add_visible_scope(self.registry.global_ref.clone())?; for imported in imported_modules { module_ref.add_visible_scope(imported.exports.clone())?; } - Ok(module_ref) + Ok(()) } - fn create_exports( + fn populate_exports( &self, + exports: &ModuleRef, module_ref: &ModuleRef, tokens: Vec, - ) -> Result { - let exports = ModuleRef::new(); + ) -> Result<()> { for token in tokens { exports.export_from(module_ref, &token)?; } - Ok(exports) - } - - fn export_global(&self, exports: &ModuleRef) -> Result<()> { - for token in exports.local_tokens()? { - self.registry.global_ref.export_from(exports, &token)?; - } Ok(()) } } @@ -271,13 +328,14 @@ impl LazyModuleRegistry { Ok(()) } - fn cache_module(&self, module: LazyLoadedModule) -> Result { + fn cache_modules(&self, pending: &[LazyLoadedModule]) -> Result<()> { let mut modules = self.write_modules()?; - if let Some(cached) = modules.get(module.name()).cloned() { - return Ok(cached); + for module in pending { + modules + .entry(module.name.clone()) + .or_insert_with(|| module.clone()); } - modules.insert(module.name.clone(), module.clone()); - Ok(module) + Ok(()) } fn read_modules( @@ -297,6 +355,40 @@ impl LazyModuleRegistry { } } +#[derive(Default)] +struct LazyModuleGraph { + registered: BTreeMap, + active: BTreeMap, + pending: Vec, + visiting: Vec, +} + +impl LazyModuleGraph { + fn validate(&self) -> Result<()> { + for module in &self.pending { + module.module_ref.validate_local_resolution_plans()?; + } + Ok(()) + } + + fn initialize(&self) -> Result<()> { + for module in &self.pending { + module.module_ref.initialize_local_singletons()?; + } + Ok(()) + } + + async fn initialize_async(&self) -> Result<()> { + for module in &self.pending { + module.module_ref.seed_local_async_singletons().await?; + } + for module in &self.pending { + module.module_ref.initialize_local_singletons()?; + } + Ok(()) + } +} + fn validate_lazy_module_name(name: &'static str) -> Result<&'static str> { if name.trim().is_empty() { return Err(BootError::EmptyModuleName); @@ -306,14 +398,49 @@ fn validate_lazy_module_name(name: &'static str) -> Result<&'static str> { fn enter_lazy_module(visiting: &mut Vec, name: &str) -> Result<()> { if let Some(index) = visiting.iter().position(|active| active == name) { - let mut chain = visiting[index..].to_vec(); - chain.push(name.to_string()); + return Err(cyclic_lazy_module_error(&visiting[index..], name)); + } + + visiting.push(name.to_string()); + Ok(()) +} + +fn cyclic_lazy_module_error(visiting: &[String], name: &str) -> BootError { + let index = visiting + .iter() + .position(|active| active == name) + .unwrap_or(0); + let mut chain = visiting[index..].to_vec(); + chain.push(name.to_string()); + BootError::Internal(format!( + "cyclic lazy module import detected: {}", + chain.join(" -> ") + )) +} + +fn reject_lazy_global(module: &dyn Module, name: &str) -> Result<()> { + if module.is_global() { return Err(BootError::Internal(format!( - "cyclic lazy module import detected: {}", - chain.join(" -> ") + "lazy-loaded global module `{name}` would change the finalized application provider graph; register global modules eagerly" ))); } + Ok(()) +} - visiting.push(name.to_string()); +fn reject_lazy_provider_enhancers( + provider: &crate::ProviderDefinition, + module_name: &str, +) -> Result<()> { + if !provider.enhancer_markers().is_empty() { + return Err(BootError::Internal(format!( + "lazy-loaded module `{module_name}` declares an application-wide provider enhancer; register modules with APP_* providers eagerly" + ))); + } Ok(()) } + +fn missing_cached_lazy_module(name: &str) -> BootError { + BootError::Internal(format!( + "lazy module `{name}` was initialized but not cached" + )) +} diff --git a/src/app/registration.rs b/src/app/registration.rs index 3831317..19d4576 100644 --- a/src/app/registration.rs +++ b/src/app/registration.rs @@ -1,5 +1,5 @@ use super::application::ModuleInstance; -use crate::pipeline::PipelineComponents; +use crate::pipeline::{PipelineComponents, ProviderEnhancerComponents}; use crate::routing::path::join_paths; use crate::{ BootError, BoxFuture, ControllerDefinition, MessagePatternDefinition, MiddlewareConsumer, @@ -12,10 +12,12 @@ use std::sync::Arc; pub(super) struct ModuleRegistry { registered: BTreeMap, active: BTreeMap, + pending: Vec, visiting: Vec, global_ref: ModuleRef, provider_overrides: BTreeMap, module_overrides: BTreeMap>, + provider_enhancers: ProviderEnhancerComponents, } pub(super) struct ModuleRegistrationSink<'a> { @@ -32,6 +34,16 @@ struct ActiveModuleRegistration { exports: ModuleRef, } +struct PendingModuleRegistration { + module: Arc, + name: String, + route_prefix: String, + module_ref: ModuleRef, + import_names: Vec, + exported_tokens: Vec, + is_global: bool, +} + impl ModuleRegistry { pub fn new( global_ref: ModuleRef, @@ -41,27 +53,22 @@ impl ModuleRegistry { Self { registered: BTreeMap::new(), active: BTreeMap::new(), + pending: Vec::new(), visiting: Vec::new(), global_ref, provider_overrides, module_overrides, + provider_enhancers: ProviderEnhancerComponents::default(), } } - pub fn register_module( - &mut self, - module: Arc, - global_pipeline: &PipelineComponents, - sink: &mut ModuleRegistrationSink<'_>, - ) -> Result { - self.register_module_with_prefix(module, global_pipeline, sink, "") + pub fn register_module(&mut self, module: Arc) -> Result { + self.register_module_with_prefix(module, "") } fn register_module_with_prefix( &mut self, module: Arc, - global_pipeline: &PipelineComponents, - sink: &mut ModuleRegistrationSink<'_>, parent_route_prefix: &str, ) -> Result { let module = self.module_override_or(module); @@ -91,7 +98,7 @@ impl ModuleRegistry { exports: exports.clone(), }, ); - let result = self.register_module_inner(module, name, global_pipeline, sink, state); + let result = self.register_module_inner(module, name, state); self.active.remove(name); self.exit_module(); result @@ -100,8 +107,6 @@ impl ModuleRegistry { fn register_forward_module_with_prefix( &mut self, module: Arc, - global_pipeline: &PipelineComponents, - sink: &mut ModuleRegistrationSink<'_>, parent_route_prefix: &str, ) -> Result { let module = self.module_override_or(module); @@ -117,23 +122,19 @@ impl ModuleRegistry { return Ok(active.clone()); } - self.register_module_with_prefix(module, global_pipeline, sink, parent_route_prefix) + self.register_module_with_prefix(module, parent_route_prefix) } pub fn register_module_async<'a>( &'a mut self, module: Arc, - global_pipeline: &'a PipelineComponents, - sink: &'a mut ModuleRegistrationSink<'_>, ) -> BoxFuture<'a, Result> { - self.register_module_async_with_prefix(module, global_pipeline, sink, "") + self.register_module_async_with_prefix(module, "") } fn register_module_async_with_prefix<'a>( &'a mut self, module: Arc, - global_pipeline: &'a PipelineComponents, - sink: &'a mut ModuleRegistrationSink<'_>, parent_route_prefix: &'a str, ) -> BoxFuture<'a, Result> { Box::pin(async move { @@ -164,9 +165,7 @@ impl ModuleRegistry { exports: exports.clone(), }, ); - let result = self - .register_module_async_inner(module, name, global_pipeline, sink, state) - .await; + let result = self.register_module_async_inner(module, name, state).await; self.active.remove(name); self.exit_module(); result @@ -176,8 +175,6 @@ impl ModuleRegistry { fn register_forward_module_async_with_prefix<'a>( &'a mut self, module: Arc, - global_pipeline: &'a PipelineComponents, - sink: &'a mut ModuleRegistrationSink<'_>, parent_route_prefix: &'a str, ) -> BoxFuture<'a, Result> { Box::pin(async move { @@ -194,13 +191,8 @@ impl ModuleRegistry { return Ok(active.clone()); } - self.register_module_async_with_prefix( - module, - global_pipeline, - sink, - parent_route_prefix, - ) - .await + self.register_module_async_with_prefix(module, parent_route_prefix) + .await }) } @@ -211,12 +203,14 @@ impl ModuleRegistry { .collect() } + pub fn provider_enhancers(&self) -> ProviderEnhancerComponents { + self.provider_enhancers.clone() + } + fn register_module_inner( &mut self, module: Arc, name: &'static str, - global_pipeline: &PipelineComponents, - sink: &mut ModuleRegistrationSink<'_>, state: ActiveModuleRegistration, ) -> Result { let ActiveModuleRegistration { @@ -226,20 +220,11 @@ impl ModuleRegistry { } = state; let mut imported_modules = Vec::new(); for imported in module.imports() { - imported_modules.push(self.register_module_with_prefix( - imported, - global_pipeline, - sink, - &route_prefix, - )?); + imported_modules.push(self.register_module_with_prefix(imported, &route_prefix)?); } for imported in module.forward_imports() { - imported_modules.push(self.register_forward_module_with_prefix( - imported, - global_pipeline, - sink, - &route_prefix, - )?); + imported_modules + .push(self.register_forward_module_with_prefix(imported, &route_prefix)?); } module_ref.add_visible_scope(self.global_ref.clone())?; @@ -249,9 +234,13 @@ impl ModuleRegistry { for provider in module.providers()? { let provider = self.provider_override_or(provider); + let enhancers = provider.enhancer_markers().to_vec(); module_ref.register(provider)?; + for enhancer in enhancers { + self.provider_enhancers + .push(enhancer.bind(module_ref.clone())); + } } - module_ref.initialize_local_singletons()?; let export_tokens = module.exports()?; for token in &export_tokens { @@ -266,70 +255,23 @@ impl ModuleRegistry { } } - module_ref.initialize_local_providers()?; - module.on_module_init(&module_ref)?; - - let mut module_pipeline = PipelineComponents::default(); - for middleware in module.middleware() { - module_pipeline.push_middleware_arc(middleware); - } - let mut middleware_consumer = MiddlewareConsumer::new(); - module.configure(&mut middleware_consumer, &module_ref)?; - - for controller in module.controllers(&module_ref)? { - let context = RouteRegistrationContext { - module_name: name, - module_ref: &module_ref, - global_pipeline, - module_pipeline: &module_pipeline, - middleware_consumer: &middleware_consumer, - route_prefix: &route_prefix, - }; - register_controller(&context, controller, sink.routes)?; - } - - for route in module.routes()? { - let context = RouteRegistrationContext { - module_name: name, - module_ref: &module_ref, - global_pipeline, - module_pipeline: &module_pipeline, - middleware_consumer: &middleware_consumer, - route_prefix: &route_prefix, - }; - sink.routes.push(context.prepare_route(route)?); - } - sink.gateways.extend( - module - .gateways(&module_ref)? - .into_iter() - .map(|gateway| gateway.with_module_name(name)), - ); - sink.message_patterns.extend( - module - .message_patterns(&module_ref)? - .into_iter() - .map(|pattern| pattern.with_module_name(name)), - ); - let import_names = imported_modules .iter() .map(|module| module.name.clone()) .collect::>(); - let route_prefix = (!route_prefix.is_empty()).then_some(route_prefix); let registered = RegisteredModule { name: name.to_string(), module_ref: module_ref.clone(), exports, }; - sink.modules.push(name.to_string()); - sink.module_instances.push(ModuleInstance { + self.pending.push(PendingModuleRegistration { module, module_ref, - imports: import_names, - exports: exported_tokens, - is_global, + name: name.to_string(), route_prefix, + import_names, + exported_tokens, + is_global, }); self.registered.insert(name.to_string(), registered.clone()); Ok(registered) @@ -339,8 +281,6 @@ impl ModuleRegistry { &'a mut self, module: Arc, name: &'static str, - global_pipeline: &'a PipelineComponents, - sink: &'a mut ModuleRegistrationSink<'_>, state: ActiveModuleRegistration, ) -> BoxFuture<'a, Result> { Box::pin(async move { @@ -352,24 +292,14 @@ impl ModuleRegistry { let mut imported_modules = Vec::new(); for imported in module.imports() { imported_modules.push( - self.register_module_async_with_prefix( - imported, - global_pipeline, - sink, - &route_prefix, - ) - .await?, + self.register_module_async_with_prefix(imported, &route_prefix) + .await?, ); } for imported in module.forward_imports() { imported_modules.push( - self.register_forward_module_async_with_prefix( - imported, - global_pipeline, - sink, - &route_prefix, - ) - .await?, + self.register_forward_module_async_with_prefix(imported, &route_prefix) + .await?, ); } @@ -380,9 +310,13 @@ impl ModuleRegistry { for provider in module.providers()? { let provider = self.provider_override_or(provider); + let enhancers = provider.enhancer_markers().to_vec(); module_ref.register_async(provider).await?; + for enhancer in enhancers { + self.provider_enhancers + .push(enhancer.bind(module_ref.clone())); + } } - module_ref.initialize_local_singletons_async().await?; let export_tokens = module.exports()?; for token in &export_tokens { @@ -397,76 +331,72 @@ impl ModuleRegistry { } } - module_ref.initialize_local_providers()?; - module.on_module_init(&module_ref)?; - - let mut module_pipeline = PipelineComponents::default(); - for middleware in module.middleware() { - module_pipeline.push_middleware_arc(middleware); - } - let mut middleware_consumer = MiddlewareConsumer::new(); - module.configure(&mut middleware_consumer, &module_ref)?; - - for controller in module.controllers(&module_ref)? { - let context = RouteRegistrationContext { - module_name: name, - module_ref: &module_ref, - global_pipeline, - module_pipeline: &module_pipeline, - middleware_consumer: &middleware_consumer, - route_prefix: &route_prefix, - }; - register_controller(&context, controller, sink.routes)?; - } - - for route in module.routes()? { - let context = RouteRegistrationContext { - module_name: name, - module_ref: &module_ref, - global_pipeline, - module_pipeline: &module_pipeline, - middleware_consumer: &middleware_consumer, - route_prefix: &route_prefix, - }; - sink.routes.push(context.prepare_route(route)?); - } - sink.gateways.extend( - module - .gateways(&module_ref)? - .into_iter() - .map(|gateway| gateway.with_module_name(name)), - ); - sink.message_patterns.extend( - module - .message_patterns(&module_ref)? - .into_iter() - .map(|pattern| pattern.with_module_name(name)), - ); - let import_names = imported_modules .iter() .map(|module| module.name.clone()) .collect::>(); - let route_prefix = (!route_prefix.is_empty()).then_some(route_prefix); let registered = RegisteredModule { name: name.to_string(), module_ref: module_ref.clone(), exports, }; - sink.modules.push(name.to_string()); - sink.module_instances.push(ModuleInstance { + self.pending.push(PendingModuleRegistration { module, module_ref, - imports: import_names, - exports: exported_tokens, - is_global, + name: name.to_string(), route_prefix, + import_names, + exported_tokens, + is_global, }); self.registered.insert(name.to_string(), registered.clone()); Ok(registered) }) } + pub fn finalize( + &mut self, + global_pipeline: &PipelineComponents, + sink: &mut ModuleRegistrationSink<'_>, + ) -> Result<()> { + for pending in &self.pending { + pending.module_ref.validate_local_resolution_plans()?; + } + for pending in &self.pending { + pending.module_ref.initialize_local_singletons()?; + } + + let pending = std::mem::take(&mut self.pending); + for pending in pending { + finalize_module(pending, global_pipeline, sink)?; + } + Ok(()) + } + + pub fn finalize_async<'a>( + &'a mut self, + global_pipeline: &'a PipelineComponents, + sink: &'a mut ModuleRegistrationSink<'_>, + ) -> BoxFuture<'a, Result<()>> { + Box::pin(async move { + for pending in &self.pending { + pending.module_ref.validate_local_resolution_plans()?; + } + for pending in &self.pending { + pending.module_ref.seed_local_async_singletons().await?; + } + for pending in &self.pending { + pending.module_ref.initialize_local_singletons()?; + } + + let pending = std::mem::take(&mut self.pending); + for pending in pending { + finalize_module(pending, global_pipeline, sink)?; + } + Ok(()) + }) + } + fn enter_module(&mut self, name: &str) -> Result<()> { if let Some(index) = self.visiting.iter().position(|active| active == name) { let mut chain = self.visiting[index..].to_vec(); @@ -486,10 +416,10 @@ impl ModuleRegistry { } fn provider_override_or(&self, provider: ProviderDefinition) -> ProviderDefinition { - self.provider_overrides - .get(provider.token()) - .cloned() - .unwrap_or(provider) + match self.provider_overrides.get(provider.token()) { + Some(replacement) => replacement.clone().with_enhancers_from(&provider), + None => provider, + } } fn module_override_or(&self, module: Arc) -> Arc { @@ -500,6 +430,75 @@ impl ModuleRegistry { } } +fn finalize_module( + pending: PendingModuleRegistration, + global_pipeline: &PipelineComponents, + sink: &mut ModuleRegistrationSink<'_>, +) -> Result<()> { + let PendingModuleRegistration { + module, + name, + route_prefix, + module_ref, + import_names, + exported_tokens, + is_global, + } = pending; + + module_ref.initialize_local_providers()?; + module.on_module_init(&module_ref)?; + + let mut module_pipeline = PipelineComponents::default(); + for middleware in module.middleware() { + module_pipeline.push_middleware_arc(middleware); + } + let mut middleware_consumer = MiddlewareConsumer::new(); + module.configure(&mut middleware_consumer, &module_ref)?; + + let context = RouteRegistrationContext { + module_name: &name, + module_ref: &module_ref, + global_pipeline, + module_pipeline: &module_pipeline, + middleware_consumer: &middleware_consumer, + route_prefix: &route_prefix, + }; + for controller in module.controllers(&module_ref)? { + register_controller(&context, controller, sink.routes)?; + } + for route in module.routes()? { + sink.routes.push(context.prepare_route(route)?); + } + sink.gateways + .extend(module.gateways(&module_ref)?.into_iter().map(|gateway| { + gateway + .with_module_name(&name) + .with_module_ref(module_ref.clone()) + })); + sink.message_patterns + .extend( + module + .message_patterns(&module_ref)? + .into_iter() + .map(|pattern| { + pattern + .with_module_name(&name) + .with_module_ref(module_ref.clone()) + }), + ); + + sink.modules.push(name); + sink.module_instances.push(ModuleInstance { + module, + module_ref, + imports: import_names, + exports: exported_tokens, + is_global, + route_prefix: (!route_prefix.is_empty()).then_some(route_prefix), + }); + Ok(()) +} + #[derive(Clone)] pub(super) struct RegisteredModule { pub name: String, diff --git a/src/http/request.rs b/src/http/request.rs index 5dccca3..cebcccb 100644 --- a/src/http/request.rs +++ b/src/http/request.rs @@ -9,7 +9,7 @@ use crate::percent::validate_percent_encoding; use crate::routing::host::normalize_host_header; #[cfg(feature = "auth")] use crate::AuthPrincipal; -use crate::{validate_value, BootError, ModuleRef, ProviderToken, Result, Validate}; +use crate::{validate_value, BootError, ContextId, ModuleRef, ProviderToken, Result, Validate}; use serde::{de::DeserializeOwned, Serialize}; use std::collections::BTreeMap; use std::fmt; @@ -103,6 +103,11 @@ impl BootRequest { self.module_ref.as_ref() } + /// Dependency-injection context attached while this request is dispatched. + pub fn context_id(&self) -> Option<&ContextId> { + self.module_ref.as_ref().and_then(ModuleRef::context_id) + } + pub fn get(&self) -> Result> where T: Send + Sync + 'static, diff --git a/src/lib.rs b/src/lib.rs index 467e553..0626f7c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -169,16 +169,17 @@ pub use openapi_security::{ OpenApiSecurityScheme, }; pub use pipeline::{ - catch_errors, CatchFilter, ExceptionFilter, ExecutionContext, ExecutionInterceptor, - ExecutionProtocol, ExecutionTransportKind, Guard, Interceptor, Middleware, MiddlewareConsumer, - MiddlewareConsumerBuilder, MiddlewareOutcome, MiddlewareRoute, Pipe, TransportExceptionFilter, - TransportExceptionResponse, TransportExecutionContext, WebSocketExceptionFilter, - WebSocketExceptionResponse, WebSocketExecutionContext, + catch_errors, CallHandler, CatchFilter, ExceptionFilter, ExecutionContext, + ExecutionInterceptor, ExecutionProtocol, ExecutionTransportKind, Guard, Interceptor, + Middleware, MiddlewareConsumer, MiddlewareConsumerBuilder, MiddlewareOutcome, MiddlewareRoute, + Pipe, TransportExceptionFilter, TransportExceptionResponse, TransportExecutionContext, + WebSocketExceptionFilter, WebSocketExceptionResponse, WebSocketExecutionContext, }; pub use provider::{ - FromModuleRef, ModuleRef, ProviderBeforeApplicationShutdown, ProviderDefinition, - ProviderOnApplicationBootstrap, ProviderOnApplicationShutdown, ProviderOnModuleDestroy, - ProviderOnModuleInit, ProviderRef, ProviderScope, ProviderToken, + ContextId, ContextIdFactory, FromModuleRef, ModuleRef, ProviderBeforeApplicationShutdown, + ProviderDefinition, ProviderDependency, ProviderOnApplicationBootstrap, + ProviderOnApplicationShutdown, ProviderOnModuleDestroy, ProviderOnModuleInit, ProviderRef, + ProviderScope, ProviderToken, }; #[cfg(feature = "queue")] pub use queue::{ diff --git a/src/pipeline/components.rs b/src/pipeline/components.rs index f9ee862..cdab5dd 100644 --- a/src/pipeline/components.rs +++ b/src/pipeline/components.rs @@ -3,9 +3,9 @@ use super::{ Interceptor, Middleware, Pipe, }; use crate::{ - BootErrorKind, TransportExceptionFilter, TransportGuard, TransportInterceptor, TransportPipe, - ValidationOptions, WebSocketExceptionFilter, WebSocketGuard, WebSocketInterceptor, - WebSocketPipe, + BootErrorKind, ContextId, ModuleRef, ProviderToken, Result, TransportExceptionFilter, + TransportGuard, TransportInterceptor, TransportPipe, ValidationOptions, + WebSocketExceptionFilter, WebSocketGuard, WebSocketInterceptor, WebSocketPipe, }; use std::any::TypeId; use std::collections::HashMap; @@ -13,21 +13,64 @@ use std::sync::Arc; pub(crate) struct PipelineComponent { type_id: TypeId, - inner: Arc, + source: PipelineComponentSource, +} + +type PipelineProviderResolver = dyn Fn(&ModuleRef, &ContextId) -> Result> + Send + Sync; + +enum PipelineComponentSource { + Static(Arc), + Provider { + owner: ModuleRef, + resolver: Arc>, + }, } impl Clone for PipelineComponent { fn clone(&self) -> Self { Self { type_id: self.type_id, - inner: Arc::clone(&self.inner), + source: match &self.source { + PipelineComponentSource::Static(inner) => { + PipelineComponentSource::Static(Arc::clone(inner)) + } + PipelineComponentSource::Provider { owner, resolver } => { + PipelineComponentSource::Provider { + owner: owner.clone(), + resolver: Arc::clone(resolver), + } + } + }, } } } impl PipelineComponent { - pub(crate) fn inner(&self) -> &T { - self.inner.as_ref() + fn from_static(type_id: TypeId, inner: Arc) -> Self { + Self { + type_id, + source: PipelineComponentSource::Static(inner), + } + } + + pub(crate) fn from_provider(type_id: TypeId, owner: ModuleRef, resolver: R) -> Self + where + R: Fn(&ModuleRef, &ContextId) -> Result> + Send + Sync + 'static, + { + Self { + type_id, + source: PipelineComponentSource::Provider { + owner, + resolver: Arc::new(resolver), + }, + } + } + + pub(crate) fn resolve(&self, context_id: &ContextId) -> Result> { + match &self.source { + PipelineComponentSource::Static(inner) => Ok(Arc::clone(inner)), + PipelineComponentSource::Provider { owner, resolver } => resolver(owner, context_id), + } } fn type_id(&self) -> TypeId { @@ -35,15 +78,168 @@ impl PipelineComponent { } } +type ProviderEnhancerBinder = dyn Fn(ModuleRef) -> BoundProviderEnhancer + Send + Sync; + +#[derive(Clone)] +pub(crate) struct ProviderEnhancerMarker { + binder: Arc, +} + +#[derive(Clone)] +pub(crate) enum BoundProviderEnhancer { + Pipe(PipelineComponent), + Guard(PipelineComponent), + Interceptor(PipelineComponent), + Filter(PipelineComponent), + WebSocketPipe(PipelineComponent), + WebSocketGuard(PipelineComponent), + WebSocketInterceptor(PipelineComponent), + WebSocketFilter(PipelineComponent), + TransportPipe(PipelineComponent), + TransportGuard(PipelineComponent), + TransportInterceptor(PipelineComponent), + TransportFilter(PipelineComponent), +} + +macro_rules! provider_enhancer_marker { + ($method:ident, $variant:ident, $provider_trait:path, $trait_object:ty) => { + pub(crate) fn $method(token: ProviderToken) -> Self + where + T: $provider_trait, + { + Self { + binder: Arc::new(move |owner| { + let token = token.clone(); + BoundProviderEnhancer::$variant(PipelineComponent::from_provider( + TypeId::of::(), + owner, + move |module_ref, context_id| { + module_ref + .resolve_token_with_context::(&token, context_id) + .map(|provider| provider as Arc<$trait_object>) + }, + )) + }), + } + } + }; +} + +impl ProviderEnhancerMarker { + provider_enhancer_marker!(pipe, Pipe, Pipe, dyn Pipe); + provider_enhancer_marker!(guard, Guard, Guard, dyn Guard); + provider_enhancer_marker!(interceptor, Interceptor, Interceptor, dyn Interceptor); + provider_enhancer_marker!(filter, Filter, ExceptionFilter, dyn ExceptionFilter); + provider_enhancer_marker!( + websocket_pipe, + WebSocketPipe, + WebSocketPipe, + dyn WebSocketPipe + ); + provider_enhancer_marker!( + websocket_guard, + WebSocketGuard, + WebSocketGuard, + dyn WebSocketGuard + ); + provider_enhancer_marker!( + websocket_interceptor, + WebSocketInterceptor, + WebSocketInterceptor, + dyn WebSocketInterceptor + ); + provider_enhancer_marker!( + websocket_filter, + WebSocketFilter, + WebSocketExceptionFilter, + dyn WebSocketExceptionFilter + ); + provider_enhancer_marker!( + transport_pipe, + TransportPipe, + TransportPipe, + dyn TransportPipe + ); + provider_enhancer_marker!( + transport_guard, + TransportGuard, + TransportGuard, + dyn TransportGuard + ); + provider_enhancer_marker!( + transport_interceptor, + TransportInterceptor, + TransportInterceptor, + dyn TransportInterceptor + ); + provider_enhancer_marker!( + transport_filter, + TransportFilter, + TransportExceptionFilter, + dyn TransportExceptionFilter + ); + + pub(crate) fn bind(&self, owner: ModuleRef) -> BoundProviderEnhancer { + (self.binder)(owner) + } +} + +#[derive(Clone, Default)] +pub(crate) struct ProviderEnhancerComponents { + pub http: PipelineComponents, + pub websocket_pipes: Vec>, + pub websocket_guards: Vec>, + pub websocket_interceptors: Vec>, + pub websocket_filters: Vec>, + pub transport_pipes: Vec>, + pub transport_guards: Vec>, + pub transport_interceptors: Vec>, + pub transport_filters: Vec>, +} + +impl ProviderEnhancerComponents { + pub(crate) fn push(&mut self, enhancer: BoundProviderEnhancer) { + match enhancer { + BoundProviderEnhancer::Pipe(component) => self.http.pipes.push(component), + BoundProviderEnhancer::Guard(component) => self.http.guards.push(component), + BoundProviderEnhancer::Interceptor(component) => { + self.http.interceptors.push(component); + } + BoundProviderEnhancer::Filter(component) => self.http.filters.push(component), + BoundProviderEnhancer::WebSocketPipe(component) => { + self.websocket_pipes.push(component); + } + BoundProviderEnhancer::WebSocketGuard(component) => { + self.websocket_guards.push(component); + } + BoundProviderEnhancer::WebSocketInterceptor(component) => { + self.websocket_interceptors.push(component); + } + BoundProviderEnhancer::WebSocketFilter(component) => { + self.websocket_filters.push(component); + } + BoundProviderEnhancer::TransportPipe(component) => { + self.transport_pipes.push(component); + } + BoundProviderEnhancer::TransportGuard(component) => { + self.transport_guards.push(component); + } + BoundProviderEnhancer::TransportInterceptor(component) => { + self.transport_interceptors.push(component); + } + BoundProviderEnhancer::TransportFilter(component) => { + self.transport_filters.push(component); + } + } + } +} + impl PipelineComponent { pub(crate) fn new

(pipe: P) -> Self where P: Pipe, { - Self { - type_id: TypeId::of::

(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::

(), Arc::new(pipe)) } fn replacement(pipe: P) -> Self @@ -51,10 +247,7 @@ impl PipelineComponent { T: Pipe, P: Pipe, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::(), Arc::new(pipe)) } } @@ -63,17 +256,11 @@ impl PipelineComponent { where G: Guard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } pub(crate) fn from_arc(guard: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: guard, - } + Self::from_static(TypeId::of::>(), guard) } fn replacement(guard: G) -> Self @@ -81,10 +268,7 @@ impl PipelineComponent { T: Guard, G: Guard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } } @@ -93,10 +277,7 @@ impl PipelineComponent { where I: Interceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } fn replacement(interceptor: I) -> Self @@ -104,10 +285,7 @@ impl PipelineComponent { T: Interceptor, I: Interceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } } @@ -116,10 +294,7 @@ impl PipelineComponent { where F: ExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } fn replacement(filter: F) -> Self @@ -127,10 +302,7 @@ impl PipelineComponent { T: ExceptionFilter, F: ExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } } @@ -139,17 +311,11 @@ impl PipelineComponent { where P: WebSocketPipe, { - Self { - type_id: TypeId::of::

(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::

(), Arc::new(pipe)) } pub(crate) fn from_arc(pipe: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: pipe, - } + Self::from_static(TypeId::of::>(), pipe) } fn replacement(pipe: P) -> Self @@ -157,10 +323,7 @@ impl PipelineComponent { T: WebSocketPipe, P: WebSocketPipe, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::(), Arc::new(pipe)) } } @@ -169,17 +332,11 @@ impl PipelineComponent { where G: WebSocketGuard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } pub(crate) fn from_arc(guard: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: guard, - } + Self::from_static(TypeId::of::>(), guard) } fn replacement(guard: G) -> Self @@ -187,10 +344,7 @@ impl PipelineComponent { T: WebSocketGuard, G: WebSocketGuard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } } @@ -199,17 +353,11 @@ impl PipelineComponent { where I: WebSocketInterceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } pub(crate) fn from_arc(interceptor: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: interceptor, - } + Self::from_static(TypeId::of::>(), interceptor) } fn replacement(interceptor: I) -> Self @@ -217,10 +365,7 @@ impl PipelineComponent { T: WebSocketInterceptor, I: WebSocketInterceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } } @@ -229,17 +374,11 @@ impl PipelineComponent { where F: WebSocketExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } pub(crate) fn from_arc(filter: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: filter, - } + Self::from_static(TypeId::of::>(), filter) } fn replacement(filter: F) -> Self @@ -247,10 +386,7 @@ impl PipelineComponent { T: WebSocketExceptionFilter, F: WebSocketExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } } @@ -259,17 +395,11 @@ impl PipelineComponent { where P: TransportPipe, { - Self { - type_id: TypeId::of::

(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::

(), Arc::new(pipe)) } pub(crate) fn from_arc(pipe: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: pipe, - } + Self::from_static(TypeId::of::>(), pipe) } fn replacement(pipe: P) -> Self @@ -277,10 +407,7 @@ impl PipelineComponent { T: TransportPipe, P: TransportPipe, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(pipe), - } + Self::from_static(TypeId::of::(), Arc::new(pipe)) } } @@ -289,17 +416,11 @@ impl PipelineComponent { where G: TransportGuard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } pub(crate) fn from_arc(guard: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: guard, - } + Self::from_static(TypeId::of::>(), guard) } fn replacement(guard: G) -> Self @@ -307,10 +428,7 @@ impl PipelineComponent { T: TransportGuard, G: TransportGuard, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(guard), - } + Self::from_static(TypeId::of::(), Arc::new(guard)) } } @@ -319,17 +437,11 @@ impl PipelineComponent { where I: TransportInterceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } pub(crate) fn from_arc(interceptor: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: interceptor, - } + Self::from_static(TypeId::of::>(), interceptor) } fn replacement(interceptor: I) -> Self @@ -337,10 +449,7 @@ impl PipelineComponent { T: TransportInterceptor, I: TransportInterceptor, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(interceptor), - } + Self::from_static(TypeId::of::(), Arc::new(interceptor)) } } @@ -349,17 +458,11 @@ impl PipelineComponent { where F: TransportExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } pub(crate) fn from_arc(filter: Arc) -> Self { - Self { - type_id: TypeId::of::>(), - inner: filter, - } + Self::from_static(TypeId::of::>(), filter) } fn replacement(filter: F) -> Self @@ -367,10 +470,7 @@ impl PipelineComponent { T: TransportExceptionFilter, F: TransportExceptionFilter, { - Self { - type_id: TypeId::of::(), - inner: Arc::new(filter), - } + Self::from_static(TypeId::of::(), Arc::new(filter)) } } @@ -386,6 +486,18 @@ pub(crate) struct PipelineComponents { } impl PipelineComponents { + pub(crate) fn append(&mut self, components: &Self) { + self.middleware + .extend(components.middleware.iter().cloned()); + self.pipes.extend(components.pipes.iter().cloned()); + self.guards.extend(components.guards.iter().cloned()); + self.interceptors + .extend(components.interceptors.iter().cloned()); + self.filters.extend(components.filters.iter().cloned()); + self.validation_enabled = self.validation_enabled || components.validation_enabled; + self.validation_options = self.validation_options.merge(components.validation_options); + } + pub fn push_middleware(&mut self, middleware: M) where M: Middleware, diff --git a/src/pipeline/context.rs b/src/pipeline/context.rs index eeb903c..7b48390 100644 --- a/src/pipeline/context.rs +++ b/src/pipeline/context.rs @@ -120,6 +120,7 @@ impl ExecutionContext { } pub(crate) fn transport( + request: BootRequest, pattern: String, kind: ExecutionTransportKind, module_name: Option, @@ -134,7 +135,7 @@ impl ExecutionContext { controller_prefix: None, serialization: SerializationOptions::default(), metadata, - request: BootRequest::new(HttpMethod::Post, "/__transport"), + request, websocket: None, transport: Some(TransportExecutionContext { pattern, kind }), } diff --git a/src/pipeline/interceptor.rs b/src/pipeline/interceptor.rs index fa9bdc6..0d2c9cb 100644 --- a/src/pipeline/interceptor.rs +++ b/src/pipeline/interceptor.rs @@ -1,9 +1,110 @@ use super::ExecutionContext; -use crate::{BootResponse, BoxFuture, Result}; +use crate::{BootError, BootResponse, BoxFuture, Result}; +use std::future::Future; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; +/// Reusable access to the next handler in an interceptor chain. +/// +/// This is the Rust equivalent of Nest's `CallHandler`. Calling [`handle`](Self::handle) +/// runs the remaining interceptors, pipes, validation, and handler. The handle is +/// reusable so an interceptor can deliberately retry the downstream pipeline +/// sequentially. Concurrent calls are rejected because they would share one +/// request-scoped provider context. +pub struct CallHandler<'a, T = BootResponse> { + call: Arc BoxFuture<'a, Result> + Send + Sync + 'a>, + running: Arc, +} + +impl<'a, T> Clone for CallHandler<'a, T> { + fn clone(&self) -> Self { + Self { + call: Arc::clone(&self.call), + running: Arc::clone(&self.running), + } + } +} + +impl<'a, T> CallHandler<'a, T> +where + T: Send + 'a, +{ + /// Build a call handler from a reusable async function. + /// + /// Frameworks normally provide the handler to an interceptor. This + /// constructor is public so interceptors and combinators can be tested in + /// isolation. + pub fn from_fn(call: F) -> Self + where + F: Fn() -> Fut + Send + Sync + 'a, + Fut: Future> + Send + 'a, + { + Self { + call: Arc::new(move || Box::pin(call())), + running: Arc::new(AtomicBool::new(false)), + } + } + + /// Run the remaining interceptor chain and underlying handler once. + /// + /// A completed or cancelled call releases the handler for a later retry. + /// Starting overlapping calls returns [`BootError::Internal`]. + pub fn handle(&self) -> BoxFuture<'a, Result> { + let call = Arc::clone(&self.call); + let running = Arc::clone(&self.running); + Box::pin(async move { + if running + .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) + .is_err() + { + return Err(BootError::Internal( + "call handler is already running".to_string(), + )); + } + + let _reset = CallHandlerReset { + running: Arc::clone(&running), + }; + call().await + }) + } +} + +struct CallHandlerReset { + running: Arc, +} + +impl Drop for CallHandlerReset { + fn drop(&mut self) { + self.running.store(false, Ordering::Release); + } +} + /// Runs around the handler for cross-cutting behavior. pub trait Interceptor: Send + Sync + 'static { + /// Run around the remaining HTTP pipeline. + /// + /// Override this method to catch or replace downstream errors, retry the + /// handler, apply a timeout, or return a response without calling `next`. + /// The default implementation preserves the legacy `before`, + /// `short_circuit`, and `after` hook behavior. + fn intercept<'a>( + &'a self, + context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + self.before(context.clone()).await?; + + if let Some(response) = self.short_circuit(context.clone()).await? { + return Ok(response); + } + + let response = next.handle().await?; + self.after(context, response).await + }) + } + fn before(&self, _context: ExecutionContext) -> BoxFuture<'static, Result<()>> { Box::pin(async { Ok(()) }) } diff --git a/src/pipeline/mod.rs b/src/pipeline/mod.rs index 937d34e..d180d9e 100644 --- a/src/pipeline/mod.rs +++ b/src/pipeline/mod.rs @@ -6,7 +6,10 @@ mod interceptor; mod middleware; mod pipe; -pub(crate) use components::{PipelineComponent, PipelineComponents, PipelineOverrides}; +pub(crate) use components::{ + PipelineComponent, PipelineComponents, PipelineOverrides, ProviderEnhancerComponents, + ProviderEnhancerMarker, +}; pub use context::{ ExecutionContext, ExecutionProtocol, ExecutionTransportKind, TransportExecutionContext, WebSocketExecutionContext, @@ -17,7 +20,7 @@ pub use filter::{ }; pub use guard::Guard; pub(crate) use interceptor::ExecutionInterceptorAdapter; -pub use interceptor::{ExecutionInterceptor, Interceptor}; +pub use interceptor::{CallHandler, ExecutionInterceptor, Interceptor}; pub use middleware::{ Middleware, MiddlewareConsumer, MiddlewareConsumerBuilder, MiddlewareOutcome, MiddlewareRoute, }; diff --git a/src/provider/cache.rs b/src/provider/cache.rs index 8a10623..a844622 100644 --- a/src/provider/cache.rs +++ b/src/provider/cache.rs @@ -1,11 +1,13 @@ use super::AnyProvider; -use std::collections::BTreeMap; +use crate::{BootError, Result}; +use std::collections::{BTreeMap, HashMap, HashSet}; use std::sync::atomic::{AtomicU64, Ordering}; -use std::sync::{Arc, RwLock}; - -pub(crate) type ProviderCache = Arc>>>; +use std::sync::{Arc, Condvar, Mutex, OnceLock}; +use std::thread::ThreadId; static NEXT_PROVIDER_CACHE_KEY: AtomicU64 = AtomicU64::new(1); +static NEXT_PROVIDER_BUILD_ID: AtomicU64 = AtomicU64::new(1); +static PROVIDER_WAIT_GRAPH: OnceLock>> = OnceLock::new(); #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] pub(crate) struct ProviderCacheKey(u64); @@ -16,6 +18,375 @@ impl ProviderCacheKey { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +pub(crate) enum ProviderInstanceKey { + Provider(ProviderCacheKey), + Transient { + provider: ProviderCacheKey, + inquirer: ProviderCacheKey, + }, +} + +impl From for ProviderInstanceKey { + fn from(value: ProviderCacheKey) -> Self { + Self::Provider(value) + } +} + +#[derive(Clone, Default)] +pub(crate) struct ProviderCache { + slots: Arc>>>, +} + +#[derive(Default)] +struct ProviderCacheSlot { + state: Mutex, + ready: Condvar, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct ProviderBuild { + id: u64, + builder: ThreadId, +} + +impl ProviderBuild { + fn new(builder: ThreadId) -> Self { + Self { + id: NEXT_PROVIDER_BUILD_ID.fetch_add(1, Ordering::Relaxed), + builder, + } + } +} + +#[derive(Default)] +enum ProviderCacheSlotState { + #[default] + Vacant, + Building(ProviderBuild), + Ready(Arc), +} + +#[derive(Debug, Clone, Copy)] +struct ProviderWaitEdge { + builder: ThreadId, + build_id: u64, +} + +struct ProviderWaitRegistration { + waiter: ThreadId, + build_id: u64, +} + +impl ProviderWaitRegistration { + fn new(build: ProviderBuild) -> Result { + let waiter = std::thread::current().id(); + let mut graph = provider_wait_graph() + .lock() + .map_err(|_| BootError::Internal("provider wait graph lock is poisoned".to_string()))?; + let mut cursor = build.builder; + let mut visited = HashSet::new(); + while let Some(edge) = graph.get(&cursor).copied() { + let next = edge.builder; + if next == waiter || !visited.insert(cursor) { + return Err(BootError::Internal( + "cyclic concurrent provider dependency detected".to_string(), + )); + } + cursor = next; + } + if cursor == waiter { + return Err(BootError::Internal( + "cyclic concurrent provider dependency detected".to_string(), + )); + } + if graph.contains_key(&waiter) { + return Err(BootError::Internal( + "provider resolver thread is already waiting".to_string(), + )); + } + graph.insert( + waiter, + ProviderWaitEdge { + builder: build.builder, + build_id: build.id, + }, + ); + Ok(Self { + waiter, + build_id: build.id, + }) + } +} + +impl Drop for ProviderWaitRegistration { + fn drop(&mut self) { + let mut graph = provider_wait_graph() + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if graph + .get(&self.waiter) + .is_some_and(|edge| edge.build_id == self.build_id) + { + graph.remove(&self.waiter); + } + } +} + +struct ProviderBuildReset { + slot: Arc, + build: ProviderBuild, + armed: bool, +} + +impl ProviderBuildReset { + fn new(slot: Arc, build: ProviderBuild) -> Self { + Self { + slot, + build, + armed: true, + } + } + + fn disarm(&mut self) { + self.armed = false; + } +} + +impl Drop for ProviderBuildReset { + fn drop(&mut self) { + if !self.armed { + return; + } + + let mut state = self + .slot + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if finish_provider_build(&mut state, self.build, ProviderCacheSlotState::Vacant) { + self.slot.ready.notify_all(); + } + } +} + +impl ProviderCache { + pub(crate) fn new() -> Self { + Self::default() + } + + pub(crate) fn contains(&self, key: impl Into) -> Result { + Ok(self.get(key)?.is_some()) + } + + pub(crate) fn get( + &self, + key: impl Into, + ) -> Result>> { + let slot = self.slot(key.into())?; + let state = slot + .state + .lock() + .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string()))?; + Ok(match &*state { + ProviderCacheSlotState::Ready(value) => Some(Arc::clone(value)), + ProviderCacheSlotState::Vacant | ProviderCacheSlotState::Building(_) => None, + }) + } + + pub(crate) fn insert( + &self, + key: impl Into, + value: Arc, + ) -> Result<()> { + let slot = self.slot(key.into())?; + let mut state = slot + .state + .lock() + .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string()))?; + let active_build = match &*state { + ProviderCacheSlotState::Building(build) => Some(build.id), + ProviderCacheSlotState::Vacant | ProviderCacheSlotState::Ready(_) => None, + }; + *state = ProviderCacheSlotState::Ready(value); + if let Some(build_id) = active_build { + clear_provider_waits(build_id); + } + slot.ready.notify_all(); + Ok(()) + } + + pub(crate) fn get_or_try_insert_with( + &self, + key: impl Into, + build: F, + ) -> Result> + where + F: FnOnce() -> Result>, + { + let slot = self.slot(key.into())?; + let mut build = Some(build); + + loop { + let mut state = slot + .state + .lock() + .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string()))?; + match &*state { + ProviderCacheSlotState::Ready(value) => return Ok(Arc::clone(value)), + ProviderCacheSlotState::Building(build) => { + let waiting = ProviderWaitRegistration::new(*build)?; + state = slot.ready.wait(state).map_err(|_| { + BootError::Internal("provider cache lock is poisoned".to_string()) + })?; + drop(state); + drop(waiting); + } + ProviderCacheSlotState::Vacant => { + let active_build = ProviderBuild::new(std::thread::current().id()); + *state = ProviderCacheSlotState::Building(active_build); + drop(state); + + let mut reset = ProviderBuildReset::new(Arc::clone(&slot), active_build); + let result = build.take().ok_or_else(|| { + BootError::Internal( + "provider cache builder was already consumed".to_string(), + ) + })?(); + + let mut state = slot.state.lock().map_err(|_| { + BootError::Internal("provider cache lock is poisoned".to_string()) + })?; + match result { + Ok(value) => { + if !finish_provider_build( + &mut state, + active_build, + ProviderCacheSlotState::Ready(Arc::clone(&value)), + ) { + return Err(BootError::Internal( + "provider cache build state changed unexpectedly".to_string(), + )); + } + reset.disarm(); + slot.ready.notify_all(); + return Ok(value); + } + Err(error) => { + if !finish_provider_build( + &mut state, + active_build, + ProviderCacheSlotState::Vacant, + ) { + return Err(BootError::Internal( + "provider cache build state changed unexpectedly".to_string(), + )); + } + reset.disarm(); + slot.ready.notify_all(); + return Err(error); + } + } + } + } + } + } + + fn slot(&self, key: ProviderInstanceKey) -> Result> { + let mut slots = self + .slots + .lock() + .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string()))?; + Ok(Arc::clone(slots.entry(key).or_default())) + } +} + +fn provider_wait_graph() -> &'static Mutex> { + PROVIDER_WAIT_GRAPH.get_or_init(|| Mutex::new(HashMap::new())) +} + +fn clear_provider_waits(build_id: u64) { + provider_wait_graph() + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .retain(|_, edge| edge.build_id != build_id); +} + +fn finish_provider_build( + state: &mut ProviderCacheSlotState, + build: ProviderBuild, + finished: ProviderCacheSlotState, +) -> bool { + if !matches!( + state, + ProviderCacheSlotState::Building(active) if active.id == build.id + ) { + return false; + } + + *state = finished; + clear_provider_waits(build.id); + true +} + pub(crate) fn new_provider_cache() -> ProviderCache { - Arc::new(RwLock::new(BTreeMap::new())) + ProviderCache::new() +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::mpsc; + use std::time::Duration; + + #[test] + fn completed_build_stops_participating_in_cycles_before_waiter_drop() { + let (id_sender, id_receiver) = mpsc::channel(); + let (build_sender, build_receiver) = mpsc::channel(); + let (result_sender, result_receiver) = mpsc::channel(); + let builder_thread = std::thread::spawn(move || { + id_sender.send(std::thread::current().id()).unwrap(); + let reverse_build = build_receiver.recv().unwrap(); + let result = ProviderWaitRegistration::new(reverse_build).map(drop); + result_sender.send(result).unwrap(); + }); + + let builder = id_receiver.recv_timeout(Duration::from_secs(2)).unwrap(); + let completed_build = ProviderBuild::new(builder); + let slot = ProviderCacheSlot::default(); + *slot.state.lock().unwrap() = ProviderCacheSlotState::Building(completed_build); + + let lingering_registration = ProviderWaitRegistration::new(completed_build).unwrap(); + { + let graph = provider_wait_graph().lock().unwrap(); + assert_eq!( + graph + .get(&std::thread::current().id()) + .map(|edge| edge.build_id), + Some(completed_build.id) + ); + } + + let value: Arc = Arc::new(()); + assert!(finish_provider_build( + &mut slot.state.lock().unwrap(), + completed_build, + ProviderCacheSlotState::Ready(value), + )); + assert!(!provider_wait_graph() + .lock() + .unwrap() + .contains_key(&std::thread::current().id())); + + build_sender + .send(ProviderBuild::new(std::thread::current().id())) + .unwrap(); + result_receiver + .recv_timeout(Duration::from_secs(2)) + .unwrap() + .unwrap(); + + drop(lingering_registration); + builder_thread.join().unwrap(); + } } diff --git a/src/provider/context_id.rs b/src/provider/context_id.rs new file mode 100644 index 0000000..a109585 --- /dev/null +++ b/src/provider/context_id.rs @@ -0,0 +1,116 @@ +use super::cache::ProviderCache; +use crate::{BootError, BootRequest, Result}; +use std::fmt; +use std::hash::{Hash, Hasher}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Weak}; + +static NEXT_CONTEXT_ID: AtomicU64 = AtomicU64::new(1); + +/// Identity and provider cache for one dependency-injection resolution context. +#[derive(Clone)] +pub struct ContextId { + id: u64, + state: ContextIdStateRef, +} + +#[derive(Clone)] +enum ContextIdStateRef { + Strong(Arc), + Weak(Weak), +} + +struct ContextIdState { + cache: ProviderCache, +} + +impl ContextId { + fn new(id: u64) -> Self { + Self { + id, + state: ContextIdStateRef::Strong(Arc::new(ContextIdState { + cache: ProviderCache::new(), + })), + } + } + + /// Numeric identity useful for diagnostics and correlation. + pub fn id(&self) -> u64 { + self.id + } + + pub(crate) fn cache(&self) -> Result { + let state = match &self.state { + ContextIdStateRef::Strong(state) => Arc::clone(state), + ContextIdStateRef::Weak(state) => state.upgrade().ok_or_else(|| { + BootError::Internal(format!( + "dependency-injection context {} is no longer active", + self.id + )) + })?, + }; + Ok(state.cache.clone()) + } + + pub(crate) fn downgrade(&self) -> Self { + let state = match &self.state { + ContextIdStateRef::Strong(state) => ContextIdStateRef::Weak(Arc::downgrade(state)), + ContextIdStateRef::Weak(state) => ContextIdStateRef::Weak(Weak::clone(state)), + }; + Self { id: self.id, state } + } +} + +impl fmt::Debug for ContextId { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ContextId") + .field("id", &self.id()) + .finish_non_exhaustive() + } +} + +impl PartialEq for ContextId { + fn eq(&self, other: &Self) -> bool { + self.id() == other.id() + } +} + +impl Eq for ContextId {} + +impl PartialOrd for ContextId { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for ContextId { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { + self.id().cmp(&other.id()) + } +} + +impl Hash for ContextId { + fn hash(&self, state: &mut H) + where + H: Hasher, + { + self.id().hash(state); + } +} + +/// Creates and discovers Nest-style dependency-injection context identities. +#[derive(Debug, Clone, Copy, Default)] +pub struct ContextIdFactory; + +impl ContextIdFactory { + /// Create an isolated dependency-injection context. + pub fn create() -> ContextId { + ContextId::new(NEXT_CONTEXT_ID.fetch_add(1, Ordering::Relaxed)) + } + + /// Return the context already attached to a request, or create a fresh one. + pub fn get_by_request(request: &BootRequest) -> ContextId { + request.context_id().cloned().unwrap_or_else(Self::create) + } +} diff --git a/src/provider/definition.rs b/src/provider/definition.rs index d39f1c6..5cdbf94 100644 --- a/src/provider/definition.rs +++ b/src/provider/definition.rs @@ -1,5 +1,10 @@ use super::{AnyProvider, ModuleRef, ProviderToken}; -use crate::{BootError, BoxFuture, Result}; +use crate::pipeline::ProviderEnhancerMarker; +use crate::{ + BootError, BoxFuture, ExceptionFilter, Guard, Interceptor, Pipe, Result, + TransportExceptionFilter, TransportGuard, TransportInterceptor, TransportPipe, + WebSocketExceptionFilter, WebSocketGuard, WebSocketInterceptor, WebSocketPipe, +}; use std::fmt; use std::future::Future; use std::sync::Arc; @@ -68,6 +73,14 @@ pub trait ProviderOnApplicationShutdown: Send + Sync + 'static { /// Builds an injectable value from the module provider graph. pub trait FromModuleRef: Sized + Send + Sync + 'static { fn from_module_ref(module_ref: &ModuleRef) -> Result; + + /// Dependencies captured while constructing this provider. + /// + /// `#[injectable]` implements this automatically. Manual implementations + /// can return `None` when the dependency graph is opaque. + fn provider_dependencies() -> Option> { + None + } } /// Lifetime strategy for provider resolution. @@ -78,6 +91,57 @@ pub enum ProviderScope { Transient, } +/// One provider dependency used to calculate contextual scope propagation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ProviderDependency { + token: ProviderToken, + optional: bool, + lazy: bool, +} + +impl ProviderDependency { + /// Declare a required typed dependency. + pub fn typed() -> Self + where + T: Send + Sync + 'static, + { + Self::named(ProviderToken::of::().as_str()) + } + + /// Declare a required named dependency. + pub fn named(token: impl Into) -> Self { + Self { + token: ProviderToken::named(token), + optional: false, + lazy: false, + } + } + + /// Mark this dependency as optional when its token is not visible. + pub fn optional(mut self) -> Self { + self.optional = true; + self + } + + /// Mark this as a lazy handle that does not participate in scope bubbling. + pub fn lazy(mut self) -> Self { + self.lazy = true; + self + } + + pub fn token(&self) -> &ProviderToken { + &self.token + } + + pub fn is_optional(&self) -> bool { + self.optional + } + + pub fn is_lazy(&self) -> bool { + self.lazy + } +} + /// A provider registration, similar to a Nest provider entry. #[derive(Clone)] pub struct ProviderDefinition { @@ -86,6 +150,8 @@ pub struct ProviderDefinition { scope: ProviderScope, lifecycle: ProviderLifecycleHooks, alias_target: Option, + dependencies: Option>, + enhancers: Arc<[ProviderEnhancerMarker]>, } #[derive(Clone)] @@ -140,6 +206,14 @@ impl fmt::Debug for ProviderDefinition { .field("scope", &self.scope) .field("async", &self.factory.is_async()) .field("alias_target", &self.alias_target) + .field("enhancers", &self.enhancers.len()) + .field( + "dependencies", + &self + .dependencies + .as_ref() + .map(|dependencies| dependencies.len()), + ) .finish_non_exhaustive() } } @@ -187,6 +261,8 @@ impl ProviderDefinition { scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), alias_target: None, + dependencies: Some(Arc::from([])), + enhancers: Arc::from([]), } } @@ -219,6 +295,8 @@ impl ProviderDefinition { scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), alias_target: None, + dependencies: None, + enhancers: Arc::from([]), } } @@ -235,6 +313,8 @@ impl ProviderDefinition { scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), alias_target: None, + dependencies: None, + enhancers: Arc::from([]), } } @@ -271,6 +351,8 @@ impl ProviderDefinition { scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), alias_target: None, + dependencies: None, + enhancers: Arc::from([]), } } @@ -289,6 +371,8 @@ impl ProviderDefinition { scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), alias_target: None, + dependencies: None, + enhancers: Arc::from([]), } } @@ -296,14 +380,22 @@ impl ProviderDefinition { where T: FromModuleRef, { - Self::factory::(T::from_module_ref) + let definition = Self::factory::(T::from_module_ref); + match T::provider_dependencies() { + Some(dependencies) => definition.with_dependencies(dependencies), + None => definition, + } } pub fn named_injectable(token: impl Into) -> Self where T: FromModuleRef, { - Self::named_factory::(token, T::from_module_ref) + let definition = Self::named_factory::(token, T::from_module_ref); + match T::provider_dependencies() { + Some(dependencies) => definition.with_dependencies(dependencies), + None => definition, + } } pub fn alias(target: ProviderToken) -> Self @@ -323,7 +415,13 @@ impl ProviderDefinition { })), scope: ProviderScope::Singleton, lifecycle: ProviderLifecycleHooks::default(), - alias_target: Some(target), + alias_target: Some(target.clone()), + dependencies: Some(Arc::from([ProviderDependency { + token: target, + optional: false, + lazy: false, + }])), + enhancers: Arc::from([]), } } @@ -419,11 +517,256 @@ impl ProviderDefinition { Self::named_injectable::(token).with_scope(ProviderScope::Request) } + /// Register an injectable provider as an application-wide HTTP guard. + pub fn app_guard() -> Self + where + T: FromModuleRef + Guard, + { + Self::injectable::().with_app_guard::() + } + + /// Register an injectable provider as an application-wide HTTP pipe. + pub fn app_pipe() -> Self + where + T: FromModuleRef + Pipe, + { + Self::injectable::().with_app_pipe::() + } + + /// Register an injectable provider as an application-wide HTTP interceptor. + pub fn app_interceptor() -> Self + where + T: FromModuleRef + Interceptor, + { + Self::injectable::().with_app_interceptor::() + } + + /// Register an injectable provider as an application-wide HTTP exception filter. + pub fn app_filter() -> Self + where + T: FromModuleRef + ExceptionFilter, + { + Self::injectable::().with_app_filter::() + } + + /// Register an injectable provider as an application-wide WebSocket guard. + pub fn app_websocket_guard() -> Self + where + T: FromModuleRef + WebSocketGuard, + { + Self::injectable::().with_app_websocket_guard::() + } + + /// Register an injectable provider as an application-wide WebSocket pipe. + pub fn app_websocket_pipe() -> Self + where + T: FromModuleRef + WebSocketPipe, + { + Self::injectable::().with_app_websocket_pipe::() + } + + /// Register an injectable provider as an application-wide WebSocket interceptor. + pub fn app_websocket_interceptor() -> Self + where + T: FromModuleRef + WebSocketInterceptor, + { + Self::injectable::().with_app_websocket_interceptor::() + } + + /// Register an injectable provider as an application-wide WebSocket exception filter. + pub fn app_websocket_filter() -> Self + where + T: FromModuleRef + WebSocketExceptionFilter, + { + Self::injectable::().with_app_websocket_filter::() + } + + /// Register an injectable provider as an application-wide transport guard. + pub fn app_transport_guard() -> Self + where + T: FromModuleRef + TransportGuard, + { + Self::injectable::().with_app_transport_guard::() + } + + /// Register an injectable provider as an application-wide transport pipe. + pub fn app_transport_pipe() -> Self + where + T: FromModuleRef + TransportPipe, + { + Self::injectable::().with_app_transport_pipe::() + } + + /// Register an injectable provider as an application-wide transport interceptor. + pub fn app_transport_interceptor() -> Self + where + T: FromModuleRef + TransportInterceptor, + { + Self::injectable::().with_app_transport_interceptor::() + } + + /// Register an injectable provider as an application-wide transport exception filter. + pub fn app_transport_filter() -> Self + where + T: FromModuleRef + TransportExceptionFilter, + { + Self::injectable::().with_app_transport_filter::() + } + + /// Mark this provider as an application-wide HTTP guard. + pub fn with_app_guard(self) -> Self + where + T: Guard, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::guard::(token)) + } + + /// Mark this provider as an application-wide HTTP pipe. + pub fn with_app_pipe(self) -> Self + where + T: Pipe, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::pipe::(token)) + } + + /// Mark this provider as an application-wide HTTP interceptor. + pub fn with_app_interceptor(self) -> Self + where + T: Interceptor, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::interceptor::(token)) + } + + /// Mark this provider as an application-wide HTTP exception filter. + pub fn with_app_filter(self) -> Self + where + T: ExceptionFilter, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::filter::(token)) + } + + /// Mark this provider as an application-wide WebSocket guard. + pub fn with_app_websocket_guard(self) -> Self + where + T: WebSocketGuard, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::websocket_guard::(token)) + } + + /// Mark this provider as an application-wide WebSocket pipe. + pub fn with_app_websocket_pipe(self) -> Self + where + T: WebSocketPipe, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::websocket_pipe::(token)) + } + + /// Mark this provider as an application-wide WebSocket interceptor. + pub fn with_app_websocket_interceptor(self) -> Self + where + T: WebSocketInterceptor, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::websocket_interceptor::(token)) + } + + /// Mark this provider as an application-wide WebSocket exception filter. + pub fn with_app_websocket_filter(self) -> Self + where + T: WebSocketExceptionFilter, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::websocket_filter::(token)) + } + + /// Mark this provider as an application-wide transport guard. + pub fn with_app_transport_guard(self) -> Self + where + T: TransportGuard, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::transport_guard::(token)) + } + + /// Mark this provider as an application-wide transport pipe. + pub fn with_app_transport_pipe(self) -> Self + where + T: TransportPipe, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::transport_pipe::(token)) + } + + /// Mark this provider as an application-wide transport interceptor. + pub fn with_app_transport_interceptor(self) -> Self + where + T: TransportInterceptor, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::transport_interceptor::(token)) + } + + /// Mark this provider as an application-wide transport exception filter. + pub fn with_app_transport_filter(self) -> Self + where + T: TransportExceptionFilter, + { + let token = self.token.clone(); + self.with_enhancer(ProviderEnhancerMarker::transport_filter::(token)) + } + pub fn with_scope(mut self, scope: ProviderScope) -> Self { self.scope = scope; self } + fn with_enhancer(mut self, enhancer: ProviderEnhancerMarker) -> Self { + let mut enhancers = self.enhancers.to_vec(); + enhancers.push(enhancer); + self.enhancers = enhancers.into(); + self + } + + /// Declare the dependencies captured by this provider factory. + pub fn with_dependencies(mut self, dependencies: I) -> Self + where + I: IntoIterator, + { + self.dependencies = Some(dependencies.into_iter().collect::>().into()); + self + } + + /// Add one required typed dependency to this provider factory. + pub fn depends_on(self) -> Self + where + T: Send + Sync + 'static, + { + self.with_dependency(ProviderDependency::typed::()) + } + + /// Add one required named dependency to this provider factory. + pub fn depends_on_named(self, token: impl Into) -> Self { + self.with_dependency(ProviderDependency::named(token)) + } + + /// Add one dependency while retaining previously declared metadata. + pub fn with_dependency(mut self, dependency: ProviderDependency) -> Self { + let mut dependencies = self + .dependencies + .as_deref() + .map(<[ProviderDependency]>::to_vec) + .unwrap_or_default(); + dependencies.push(dependency); + self.dependencies = Some(dependencies.into()); + self + } + pub fn with_on_module_init(mut self) -> Self where T: ProviderOnModuleInit, @@ -500,6 +843,11 @@ impl ProviderDefinition { self.scope } + /// Declared dependency metadata, or `None` for an opaque factory. + pub fn dependencies(&self) -> Option<&[ProviderDependency]> { + self.dependencies.as_deref() + } + pub(super) fn is_alias(&self) -> bool { self.alias_target.is_some() } @@ -516,6 +864,15 @@ impl ProviderDefinition { &self.lifecycle } + pub(crate) fn enhancer_markers(&self) -> &[ProviderEnhancerMarker] { + &self.enhancers + } + + pub(crate) fn with_enhancers_from(mut self, source: &Self) -> Self { + self.enhancers = Arc::clone(&source.enhancers); + self + } + pub(super) fn build(&self, module_ref: &ModuleRef) -> Result> { match &self.factory { ProviderFactoryKind::Sync(factory) => factory(module_ref), diff --git a/src/provider/entry.rs b/src/provider/entry.rs index a4c4c9c..280b424 100644 --- a/src/provider/entry.rs +++ b/src/provider/entry.rs @@ -1,13 +1,43 @@ -use super::cache::{ProviderCache, ProviderCacheKey}; +use super::cache::{ProviderCache, ProviderCacheKey, ProviderInstanceKey}; use super::module_ref::ModuleRef; use super::resolution::{ - enter_resolution_stack, exit_resolution_stack, new_resolution_stack, ProviderResolutionStack, + ensure_not_resolving, enter_resolution_stack, new_resolution_stack, resolution_chain_with, + ProviderResolutionStack, }; -use super::{AnyProvider, ProviderDefinition, ProviderScope, ProviderToken}; +use super::{AnyProvider, ContextId, ProviderDefinition, ProviderScope, ProviderToken}; use crate::{BootError, Result}; -use std::collections::BTreeMap; +use std::collections::{BTreeMap, BTreeSet}; use std::sync::Arc; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct ProviderResolutionPlan { + scope: ProviderScope, + contextual: bool, +} + +impl ProviderResolutionPlan { + fn declared(scope: ProviderScope) -> Self { + Self { + scope, + contextual: scope == ProviderScope::Request, + } + } + + pub(crate) fn scope(self) -> ProviderScope { + self.scope + } + + pub(crate) fn is_contextual(self) -> bool { + self.contextual + } +} + +#[derive(Default)] +struct ProviderPlanState { + visiting: BTreeSet, + memo: BTreeMap, +} + #[derive(Clone)] pub(crate) struct ProviderEntry { cache_key: ProviderCacheKey, @@ -37,43 +67,167 @@ impl ProviderEntry { self.definition.scope() } - pub(crate) fn is_local_singleton(&self) -> bool { - self.scope() == ProviderScope::Singleton && !self.definition.is_alias() + pub(crate) fn cache_key(&self) -> ProviderCacheKey { + self.cache_key + } + + pub(crate) fn token(&self) -> &ProviderToken { + self.definition.token() + } + + pub(crate) fn is_alias(&self) -> bool { + self.definition.is_alias() } pub(crate) fn is_async_factory(&self) -> bool { self.definition.is_async_factory() } + pub(crate) fn has_lifecycle_hooks(&self) -> bool { + self.definition.lifecycle().has_hooks() + } + + pub(crate) fn resolution_plan(&self, module_ref: &ModuleRef) -> Result { + self.resolution_plan_inner(module_ref, &mut ProviderPlanState::default()) + } + + pub(crate) fn owner_module_ref(&self, fallback: &ModuleRef) -> ModuleRef { + self.owner.clone().unwrap_or_else(|| fallback.clone()) + } + + /// Resolve the eager edges used to order async singleton factories. + /// + /// Alias targets are authoritative even if callers replaced the alias's + /// dependency metadata. Required lazy edges are still checked for + /// visibility by resolution planning, but are intentionally excluded here. + pub(crate) fn eager_dependency_entries( + &self, + module_ref: &ModuleRef, + ) -> Result> { + let base_ref = self.owner.as_ref().unwrap_or(module_ref); + if let Some(target) = self.definition.alias_target() { + return base_ref + .get_entry(target)? + .map(|entry| vec![entry]) + .ok_or_else(|| BootError::MissingProvider(target.to_string())); + } + + let mut entries = Vec::new(); + if let Some(dependencies) = self.definition.dependencies() { + for dependency in dependencies { + let Some(entry) = base_ref.get_entry(dependency.token())? else { + if dependency.is_optional() { + continue; + } + return Err(BootError::MissingProvider(dependency.token().to_string())); + }; + if !dependency.is_lazy() { + entries.push(entry); + } + } + } + Ok(entries) + } + + fn resolution_plan_inner( + &self, + module_ref: &ModuleRef, + state: &mut ProviderPlanState, + ) -> Result { + if let Some(plan) = state.memo.get(&self.cache_key) { + return Ok(*plan); + } + + let declared = ProviderResolutionPlan::declared(self.scope()); + if !state.visiting.insert(self.cache_key) { + return Ok(declared); + } + + let base_ref = self.owner.as_ref().unwrap_or(module_ref); + let result = if let Some(target) = self.definition.alias_target() { + let target = base_ref + .get_entry(target)? + .ok_or_else(|| BootError::MissingProvider(target.to_string()))?; + target.resolution_plan_inner(base_ref, state) + } else { + let mut plan = declared; + if let Some(dependencies) = self.definition.dependencies() { + for dependency in dependencies { + let Some(entry) = base_ref.get_entry(dependency.token())? else { + if dependency.is_optional() { + continue; + } + return Err(BootError::MissingProvider(dependency.token().to_string())); + }; + if !dependency.is_lazy() + && entry.resolution_plan_inner(base_ref, state)?.contextual + { + plan.contextual = true; + } + } + } + Ok(plan) + }; + + state.visiting.remove(&self.cache_key); + if let Ok(plan) = result { + state.memo.insert(self.cache_key, plan); + } + result + } + pub(crate) fn resolve( &self, module_ref: &ModuleRef, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, + inquirer: Option, resolution_stack: &ProviderResolutionStack, alias_path: &mut Vec, ) -> Result> { let base_ref = self.owner.as_ref().unwrap_or(module_ref); - if let Some(target) = self.definition.alias_target() { + if self.definition.alias_target().is_some() { return self.resolve_alias( base_ref, - target, - request_cache, + context_id, + transient_cache, + inquirer, resolution_stack, alias_path, ); } - match self.scope() { - ProviderScope::Singleton => self.resolve_singleton(base_ref, resolution_stack), - ProviderScope::Transient => { - let factory_ref = - self.factory_ref(base_ref, request_cache.clone(), resolution_stack); - self.build_with_resolution_stack(&factory_ref, resolution_stack) + let plan = self.resolution_plan(base_ref)?; + match plan.scope() { + ProviderScope::Singleton if !plan.is_contextual() => { + self.resolve_singleton(base_ref, transient_cache, resolution_stack) } - ProviderScope::Request => { - let factory_ref = - self.factory_ref(base_ref, request_cache.clone(), resolution_stack); - self.resolve_request(&factory_ref, request_cache, resolution_stack) + ProviderScope::Singleton | ProviderScope::Request => { + let factory_ref = self.factory_ref( + base_ref, + context_id.clone(), + transient_cache, + resolution_stack, + ); + self.resolve_contextual(&factory_ref, context_id, resolution_stack) + } + ProviderScope::Transient => { + if plan.is_contextual() && context_id.is_none() { + return Err(self.missing_request_context_error(resolution_stack)); + } + let factory_ref = self.factory_ref( + base_ref, + context_id.clone(), + transient_cache.clone(), + resolution_stack, + ); + self.resolve_transient( + &factory_ref, + context_id, + transient_cache, + inquirer, + resolution_stack, + ) } } } @@ -81,11 +235,15 @@ impl ProviderEntry { fn resolve_alias( &self, module_ref: &ModuleRef, - target: &ProviderToken, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, + inquirer: Option, resolution_stack: &ProviderResolutionStack, alias_path: &mut Vec, ) -> Result> { + let target = self.definition.alias_target().ok_or_else(|| { + BootError::Internal("provider alias is missing its target".to_string()) + })?; if alias_path.contains(self.definition.token()) { alias_path.push(self.definition.token().clone()); let chain = alias_path @@ -99,9 +257,11 @@ impl ProviderEntry { } alias_path.push(self.definition.token().clone()); - let value = module_ref.get_any_with_request_cache_inner( + let value = module_ref.get_any_with_context_inner( target, - request_cache, + context_id, + transient_cache, + inquirer, resolution_stack, alias_path, )?; @@ -113,65 +273,94 @@ impl ProviderEntry { pub(crate) fn resolve_singleton( &self, module_ref: &ModuleRef, + transient_cache: ProviderCache, resolution_stack: &ProviderResolutionStack, ) -> Result> { - if let Some(value) = self - .read_cache(&self.singleton)? - .get(&self.cache_key) - .cloned() - { - return Ok(value); - } + ensure_not_resolving(resolution_stack, self.cache_key, self.definition.token())?; + self.singleton.get_or_try_insert_with(self.cache_key, || { + let factory_ref = self.factory_ref(module_ref, None, transient_cache, resolution_stack); + self.build_with_resolution_stack(&factory_ref, resolution_stack) + }) + } - let factory_ref = module_ref.with_resolution_stack(Arc::clone(resolution_stack)); - let value = self.build_with_resolution_stack(&factory_ref, resolution_stack)?; - self.write_cache(&self.singleton)? - .insert(self.cache_key, Arc::clone(&value)); - Ok(value) + fn resolve_contextual( + &self, + module_ref: &ModuleRef, + context_id: Option, + resolution_stack: &ProviderResolutionStack, + ) -> Result> { + let Some(context_id) = context_id else { + return Err(self.missing_request_context_error(resolution_stack)); + }; + ensure_not_resolving(resolution_stack, self.cache_key, self.definition.token())?; + context_id + .cache()? + .get_or_try_insert_with(self.cache_key, || { + self.build_with_resolution_stack(module_ref, resolution_stack) + }) } - fn resolve_request( + fn resolve_transient( &self, module_ref: &ModuleRef, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, + inquirer: Option, resolution_stack: &ProviderResolutionStack, ) -> Result> { - let Some(request_cache) = request_cache else { + let inquirer = match (context_id.as_ref(), inquirer) { + (Some(_), None) => Some(self.cache_key), + (_, inquirer) => inquirer, + }; + let Some(inquirer) = inquirer else { return self.build_with_resolution_stack(module_ref, resolution_stack); }; - if let Some(value) = self - .read_cache(&request_cache)? - .get(&self.cache_key) - .cloned() - { - return Ok(value); - } + ensure_not_resolving(resolution_stack, self.cache_key, self.definition.token())?; + let cache = match context_id.as_ref() { + Some(context_id) => context_id.cache()?, + None => transient_cache, + }; + cache.get_or_try_insert_with( + ProviderInstanceKey::Transient { + provider: self.cache_key, + inquirer, + }, + || self.build_with_resolution_stack(module_ref, resolution_stack), + ) + } - let value = self.build_with_resolution_stack(module_ref, resolution_stack)?; - self.write_cache(&request_cache)? - .insert(self.cache_key, Arc::clone(&value)); - Ok(value) + fn missing_request_context_error( + &self, + resolution_stack: &ProviderResolutionStack, + ) -> BootError { + let chain = resolution_chain_with(resolution_stack, self.definition.token()); + BootError::Internal(format!( + "contextual provider chain `{chain}` requires an active request scope; use ModuleRef::resolve(...) for an isolated context or declare factory dependencies so request scope can propagate" + )) } pub(crate) async fn seed_singleton_async(&self, module_ref: ModuleRef) -> Result<()> { - let resolution_stack = new_resolution_stack(); - let module_ref = module_ref.with_resolution_stack(Arc::clone(&resolution_stack)); - enter_resolution_stack(&resolution_stack, self.definition.token())?; - let result = self.definition.build_async(module_ref).await; - let exit_result = exit_resolution_stack(&resolution_stack); - let value = match (result, exit_result) { - (Ok(value), Ok(())) => value, - (Err(error), _) => return Err(error), - (Ok(_), Err(error)) => return Err(error), - }; + if self.singleton.contains(self.cache_key)? { + return Ok(()); + } + let (resolution_stack, _guard) = enter_resolution_stack( + &new_resolution_stack(), + self.cache_key, + self.definition.token(), + )?; + let module_ref = self.factory_ref( + &module_ref, + None, + module_ref.transient_cache(), + &resolution_stack, + ); + let value = self.definition.build_async(module_ref).await?; self.seed_singleton(value) } fn seed_singleton(&self, value: Arc) -> Result<()> { - self.write_cache(&self.singleton)? - .insert(self.cache_key, value); - Ok(()) + self.singleton.insert(self.cache_key, value) } pub(crate) fn on_module_init(&self, module_ref: &ModuleRef) -> Result<()> { @@ -180,7 +369,8 @@ impl ProviderEntry { }; let resolution_stack = new_resolution_stack(); - let value = self.resolve_singleton(module_ref, &resolution_stack)?; + let value = + self.resolve_singleton(module_ref, module_ref.transient_cache(), &resolution_stack)?; hook(value, module_ref) } @@ -190,7 +380,8 @@ impl ProviderEntry { }; let resolution_stack = new_resolution_stack(); - let value = self.resolve_singleton(&module_ref, &resolution_stack)?; + let value = + self.resolve_singleton(&module_ref, module_ref.transient_cache(), &resolution_stack)?; hook(value, module_ref).await } @@ -204,7 +395,8 @@ impl ProviderEntry { }; let resolution_stack = new_resolution_stack(); - let value = self.resolve_singleton(&module_ref, &resolution_stack)?; + let value = + self.resolve_singleton(&module_ref, module_ref.transient_cache(), &resolution_stack)?; hook(value, module_ref, signal).await } @@ -218,7 +410,8 @@ impl ProviderEntry { }; let resolution_stack = new_resolution_stack(); - let value = self.resolve_singleton(&module_ref, &resolution_stack)?; + let value = + self.resolve_singleton(&module_ref, module_ref.transient_cache(), &resolution_stack)?; hook(value, module_ref, signal).await } @@ -232,20 +425,25 @@ impl ProviderEntry { }; let resolution_stack = new_resolution_stack(); - let value = self.resolve_singleton(&module_ref, &resolution_stack)?; + let value = + self.resolve_singleton(&module_ref, module_ref.transient_cache(), &resolution_stack)?; hook(value, module_ref, signal).await } fn factory_ref( &self, module_ref: &ModuleRef, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, resolution_stack: &ProviderResolutionStack, ) -> ModuleRef { - let factory_ref = module_ref.with_resolution_stack(Arc::clone(resolution_stack)); - match request_cache { - Some(request_cache) => factory_ref.with_request_cache(request_cache), - None => factory_ref, + let factory_ref = module_ref + .with_transient_cache(transient_cache) + .with_inquirer(self.cache_key) + .with_resolution_stack(Arc::clone(resolution_stack)); + match context_id { + Some(context_id) => factory_ref.weak_context_scope(&context_id), + None => factory_ref.without_context(), } } @@ -254,31 +452,9 @@ impl ProviderEntry { module_ref: &ModuleRef, resolution_stack: &ProviderResolutionStack, ) -> Result> { - enter_resolution_stack(resolution_stack, self.definition.token())?; - let result = self.definition.build(module_ref); - let exit_result = exit_resolution_stack(resolution_stack); - match (result, exit_result) { - (Ok(value), Ok(())) => Ok(value), - (Err(error), _) => Err(error), - (Ok(_), Err(error)) => Err(error), - } - } - - fn read_cache<'a>( - &self, - cache: &'a ProviderCache, - ) -> Result>>> { - cache - .read() - .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string())) - } - - fn write_cache<'a>( - &self, - cache: &'a ProviderCache, - ) -> Result>>> { - cache - .write() - .map_err(|_| BootError::Internal("provider cache lock is poisoned".to_string())) + let (resolution_stack, _guard) = + enter_resolution_stack(resolution_stack, self.cache_key, self.definition.token())?; + self.definition + .build(&module_ref.with_resolution_stack(resolution_stack)) } } diff --git a/src/provider/mod.rs b/src/provider/mod.rs index 6084896..a5658fb 100644 --- a/src/provider/mod.rs +++ b/src/provider/mod.rs @@ -1,4 +1,5 @@ mod cache; +mod context_id; mod definition; mod entry; mod module_ref; @@ -8,8 +9,9 @@ mod token; use std::any::Any; +pub use context_id::{ContextId, ContextIdFactory}; pub use definition::{ - FromModuleRef, ProviderBeforeApplicationShutdown, ProviderDefinition, + FromModuleRef, ProviderBeforeApplicationShutdown, ProviderDefinition, ProviderDependency, ProviderOnApplicationBootstrap, ProviderOnApplicationShutdown, ProviderOnModuleDestroy, ProviderOnModuleInit, ProviderScope, }; diff --git a/src/provider/module_ref.rs b/src/provider/module_ref.rs index c115338..d8c8041 100644 --- a/src/provider/module_ref.rs +++ b/src/provider/module_ref.rs @@ -1,10 +1,16 @@ -use super::cache::{new_provider_cache, ProviderCache}; +use super::cache::{new_provider_cache, ProviderCache, ProviderCacheKey}; use super::entry::ProviderEntry; use super::provider_ref::ProviderRef; -use super::resolution::{new_resolution_stack, ProviderResolutionStack}; -use super::{AnyProvider, FromModuleRef, ProviderDefinition, ProviderScope, ProviderToken}; -use crate::{BootError, Result}; -use std::collections::BTreeMap; +use super::resolution::{ + enter_resolution_stack, new_resolution_stack, resolution_stack_is_empty, + ProviderResolutionStack, +}; +use super::{ + AnyProvider, ContextId, ContextIdFactory, FromModuleRef, ProviderDefinition, ProviderScope, + ProviderToken, +}; +use crate::{BootError, BoxFuture, Result}; +use std::collections::{BTreeMap, BTreeSet}; use std::fmt; use std::sync::{Arc, RwLock}; @@ -14,10 +20,62 @@ pub struct ModuleRef { providers: Arc>>, provider_order: Arc>>, visible_scopes: Arc>>, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, + inquirer: Option, resolution_stack: Option, } +#[derive(Default)] +struct AsyncProviderSeedState { + complete: BTreeSet, + visiting: Vec<(ProviderCacheKey, ProviderToken)>, +} + +impl AsyncProviderSeedState { + fn enter(&mut self, entry: &ProviderEntry) -> Result { + let cache_key = entry.cache_key(); + if self.complete.contains(&cache_key) { + return Ok(false); + } + if let Some(index) = self + .visiting + .iter() + .position(|(active, _)| *active == cache_key) + { + let mut chain = self.visiting[index..] + .iter() + .map(|(_, token)| token.to_string()) + .collect::>(); + chain.push(entry.token().to_string()); + return Err(BootError::Internal(format!( + "cyclic async provider dependency detected: {}", + chain.join(" -> ") + ))); + } + + self.visiting.push((cache_key, entry.token().clone())); + Ok(true) + } + + fn exit(&mut self, cache_key: ProviderCacheKey, complete: bool) -> Result<()> { + let Some((active, _)) = self.visiting.pop() else { + return Err(BootError::Internal( + "async provider dependency stack underflow".to_string(), + )); + }; + if active != cache_key { + return Err(BootError::Internal( + "async provider dependency stack is inconsistent".to_string(), + )); + } + if complete { + self.complete.insert(cache_key); + } + Ok(()) + } +} + impl fmt::Debug for ModuleRef { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { let len = self @@ -43,29 +101,103 @@ impl ModuleRef { } pub fn request_scope(&self) -> Self { - self.with_request_cache(new_provider_cache()) + self.context_scope(&ContextIdFactory::create()) + } + + /// Bind this module view to an existing dependency-injection context. + pub fn context_scope(&self, context_id: &ContextId) -> Self { + Self { + providers: Arc::clone(&self.providers), + provider_order: Arc::clone(&self.provider_order), + visible_scopes: Arc::clone(&self.visible_scopes), + context_id: Some(context_id.clone()), + transient_cache: self.transient_cache.clone(), + inquirer: self.inquirer, + resolution_stack: self.resolution_stack.clone(), + } } - pub(crate) fn with_request_cache(&self, request_cache: ProviderCache) -> Self { + pub(crate) fn weak_context_scope(&self, context_id: &ContextId) -> Self { Self { providers: Arc::clone(&self.providers), provider_order: Arc::clone(&self.provider_order), visible_scopes: Arc::clone(&self.visible_scopes), - request_cache: Some(request_cache), + context_id: Some(context_id.downgrade()), + transient_cache: self.transient_cache.clone(), + inquirer: self.inquirer, resolution_stack: self.resolution_stack.clone(), } } + /// Return the dependency-injection context attached to this module view. + pub fn context_id(&self) -> Option<&ContextId> { + self.context_id.as_ref() + } + pub(crate) fn with_resolution_stack(&self, resolution_stack: ProviderResolutionStack) -> Self { Self { providers: Arc::clone(&self.providers), provider_order: Arc::clone(&self.provider_order), visible_scopes: Arc::clone(&self.visible_scopes), - request_cache: self.request_cache.clone(), + context_id: self.context_id.clone(), + transient_cache: self.transient_cache.clone(), + inquirer: self.inquirer, resolution_stack: Some(resolution_stack), } } + pub(crate) fn without_resolution_stack(&self) -> Self { + Self { + providers: Arc::clone(&self.providers), + provider_order: Arc::clone(&self.provider_order), + visible_scopes: Arc::clone(&self.visible_scopes), + context_id: self.context_id.clone(), + transient_cache: self.transient_cache.clone(), + inquirer: self.inquirer, + resolution_stack: None, + } + } + + pub(crate) fn with_inquirer(&self, inquirer: ProviderCacheKey) -> Self { + Self { + providers: Arc::clone(&self.providers), + provider_order: Arc::clone(&self.provider_order), + visible_scopes: Arc::clone(&self.visible_scopes), + context_id: self.context_id.clone(), + transient_cache: self.transient_cache.clone(), + inquirer: Some(inquirer), + resolution_stack: self.resolution_stack.clone(), + } + } + + pub(crate) fn without_context(&self) -> Self { + Self { + providers: Arc::clone(&self.providers), + provider_order: Arc::clone(&self.provider_order), + visible_scopes: Arc::clone(&self.visible_scopes), + context_id: None, + transient_cache: self.transient_cache.clone(), + inquirer: self.inquirer, + resolution_stack: self.resolution_stack.clone(), + } + } + + pub(crate) fn with_transient_cache(&self, transient_cache: ProviderCache) -> Self { + Self { + providers: Arc::clone(&self.providers), + provider_order: Arc::clone(&self.provider_order), + visible_scopes: Arc::clone(&self.visible_scopes), + context_id: self.context_id.clone(), + transient_cache, + inquirer: self.inquirer, + resolution_stack: self.resolution_stack.clone(), + } + } + + pub(crate) fn transient_cache(&self) -> ProviderCache { + self.transient_cache.clone() + } + pub fn register(&self, definition: ProviderDefinition) -> Result<()> { let token = definition.token().clone(); self.validate_registration(&token, &definition)?; @@ -183,6 +315,14 @@ impl ModuleRef { self.resolve_token::(&ProviderToken::of::()) } + /// Resolve a typed provider in a caller-supplied resolution context. + pub fn resolve_with_context(&self, context_id: &ContextId) -> Result> + where + T: Send + Sync + 'static, + { + self.resolve_token_with_context::(&ProviderToken::of::(), context_id) + } + /// Resolve a named provider in a fresh resolution context. pub fn resolve_named(&self, token: &str) -> Result> where @@ -191,6 +331,18 @@ impl ModuleRef { self.resolve_token::(&ProviderToken::named(token)) } + /// Resolve a named provider in a caller-supplied resolution context. + pub fn resolve_named_with_context( + &self, + token: &str, + context_id: &ContextId, + ) -> Result> + where + T: Send + Sync + 'static, + { + self.resolve_token_with_context::(&ProviderToken::named(token), context_id) + } + /// Resolve a typed provider in a fresh resolution context when it exists. pub fn resolve_optional(&self) -> Result>> where @@ -199,6 +351,14 @@ impl ModuleRef { self.resolve_optional_token::(&ProviderToken::of::()) } + /// Resolve an optional typed provider in a caller-supplied context. + pub fn resolve_optional_with_context(&self, context_id: &ContextId) -> Result>> + where + T: Send + Sync + 'static, + { + self.resolve_optional_token_with_context::(&ProviderToken::of::(), context_id) + } + /// Resolve a named provider in a fresh resolution context when it exists. pub fn resolve_optional_named(&self, token: &str) -> Result>> where @@ -207,12 +367,32 @@ impl ModuleRef { self.resolve_optional_token::(&ProviderToken::named(token)) } + /// Resolve an optional named provider in a caller-supplied context. + pub fn resolve_optional_named_with_context( + &self, + token: &str, + context_id: &ContextId, + ) -> Result>> + where + T: Send + Sync + 'static, + { + self.resolve_optional_token_with_context::(&ProviderToken::named(token), context_id) + } + /// Create an injectable value without registering it in the provider graph. pub fn create(&self) -> Result where T: FromModuleRef, { - T::from_module_ref(self) + let inquirer = ProviderCacheKey::next(); + let token = ProviderToken::of::(); + let (resolution_stack, _guard) = + enter_resolution_stack(&new_resolution_stack(), inquirer, &token)?; + let factory_ref = self + .with_transient_cache(new_provider_cache()) + .with_inquirer(inquirer) + .with_resolution_stack(resolution_stack); + T::from_module_ref(&factory_ref) } /// Create an injectable `Arc` without registering it in the provider graph. @@ -238,6 +418,49 @@ impl ModuleRef { self.contains(&ProviderToken::named(token)) } + /// Return whether a typed provider requires a request-resolution context. + /// + /// This includes explicitly request-scoped providers and singleton or + /// transient providers whose declared dependency tree reaches one. + pub fn provider_is_contextual(&self) -> Result + where + T: Send + Sync + 'static, + { + self.token_is_contextual(&ProviderToken::of::()) + } + + /// Return the provider's declared cache scope after following aliases. + pub fn provider_scope(&self) -> Result + where + T: Send + Sync + 'static, + { + self.token_scope(&ProviderToken::of::()) + } + + /// Return whether a named provider requires a request-resolution context. + pub fn named_provider_is_contextual(&self, token: &str) -> Result { + self.token_is_contextual(&ProviderToken::named(token)) + } + + /// Return a named provider's declared cache scope after following aliases. + pub fn named_provider_scope(&self, token: &str) -> Result { + self.token_scope(&ProviderToken::named(token)) + } + + pub(crate) fn token_is_contextual(&self, token: &ProviderToken) -> Result { + let entry = self + .get_entry(token)? + .ok_or_else(|| BootError::MissingProvider(token.to_string()))?; + Ok(entry.resolution_plan(self)?.is_contextual()) + } + + pub(crate) fn token_scope(&self, token: &ProviderToken) -> Result { + let entry = self + .get_entry(token)? + .ok_or_else(|| BootError::MissingProvider(token.to_string()))?; + Ok(entry.resolution_plan(self)?.scope()) + } + pub fn tokens(&self) -> Result> { let mut tokens = BTreeMap::new(); self.collect_tokens(&mut tokens)?; @@ -299,22 +522,99 @@ impl ModuleRef { pub(crate) fn initialize_local_singletons(&self) -> Result<()> { for entry in self.local_entries()? { - if entry.is_local_singleton() { + let plan = entry.resolution_plan(self)?; + self.validate_resolution_plan(&entry, plan)?; + if plan.scope() == ProviderScope::Singleton + && !plan.is_contextual() + && !entry.is_alias() + { let resolution_stack = new_resolution_stack(); - entry.resolve_singleton(self, &resolution_stack)?; + entry.resolve_singleton(self, self.transient_cache(), &resolution_stack)?; } } Ok(()) } - pub(crate) async fn initialize_local_singletons_async(&self) -> Result<()> { + pub(crate) fn validate_local_resolution_plans(&self) -> Result<()> { for entry in self.local_entries()? { - if entry.is_local_singleton() && entry.is_async_factory() { - entry.seed_singleton_async(self.clone()).await?; + let plan = entry.resolution_plan(self)?; + self.validate_resolution_plan(&entry, plan)?; + } + Ok(()) + } + + pub(crate) async fn seed_local_async_singletons(&self) -> Result<()> { + let entries = self.local_entries()?; + for entry in &entries { + let plan = entry.resolution_plan(self)?; + self.validate_resolution_plan(entry, plan)?; + } + + let mut state = AsyncProviderSeedState::default(); + for entry in entries { + let plan = entry.resolution_plan(self)?; + if plan.scope() == ProviderScope::Singleton + && !plan.is_contextual() + && !entry.is_alias() + && entry.is_async_factory() + { + self.seed_async_dependency_tree(entry, &mut state).await?; } } + Ok(()) + } - self.initialize_local_singletons() + fn seed_async_dependency_tree<'a>( + &'a self, + entry: ProviderEntry, + state: &'a mut AsyncProviderSeedState, + ) -> BoxFuture<'a, Result<()>> { + Box::pin(async move { + let cache_key = entry.cache_key(); + if !state.enter(&entry)? { + return Ok(()); + } + + let owner = entry.owner_module_ref(self); + let result = async { + let plan = entry.resolution_plan(&owner)?; + owner.validate_resolution_plan(&entry, plan)?; + for dependency in entry.eager_dependency_entries(&owner)? { + owner.seed_async_dependency_tree(dependency, state).await?; + } + if entry.is_async_factory() { + entry.seed_singleton_async(owner).await?; + } + Ok(()) + } + .await; + let exit_result = state.exit(cache_key, result.is_ok()); + match (result, exit_result) { + (Ok(()), Ok(())) => Ok(()), + (Err(error), _) => Err(error), + (Ok(()), Err(error)) => Err(error), + } + }) + } + + fn validate_resolution_plan( + &self, + entry: &ProviderEntry, + plan: super::entry::ProviderResolutionPlan, + ) -> Result<()> { + if plan.is_contextual() && entry.is_async_factory() { + return Err(BootError::Internal(format!( + "async provider `{}` cannot depend on a request-context provider", + entry.token() + ))); + } + if plan.is_contextual() && entry.has_lifecycle_hooks() { + return Err(BootError::Internal(format!( + "provider `{}` cannot use singleton lifecycle hooks because request scope propagated through its dependencies", + entry.token() + ))); + } + Ok(()) } pub(crate) fn initialize_local_providers(&self) -> Result<()> { @@ -395,47 +695,84 @@ impl ModuleRef { where T: Send + Sync + 'static, { - self.request_scope().get_token(token) + self.resolve_token_with_context(token, &ContextIdFactory::create()) + } + + pub(crate) fn resolve_token_with_context( + &self, + token: &ProviderToken, + context_id: &ContextId, + ) -> Result> + where + T: Send + Sync + 'static, + { + self.context_scope(context_id).get_token(token) } pub(crate) fn resolve_optional_token(&self, token: &ProviderToken) -> Result>> where T: Send + Sync + 'static, { - self.request_scope().get_optional_token(token) + self.resolve_optional_token_with_context(token, &ContextIdFactory::create()) + } + + pub(crate) fn resolve_optional_token_with_context( + &self, + token: &ProviderToken, + context_id: &ContextId, + ) -> Result>> + where + T: Send + Sync + 'static, + { + self.context_scope(context_id).get_optional_token(token) } fn get_any(&self, token: &ProviderToken) -> Result>> { let mut alias_path = Vec::new(); - let resolution_stack = self - .resolution_stack - .clone() - .unwrap_or_else(new_resolution_stack); - self.get_any_with_request_cache_inner( + let resolution_stack = match &self.resolution_stack { + Some(resolution_stack) if !resolution_stack_is_empty(resolution_stack) => { + Arc::clone(resolution_stack) + } + Some(_) | None => new_resolution_stack(), + }; + self.get_any_with_context_inner( token, - self.request_cache.clone(), + self.context_id.clone(), + self.transient_cache.clone(), + self.inquirer, &resolution_stack, &mut alias_path, ) } - pub(crate) fn get_any_with_request_cache_inner( + pub(crate) fn get_any_with_context_inner( &self, token: &ProviderToken, - request_cache: Option, + context_id: Option, + transient_cache: ProviderCache, + inquirer: Option, resolution_stack: &ProviderResolutionStack, alias_path: &mut Vec, ) -> Result>> { if let Some(entry) = self.read_providers()?.get(token).cloned() { return entry - .resolve(self, request_cache, resolution_stack, alias_path) + .resolve( + self, + context_id, + transient_cache, + inquirer, + resolution_stack, + alias_path, + ) .map(Some); } for scope in self.visible_scopes()? { - if let Some(value) = scope.get_any_with_request_cache_inner( + if let Some(value) = scope.get_any_with_context_inner( token, - request_cache.clone(), + context_id.clone(), + transient_cache.clone(), + inquirer, resolution_stack, alias_path, )? { @@ -446,9 +783,9 @@ impl ModuleRef { Ok(None) } - fn get_entry(&self, token: &ProviderToken) -> Result> { + pub(crate) fn get_entry(&self, token: &ProviderToken) -> Result> { if let Some(entry) = self.read_providers()?.get(token).cloned() { - return Ok(Some(entry)); + return Ok(Some(entry.with_owner(self.clone()))); } for scope in self.visible_scopes()? { diff --git a/src/provider/provider_ref.rs b/src/provider/provider_ref.rs index 44eee02..2d5309d 100644 --- a/src/provider/provider_ref.rs +++ b/src/provider/provider_ref.rs @@ -1,4 +1,4 @@ -use super::{ModuleRef, ProviderToken}; +use super::{ContextId, ModuleRef, ProviderToken}; use crate::Result; use std::fmt; use std::marker::PhantomData; @@ -26,7 +26,7 @@ where { pub fn new(module_ref: ModuleRef, token: ProviderToken) -> Self { Self { - module_ref, + module_ref: module_ref.without_resolution_stack(), token, _marker: PhantomData, } @@ -52,7 +52,17 @@ where self.module_ref.resolve_token(&self.token) } + pub fn resolve_with_context(&self, context_id: &ContextId) -> Result> { + self.module_ref + .resolve_token_with_context(&self.token, context_id) + } + pub fn resolve_optional(&self) -> Result>> { self.module_ref.resolve_optional_token(&self.token) } + + pub fn resolve_optional_with_context(&self, context_id: &ContextId) -> Result>> { + self.module_ref + .resolve_optional_token_with_context(&self.token, context_id) + } } diff --git a/src/provider/resolution.rs b/src/provider/resolution.rs index 287d710..9390150 100644 --- a/src/provider/resolution.rs +++ b/src/provider/resolution.rs @@ -1,43 +1,105 @@ +use super::cache::ProviderCacheKey; use super::ProviderToken; use crate::{BootError, Result}; -use std::sync::{Arc, RwLock}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; -pub(crate) type ProviderResolutionStack = Arc>>; +#[derive(Clone)] +pub(crate) struct ProviderResolutionFrame { + cache_key: ProviderCacheKey, + token: ProviderToken, + active: Arc, +} + +impl ProviderResolutionFrame { + fn is_active(&self) -> bool { + self.active.load(Ordering::Acquire) + } +} + +pub(crate) type ProviderResolutionStack = Arc>; + +pub(crate) struct ProviderResolutionGuard { + active: Arc, +} + +impl Drop for ProviderResolutionGuard { + fn drop(&mut self) { + self.active.store(false, Ordering::Release); + } +} pub(crate) fn new_resolution_stack() -> ProviderResolutionStack { - Arc::new(RwLock::new(Vec::new())) + Arc::new(Vec::new()) } -pub(crate) fn enter_resolution_stack( +pub(crate) fn resolution_stack_is_empty(resolution_stack: &ProviderResolutionStack) -> bool { + !resolution_stack + .iter() + .any(ProviderResolutionFrame::is_active) +} + +pub(crate) fn ensure_not_resolving( resolution_stack: &ProviderResolutionStack, + cache_key: ProviderCacheKey, token: &ProviderToken, ) -> Result<()> { - let mut stack = resolution_stack.write().map_err(|_| { - BootError::Internal("provider resolution stack lock is poisoned".to_string()) - })?; - if let Some(index) = stack.iter().position(|active| active == token) { - let mut chain = stack[index..].to_vec(); - chain.push(token.clone()); - let chain = chain + let active_frames = resolution_stack + .iter() + .filter(|frame| frame.is_active()) + .collect::>(); + if let Some(index) = active_frames + .iter() + .position(|active| active.cache_key == cache_key) + { + let mut chain = active_frames[index..] .iter() - .map(ToString::to_string) - .collect::>() - .join(" -> "); + .map(|frame| frame.token.to_string()) + .collect::>(); + chain.push(token.to_string()); return Err(BootError::Internal(format!( - "cyclic provider dependency detected: {chain}" + "cyclic provider dependency detected: {}", + chain.join(" -> ") ))); } - stack.push(token.clone()); Ok(()) } -pub(crate) fn exit_resolution_stack(resolution_stack: &ProviderResolutionStack) -> Result<()> { - let mut stack = resolution_stack.write().map_err(|_| { - BootError::Internal("provider resolution stack lock is poisoned".to_string()) - })?; - stack - .pop() - .ok_or_else(|| BootError::Internal("provider resolution stack underflow".to_string()))?; - Ok(()) +pub(crate) fn enter_resolution_stack( + resolution_stack: &ProviderResolutionStack, + cache_key: ProviderCacheKey, + token: &ProviderToken, +) -> Result<(ProviderResolutionStack, ProviderResolutionGuard)> { + ensure_not_resolving(resolution_stack, cache_key, token)?; + + let active = Arc::new(AtomicBool::new(true)); + let mut frames = resolution_stack + .iter() + .filter(|frame| frame.is_active()) + .cloned() + .collect::>(); + frames.push(ProviderResolutionFrame { + cache_key, + token: token.clone(), + active: Arc::clone(&active), + }); + Ok((Arc::new(frames), ProviderResolutionGuard { active })) +} + +pub(crate) fn resolution_chain_with( + resolution_stack: &ProviderResolutionStack, + token: &ProviderToken, +) -> String { + let mut chain = resolution_stack + .iter() + .filter(|frame| frame.is_active()) + .map(|frame| frame.token.clone()) + .collect::>(); + chain.push(token.clone()); + chain + .iter() + .map(ToString::to_string) + .collect::>() + .join(" -> ") } diff --git a/src/routing/handler.rs b/src/routing/handler.rs index 5780175..62d4a9b 100644 --- a/src/routing/handler.rs +++ b/src/routing/handler.rs @@ -1,5 +1,8 @@ +use super::route::RouteDefinition; use crate::{BootError, BootRequest, BootResponse, BoxFuture, ModuleRef, Result}; use std::future::Future; +use std::marker::PhantomData; +use std::sync::Arc; /// Type-erased route handler used by adapters. pub trait RouteHandler: Send + Sync + 'static { @@ -46,3 +49,41 @@ where } } } + +pub(crate) struct ProviderRouteHandler { + factory: F, + marker: PhantomData T>, +} + +impl ProviderRouteHandler { + pub(crate) fn new(factory: F) -> Self { + Self { + factory, + marker: PhantomData, + } + } +} + +impl RouteHandler for ProviderRouteHandler +where + T: Send + Sync + 'static, + F: Fn(Arc) -> Result + Send + Sync + 'static, +{ + fn call(&self, request: BootRequest) -> BoxFuture<'static, Result> { + let Some(module_ref) = request.module_ref().cloned() else { + return Box::pin(async { + Err(BootError::Internal( + "provider-backed route requires a module context".to_string(), + )) + }); + }; + + let route = module_ref + .get::() + .and_then(|controller| (self.factory)(controller)); + match route { + Ok(route) => route.handler.call(request), + Err(error) => Box::pin(async move { Err(error) }), + } + } +} diff --git a/src/routing/route/definition.rs b/src/routing/route/definition.rs index 5adfaf7..ed59b22 100644 --- a/src/routing/route/definition.rs +++ b/src/routing/route/definition.rs @@ -11,7 +11,7 @@ use std::sync::Arc; #[cfg(feature = "cache")] use std::time::Duration; -use crate::routing::handler::{RequestScopedRouteHandler, RouteHandler}; +use crate::routing::handler::{ProviderRouteHandler, RequestScopedRouteHandler, RouteHandler}; use crate::routing::host::{ host_param_names, host_shape_key, host_specificity, match_host_params, match_host_shape, validate_host_pattern, @@ -29,7 +29,7 @@ pub struct RouteDefinition { pub(super) method: HttpMethod, pub(super) path: String, pub(super) host: Option, - pub(super) handler: Arc, + pub(crate) handler: Arc, pub(super) middleware: Vec>, pub(super) pipes: Vec>, pub(super) guards: Vec>, @@ -87,6 +87,20 @@ impl RouteDefinition { Self::new(method, path, RequestScopedRouteHandler::new(factory)) } + /// Build a route whose handler resolves a provider from the current request + /// scope before constructing the concrete route handler. + pub fn new_provider( + method: HttpMethod, + path: impl Into, + factory: F, + ) -> Result + where + T: Send + Sync + 'static, + F: Fn(Arc) -> Result + Send + Sync + 'static, + { + Self::new(method, path, ProviderRouteHandler::::new(factory)) + } + pub fn method(&self) -> HttpMethod { self.method } diff --git a/src/routing/route/execution.rs b/src/routing/route/execution.rs index 57fe3bf..cf94ab2 100644 --- a/src/routing/route/execution.rs +++ b/src/routing/route/execution.rs @@ -1,15 +1,49 @@ use super::definition::RouteDefinition; use crate::routing::path::{join_paths, route_shape_key, route_specificity}; -use crate::{BootError, BootRequest, BootResponse, ExecutionContext, MiddlewareOutcome, Result}; +use crate::{ + BootError, BootRequest, BootResponse, CallHandler, ContextId, ContextIdFactory, + ExecutionContext, MiddlewareOutcome, Result, +}; +use std::sync::{Arc, Mutex}; + +#[derive(Clone)] +struct PipelineErrorContext { + context: Arc>, +} + +impl PipelineErrorContext { + fn new(context: ExecutionContext) -> Self { + Self { + context: Arc::new(Mutex::new(context)), + } + } + + fn replace(&self, context: ExecutionContext) { + *self + .context + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = context; + } + + fn snapshot(&self) -> ExecutionContext { + self.context + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } +} impl RouteDefinition { pub async fn call(&self, mut request: BootRequest) -> Result { + let context_id = ContextIdFactory::create(); + request = self.attach_request_context(request, &context_id); if !self.method.matches(request.method) { let message = format!("{} {}", request.method.as_str(), request.path); return self .handle_error( self.execution_context(request), BootError::MethodNotAllowed(message), + &context_id, ) .await; } @@ -22,12 +56,13 @@ impl RouteDefinition { .handle_error( self.execution_context(request), BootError::NotFound(message), + &context_id, ) .await; } Err(error) => { return self - .handle_error(self.execution_context(request), error) + .handle_error(self.execution_context(request), error, &context_id) .await; } }; @@ -40,20 +75,17 @@ impl RouteDefinition { .handle_error( self.execution_context(request), BootError::NotFound(message), + &context_id, ) .await; } Err(error) => { return self - .handle_error(self.execution_context(request), error) + .handle_error(self.execution_context(request), error, &context_id) .await; } }; request = request.with_host_params(host_params); - if let Some(module_ref) = &self.module_ref { - request = request.with_module_ref(module_ref.request_scope()); - } - #[cfg(feature = "request-context")] { let context = crate::RequestContext::from_route_request( @@ -63,24 +95,31 @@ impl RouteDefinition { self.controller_prefix.clone(), self.metadata.clone(), ); - return crate::RequestContext::scope(context, self.call_pipeline(request)).await; + return crate::RequestContext::scope(context, self.call_pipeline(request, context_id)) + .await; } #[cfg(not(feature = "request-context"))] { - self.call_pipeline(request).await + self.call_pipeline(request, context_id).await } } - async fn call_pipeline(&self, mut request: BootRequest) -> Result { + async fn call_pipeline( + &self, + mut request: BootRequest, + context_id: ContextId, + ) -> Result { for middleware in &self.middleware { let context_request = request.clone(); request = match middleware.handle(request).await { - Ok(MiddlewareOutcome::Continue(request)) => request, + Ok(MiddlewareOutcome::Continue(request)) => { + self.attach_request_context(request, &context_id) + } Ok(MiddlewareOutcome::Respond(response)) => return Ok(response), Err(error) => { return self - .handle_error(self.execution_context(context_request), error) + .handle_error(self.execution_context(context_request), error, &context_id) .await; } }; @@ -89,48 +128,92 @@ impl RouteDefinition { let context = self.execution_context(request.clone()); for guard in &self.guards { - let can_activate = match guard.inner().can_activate(context.clone()).await { + let guard = match guard.resolve(&context_id) { + Ok(guard) => guard, + Err(error) => { + return self.handle_error(context.clone(), error, &context_id).await; + } + }; + let can_activate = match guard.can_activate(context.clone()).await { Ok(can_activate) => can_activate, - Err(error) => return self.handle_error(context.clone(), error).await, + Err(error) => { + return self.handle_error(context.clone(), error, &context_id).await; + } }; if !can_activate { let message = format!("{} {}", context.method.as_str(), context.request_path); return self - .handle_error(context, BootError::Forbidden(message)) + .handle_error(context, BootError::Forbidden(message), &context_id) .await; } } - for (index, interceptor) in self.interceptors.iter().enumerate() { - if let Err(error) = interceptor.inner().before(context.clone()).await { - return self.handle_error(context.clone(), error).await; + let mut resolved_interceptors = Vec::with_capacity(self.interceptors.len()); + for interceptor in &self.interceptors { + match interceptor.resolve(&context_id) { + Ok(interceptor) => resolved_interceptors.push(interceptor), + Err(error) => { + return self.handle_error(context.clone(), error, &context_id).await; + } } + } - let mut response = match interceptor.inner().short_circuit(context.clone()).await { - Ok(Some(response)) => response, - Ok(None) => continue, - Err(error) => return self.handle_error(context.clone(), error).await, - }; + let error_context = PipelineErrorContext::new(context.clone()); + let terminal_context = context.clone(); + let terminal_error_context = error_context.clone(); + let handler_context_id = context_id.clone(); + let mut next = CallHandler::from_fn(move || { + terminal_error_context.replace(terminal_context.clone()); + self.call_handler_pipeline( + request.clone(), + terminal_error_context.clone(), + handler_context_id.clone(), + ) + }); + for interceptor in resolved_interceptors.iter().rev() { + let interceptor_context = context.clone(); + let success_context = context.clone(); + let interceptor_error_context = error_context.clone(); + let downstream = next.clone(); + next = CallHandler::from_fn(move || { + interceptor_error_context.replace(interceptor_context.clone()); + let future = interceptor.intercept(interceptor_context.clone(), downstream.clone()); + let success_context = success_context.clone(); + let interceptor_error_context = interceptor_error_context.clone(); + async move { + let result = future.await; + if result.is_ok() { + interceptor_error_context.replace(success_context); + } + result + } + }); + } - for interceptor in self.interceptors[..index].iter().rev() { - response = match interceptor.inner().after(context.clone(), response).await { - Ok(response) => response, - Err(error) => return self.handle_error(context.clone(), error).await, - }; + match next.handle().await { + Ok(response) => Ok(response), + Err(error) => { + self.handle_error(error_context.snapshot(), error, &context_id) + .await } - - return Ok(response); } + } + async fn call_handler_pipeline( + &self, + mut request: BootRequest, + error_context: PipelineErrorContext, + context_id: ContextId, + ) -> Result { for pipe in &self.pipes { let context_request = request.clone(); - request = match pipe.inner().transform(request).await { - Ok(request) => request, + let pipe = pipe.resolve(&context_id)?; + request = match pipe.transform(request).await { + Ok(request) => self.attach_request_context(request, &context_id), Err(error) => { - return self - .handle_error(self.execution_context(context_request), error) - .await; + error_context.replace(self.execution_context(context_request)); + return Err(error); } }; } @@ -141,27 +224,14 @@ impl RouteDefinition { request = match validator(request, self.validation_options) { Ok(request) => request, Err(error) => { - return self - .handle_error(self.execution_context(context_request), error) - .await; + error_context.replace(self.execution_context(context_request)); + return Err(error); } }; } } - let mut response = match self.handler.call(request).await { - Ok(response) => response, - Err(error) => return self.handle_error(context, error).await, - }; - - for interceptor in self.interceptors.iter().rev() { - response = match interceptor.inner().after(context.clone(), response).await { - Ok(response) => response, - Err(error) => return self.handle_error(context.clone(), error).await, - }; - } - - Ok(response) + self.handler.call(request).await } /// Dispatch a request through this route and convert unhandled errors into Boot HTTP responses. @@ -227,10 +297,11 @@ impl RouteDefinition { &self, context: ExecutionContext, error: BootError, + context_id: &ContextId, ) -> Result { for filter in self.filters.iter().rev() { + let filter = filter.resolve(context_id)?; if let Some(response) = filter - .inner() .catch(context.clone(), error.clone_for_filter()) .await? { @@ -239,4 +310,11 @@ impl RouteDefinition { } Err(error) } + + fn attach_request_context(&self, request: BootRequest, context_id: &ContextId) -> BootRequest { + match &self.module_ref { + Some(module_ref) => request.with_module_ref(module_ref.context_scope(context_id)), + None => request, + } + } } diff --git a/src/testing.rs b/src/testing.rs index 2d84c95..964515e 100644 --- a/src/testing.rs +++ b/src/testing.rs @@ -314,6 +314,8 @@ impl TestingModuleBuilder { .into_iter() .map(|pattern| pattern.with_pipeline_overrides(&self.pipeline_overrides)) .collect(); + self.pipeline_overrides + .apply_to_filters(&mut app.global_filters); } Ok(TestingModule { app }) } @@ -336,6 +338,8 @@ impl TestingModuleBuilder { .into_iter() .map(|pattern| pattern.with_pipeline_overrides(&self.pipeline_overrides)) .collect(); + self.pipeline_overrides + .apply_to_filters(&mut app.global_filters); } Ok(TestingModule { app }) } diff --git a/src/transport/mod.rs b/src/transport/mod.rs index fff18ac..f87315c 100644 --- a/src/transport/mod.rs +++ b/src/transport/mod.rs @@ -1,16 +1,16 @@ -use crate::pipeline::{PipelineComponent, PipelineOverrides}; use crate::{ - catch_errors, validate_json_value_with_options, validate_value, BootError, BootErrorKind, - BoxFuture, ExecutionContext, ExecutionInterceptor, ExecutionTransportKind, Guard, Result, - TransportExceptionFilter, Validate, ValidationOptions, ValidationSchema, + validate_value, BootError, BootRequest, BoxFuture, CallHandler, ExecutionContext, + ExecutionInterceptor, ExecutionTransportKind, Guard, Result, Validate, }; use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; use serde_json::Value; -use std::collections::BTreeMap; use std::future::Future; use std::ops::Deref; -use std::sync::Arc; + +mod pattern; + +pub use self::pattern::MessagePatternDefinition; #[cfg(feature = "grpc-transport")] mod grpc; @@ -273,12 +273,13 @@ pub struct TransportContext { } impl TransportContext { - fn new(definition: &MessagePatternDefinition, pattern: &str) -> Self { + fn new(definition: &MessagePatternDefinition, pattern: &str, request: BootRequest) -> Self { let pattern = pattern.to_string(); - let kind = definition.kind; - let module_name = definition.module_name.clone(); - let metadata = definition.metadata.clone(); + let kind = definition.kind(); + let module_name = definition.module_name().map(str::to_string); + let metadata = definition.metadata().clone(); let execution_context = ExecutionContext::transport( + request, pattern.clone(), ExecutionTransportKind::from(kind), module_name.clone(), @@ -363,6 +364,23 @@ where /// Around-handler hook for transport message patterns. pub trait TransportInterceptor: Send + Sync + 'static { + /// Run around the remaining transport pipeline. + /// + /// Override this method to recover downstream errors, retry the remaining + /// pipeline, or return a reply without calling `next`. The default + /// implementation preserves the legacy `before` and `after` hook behavior. + fn intercept<'a>( + &'a self, + context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.before(context.clone()).await?; + let reply = next.handle().await?; + self.after(context, reply).await + }) + } + fn before(&self, _context: TransportContext) -> BoxFuture<'static, Result<()>> { Box::pin(async { Ok(()) }) } @@ -401,509 +419,6 @@ where } } -type TransportHandlerFuture = BoxFuture<'static, Result>>; -type MessageValidator = - Arc Result + Send + Sync>; - -trait TransportMessageHandler: Send + Sync + 'static { - fn call(&self, message: TransportMessage) -> TransportHandlerFuture; -} - -struct TransportHandlerAdapter { - handler: H, -} - -impl TransportMessageHandler for TransportHandlerAdapter -where - H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - R: IntoTransportReply + Send + 'static, -{ - fn call(&self, message: TransportMessage) -> TransportHandlerFuture { - let future = (self.handler)(message); - Box::pin(async move { Ok(future.await?.into_transport_reply()) }) - } -} - -/// Framework-neutral message pattern handler definition. -#[derive(Clone)] -pub struct MessagePatternDefinition { - pattern: String, - kind: MessagePatternKind, - handler: Arc, - pipes: Vec>, - guards: Vec>, - interceptors: Vec>, - filters: Vec>, - validators: Vec, - validation_enabled: bool, - validation_disabled: bool, - validation_options: ValidationOptions, - metadata: BTreeMap, - module_name: Option, -} - -impl MessagePatternDefinition { - pub fn request(pattern: impl Into, handler: H) -> Result - where - H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - R: IntoTransportReply + Send + 'static, - { - Self::new(pattern, MessagePatternKind::RequestResponse, handler) - } - - pub fn event(pattern: impl Into, handler: H) -> Result - where - H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - { - Self::new(pattern, MessagePatternKind::Event, handler) - } - - pub fn request_json(pattern: impl Into, handler: H) -> Result - where - T: DeserializeOwned + Send + 'static, - H: Fn(T) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - R: Serialize + Send + 'static, - { - Self::request(pattern, move |message: TransportMessage| { - let payload = message.data_as::(); - let future = payload.map(&handler); - async move { - let response = future?.await?; - TransportReply::json(&response) - } - }) - } - - pub fn request_validated_json( - pattern: impl Into, - handler: H, - ) -> Result - where - T: DeserializeOwned + Validate + Send + 'static, - H: Fn(T) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - R: Serialize + Send + 'static, - { - Self::request(pattern, move |message: TransportMessage| { - let payload = message.validated_data::(); - let future = payload.map(&handler); - async move { - let response = future?.await?; - TransportReply::json(&response) - } - }) - } - - pub fn event_json(pattern: impl Into, handler: H) -> Result - where - T: DeserializeOwned + Send + 'static, - H: Fn(T) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - { - Self::event(pattern, move |message: TransportMessage| { - let payload = message.data_as::(); - let future = payload.map(&handler); - async move { future?.await } - }) - } - - pub fn event_validated_json(pattern: impl Into, handler: H) -> Result - where - T: DeserializeOwned + Validate + Send + 'static, - H: Fn(T) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - { - Self::event(pattern, move |message: TransportMessage| { - let payload = message.validated_data::(); - let future = payload.map(&handler); - async move { future?.await } - }) - } - - fn new( - pattern: impl Into, - kind: MessagePatternKind, - handler: H, - ) -> Result - where - H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - R: IntoTransportReply + Send + 'static, - { - let pattern = pattern.into(); - validate_pattern(&pattern)?; - Ok(Self { - pattern, - kind, - handler: Arc::new(TransportHandlerAdapter { handler }), - pipes: Vec::new(), - guards: Vec::new(), - interceptors: Vec::new(), - filters: Vec::new(), - validators: Vec::new(), - validation_enabled: false, - validation_disabled: false, - validation_options: ValidationOptions::default(), - metadata: BTreeMap::new(), - module_name: None, - }) - } - - pub fn pattern(&self) -> &str { - &self.pattern - } - - pub fn kind(&self) -> MessagePatternKind { - self.kind - } - - pub fn module_name(&self) -> Option<&str> { - self.module_name.as_deref() - } - - pub fn metadata(&self) -> &BTreeMap { - &self.metadata - } - - pub fn metadata_value(&self, key: &str) -> Option<&Value> { - self.metadata.get(key) - } - - pub fn with_metadata(self, key: impl Into, value: V) -> Result - where - V: Serialize, - { - let key = key.into(); - let value = serde_json::to_value(value).map_err(|error| { - BootError::Internal(format!( - "failed to serialize message pattern metadata `{key}`: {error}" - )) - })?; - Ok(self.with_metadata_value(key, value)) - } - - pub fn with_metadata_value(mut self, key: impl Into, value: Value) -> Self { - self.metadata.insert(key.into(), value); - self - } - - pub fn with_pipe

(mut self, pipe: P) -> Self - where - P: TransportPipe, - { - self.pipes - .push(PipelineComponent::::new(pipe)); - self - } - - pub fn with_guard(mut self, guard: G) -> Self - where - G: TransportGuard, - { - self.guards - .push(PipelineComponent::::new(guard)); - self - } - - pub fn with_execution_guard(mut self, guard: G) -> Self - where - G: Guard, - { - self.guards - .push(PipelineComponent::::new( - ExecutionTransportGuard { inner: guard }, - )); - self - } - - pub(crate) fn with_execution_pipeline_prefix( - mut self, - guards: &[Arc], - interceptors: &[Arc], - ) -> Self { - self.guards = prepend_execution_guards(guards, self.guards); - self.interceptors = prepend_execution_interceptors(interceptors, self.interceptors); - self - } - - pub(crate) fn with_guard_prefix(mut self, guards: &[Arc]) -> Self { - let mut merged = guards - .iter() - .cloned() - .map(PipelineComponent::::from_arc) - .collect::>(); - merged.extend(self.guards); - self.guards = merged; - self - } - - pub(crate) fn with_interceptor_prefix( - mut self, - interceptors: &[Arc], - ) -> Self { - let mut merged = interceptors - .iter() - .cloned() - .map(PipelineComponent::::from_arc) - .collect::>(); - merged.extend(self.interceptors); - self.interceptors = merged; - self - } - - pub(crate) fn with_pipe_prefix(mut self, pipes: &[Arc]) -> Self { - let mut merged = pipes - .iter() - .cloned() - .map(PipelineComponent::::from_arc) - .collect::>(); - merged.extend(self.pipes); - self.pipes = merged; - self - } - - pub(crate) fn with_filter_prefix( - mut self, - filters: &[Arc], - ) -> Self { - let mut merged = filters - .iter() - .cloned() - .map(PipelineComponent::::from_arc) - .collect::>(); - merged.extend(self.filters); - self.filters = merged; - self - } - - pub fn with_interceptor(mut self, interceptor: I) -> Self - where - I: TransportInterceptor, - { - self.interceptors - .push(PipelineComponent::::new( - interceptor, - )); - self - } - - pub fn with_execution_interceptor(mut self, interceptor: I) -> Self - where - I: ExecutionInterceptor, - { - self.interceptors - .push(PipelineComponent::::new( - ExecutionTransportInterceptor { inner: interceptor }, - )); - self - } - - pub fn with_filter(mut self, filter: F) -> Self - where - F: TransportExceptionFilter, - { - self.filters - .push(PipelineComponent::::new( - filter, - )); - self - } - - pub fn with_catch_filter(self, kinds: I, filter: F) -> Self - where - I: IntoIterator, - F: TransportExceptionFilter, - { - self.with_filter(catch_errors(kinds, filter)) - } - - pub(crate) fn with_pipeline_overrides(mut self, overrides: &PipelineOverrides) -> Self { - overrides.apply_to_transport_pipes(&mut self.pipes); - overrides.apply_to_transport_guards(&mut self.guards); - overrides.apply_to_transport_interceptors(&mut self.interceptors); - overrides.apply_to_transport_filters(&mut self.filters); - self - } - - pub fn with_validation(mut self) -> Self { - self.validation_enabled = true; - self.validation_disabled = false; - self - } - - pub fn with_validation_options(mut self, options: ValidationOptions) -> Self { - self.validation_enabled = true; - self.validation_disabled = false; - self.validation_options = self.validation_options.merge(options); - self - } - - pub fn without_validation(mut self) -> Self { - self.validation_enabled = false; - self.validation_disabled = true; - self - } - - pub(crate) fn with_validation_prefix( - mut self, - validation_enabled: bool, - validation_options: ValidationOptions, - ) -> Self { - if !self.validation_disabled { - self.validation_enabled = validation_enabled || self.validation_enabled; - self.validation_options = validation_options.merge(self.validation_options); - } - self - } - - pub fn with_payload_validation(mut self) -> Self - where - T: DeserializeOwned + Validate + 'static, - { - self.validators.push(Arc::new(|message, _| { - message.validated_data::().map(|_| message) - })); - self.with_validation() - } - - pub fn with_payload_validation_options(mut self, options: ValidationOptions) -> Self - where - T: DeserializeOwned + Serialize + Validate + ValidationSchema + 'static, - { - self.validators - .push(Arc::new(move |mut message, inherited_options| { - let options = inherited_options.merge(options); - let data = validate_json_value_with_options::( - message.data.clone(), - options, - "message property", - )?; - if options.transform || options.whitelist { - message.data = data; - } - Ok(message) - })); - self.with_validation() - } - - pub async fn dispatch(&self, message: TransportMessage) -> Result> { - let context = TransportContext::new(self, &message.pattern); - match self.dispatch_pipeline(message, context.clone()).await { - Ok(reply) => Ok(reply), - Err(error) => self.handle_error(context, error).await, - } - } - - async fn dispatch_pipeline( - &self, - mut message: TransportMessage, - context: TransportContext, - ) -> Result> { - if message.pattern != self.pattern { - return Err(BootError::NotFound(format!( - "message pattern {}", - message.pattern - ))); - } - - for guard in &self.guards { - let can_activate = guard.inner().can_activate(context.clone()).await?; - if !can_activate { - return Err(BootError::Forbidden(format!( - "message pattern {}", - message.pattern - ))); - } - } - - for interceptor in &self.interceptors { - interceptor.inner().before(context.clone()).await?; - } - - for pipe in &self.pipes { - message = pipe.inner().transform(message).await?; - } - - if self.validation_enabled { - for validator in &self.validators { - message = validator(message, self.validation_options)?; - } - } - - let mut reply = self.handler.call(message).await?; - if self.kind == MessagePatternKind::Event { - reply = None; - } - - for interceptor in self.interceptors.iter().rev() { - reply = interceptor.inner().after(context.clone(), reply).await?; - } - Ok(reply) - } - - async fn handle_error( - &self, - context: TransportContext, - error: BootError, - ) -> Result> { - for filter in self.filters.iter().rev() { - if let Some(response) = filter - .inner() - .catch(context.clone(), error.clone_for_filter()) - .await? - { - return Ok(if self.kind == MessagePatternKind::Event { - None - } else { - response.into_reply() - }); - } - } - Err(error) - } - - pub(crate) fn with_module_name(mut self, module_name: &str) -> Self { - self.module_name = Some(module_name.to_string()); - self - } -} - -fn prepend_execution_guards( - prefix: &[Arc], - values: Vec>, -) -> Vec> { - let mut merged = prefix - .iter() - .cloned() - .map(|guard| { - PipelineComponent::::new(ExecutionTransportGuard { inner: guard }) - }) - .collect::>(); - merged.extend(values); - merged -} - -fn prepend_execution_interceptors( - prefix: &[Arc], - values: Vec>, -) -> Vec> { - let mut merged = prefix - .iter() - .cloned() - .map(|interceptor| { - PipelineComponent::::new(ExecutionTransportInterceptor { - inner: interceptor, - }) - }) - .collect::>(); - merged.extend(values); - merged -} - /// Adapter trait for message transports such as in-process, Redis, NATS, or Kafka. pub trait MessageTransport { type Output; @@ -950,12 +465,3 @@ impl InProcessTransportClient { self.app.emit_message(message).await } } - -fn validate_pattern(pattern: &str) -> Result<()> { - if pattern.trim().is_empty() { - return Err(BootError::BadRequest( - "message pattern cannot be empty".to_string(), - )); - } - Ok(()) -} diff --git a/src/transport/pattern.rs b/src/transport/pattern.rs new file mode 100644 index 0000000..2108d43 --- /dev/null +++ b/src/transport/pattern.rs @@ -0,0 +1,762 @@ +use super::{ + ExecutionTransportGuard, ExecutionTransportInterceptor, IntoTransportReply, MessagePatternKind, + TransportContext, TransportGuard, TransportInterceptor, TransportMessage, TransportPipe, + TransportReply, +}; +use crate::pipeline::{PipelineComponent, PipelineOverrides, ProviderEnhancerComponents}; +use crate::{ + catch_errors, validate_json_value_with_options, BootError, BootErrorKind, BootRequest, + BoxFuture, CallHandler, ContextId, ContextIdFactory, ExecutionInterceptor, Guard, HttpMethod, + ModuleRef, Result, TransportExceptionFilter, Validate, ValidationOptions, ValidationSchema, +}; +use serde::de::DeserializeOwned; +use serde::Serialize; +use serde_json::Value; +use std::collections::BTreeMap; +use std::future::Future; +use std::sync::{Arc, Mutex}; + +type TransportHandlerFuture = BoxFuture<'static, Result>>; +type MessageValidator = + Arc Result + Send + Sync>; + +trait TransportMessageHandler: Send + Sync + 'static { + fn call(&self, message: TransportMessage) -> TransportHandlerFuture; +} + +struct TransportHandlerAdapter { + handler: H, +} + +impl TransportMessageHandler for TransportHandlerAdapter +where + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: IntoTransportReply + Send + 'static, +{ + fn call(&self, message: TransportMessage) -> TransportHandlerFuture { + let future = (self.handler)(message); + Box::pin(async move { Ok(future.await?.into_transport_reply()) }) + } +} + +type ScopedTransportHandlerFactory = + dyn Fn(&ModuleRef) -> Result> + Send + Sync; + +#[derive(Clone)] +enum TransportHandlerDefinition { + Static(Arc), + Scoped(Arc), +} + +impl TransportHandlerDefinition { + fn resolve( + &self, + module_ref: Option<&ModuleRef>, + pattern: &str, + ) -> Result> { + match self { + Self::Static(handler) => Ok(Arc::clone(handler)), + Self::Scoped(factory) => { + let module_ref = module_ref.ok_or_else(|| { + BootError::Internal(format!( + "scoped transport message pattern `{pattern}` requires a declaring or default module context" + )) + })?; + factory(module_ref) + } + } + } + + fn is_scoped(&self) -> bool { + matches!(self, Self::Scoped(_)) + } +} + +#[derive(Default)] +struct DispatchHandlerCache { + handler: Mutex>>, +} + +impl DispatchHandlerCache { + fn resolve( + &self, + definition: &TransportHandlerDefinition, + module_ref: Option<&ModuleRef>, + pattern: &str, + ) -> Result> { + let mut cached = self + .handler + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(handler) = cached.as_ref() { + return Ok(Arc::clone(handler)); + } + + let handler = definition.resolve(module_ref, pattern)?; + *cached = Some(Arc::clone(&handler)); + Ok(handler) + } +} + +#[derive(Clone)] +struct MessageDispatchState { + context_id: ContextId, + module_ref: Option, + handler: Arc, +} + +impl MessageDispatchState { + fn new(module_ref: Option<&ModuleRef>) -> Self { + let context_id = ContextIdFactory::create(); + let module_ref = module_ref.map(|module_ref| module_ref.context_scope(&context_id)); + Self { + context_id, + module_ref, + handler: Arc::new(DispatchHandlerCache::default()), + } + } + + fn request(&self) -> BootRequest { + let request = BootRequest::new(HttpMethod::Post, "/__transport"); + match &self.module_ref { + Some(module_ref) => request.with_module_ref(module_ref.clone()), + None => request, + } + } +} + +/// Framework-neutral message pattern handler definition. +#[derive(Clone)] +pub struct MessagePatternDefinition { + pattern: String, + kind: MessagePatternKind, + handler: TransportHandlerDefinition, + pipes: Vec>, + guards: Vec>, + interceptors: Vec>, + filters: Vec>, + validators: Vec, + validation_enabled: bool, + validation_disabled: bool, + validation_options: ValidationOptions, + metadata: BTreeMap, + module_name: Option, + module_ref: Option, +} + +impl MessagePatternDefinition { + pub fn request(pattern: impl Into, handler: H) -> Result + where + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: IntoTransportReply + Send + 'static, + { + Self::new(pattern, MessagePatternKind::RequestResponse, handler) + } + + /// Build a request-response pattern whose handler is created from the + /// current message dispatch's dependency-injection scope. + pub fn request_scoped(pattern: impl Into, factory: F) -> Result + where + F: Fn(&ModuleRef) -> Result + Send + Sync + 'static, + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: IntoTransportReply + Send + 'static, + { + Self::new_scoped(pattern, MessagePatternKind::RequestResponse, factory) + } + + pub fn event(pattern: impl Into, handler: H) -> Result + where + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + Self::new(pattern, MessagePatternKind::Event, handler) + } + + /// Build an event pattern whose handler is created from the current + /// message dispatch's dependency-injection scope. + pub fn event_scoped(pattern: impl Into, factory: F) -> Result + where + F: Fn(&ModuleRef) -> Result + Send + Sync + 'static, + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + Self::new_scoped(pattern, MessagePatternKind::Event, factory) + } + + pub fn request_json(pattern: impl Into, handler: H) -> Result + where + T: DeserializeOwned + Send + 'static, + H: Fn(T) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: Serialize + Send + 'static, + { + Self::request(pattern, move |message: TransportMessage| { + let payload = message.data_as::(); + let future = payload.map(&handler); + async move { + let response = future?.await?; + TransportReply::json(&response) + } + }) + } + + pub fn request_validated_json( + pattern: impl Into, + handler: H, + ) -> Result + where + T: DeserializeOwned + Validate + Send + 'static, + H: Fn(T) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: Serialize + Send + 'static, + { + Self::request(pattern, move |message: TransportMessage| { + let payload = message.validated_data::(); + let future = payload.map(&handler); + async move { + let response = future?.await?; + TransportReply::json(&response) + } + }) + } + + pub fn event_json(pattern: impl Into, handler: H) -> Result + where + T: DeserializeOwned + Send + 'static, + H: Fn(T) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + Self::event(pattern, move |message: TransportMessage| { + let payload = message.data_as::(); + let future = payload.map(&handler); + async move { future?.await } + }) + } + + pub fn event_validated_json(pattern: impl Into, handler: H) -> Result + where + T: DeserializeOwned + Validate + Send + 'static, + H: Fn(T) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + Self::event(pattern, move |message: TransportMessage| { + let payload = message.validated_data::(); + let future = payload.map(&handler); + async move { future?.await } + }) + } + + fn new( + pattern: impl Into, + kind: MessagePatternKind, + handler: H, + ) -> Result + where + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: IntoTransportReply + Send + 'static, + { + Self::from_handler( + pattern, + kind, + TransportHandlerDefinition::Static(Arc::new(TransportHandlerAdapter { handler })), + ) + } + + fn new_scoped( + pattern: impl Into, + kind: MessagePatternKind, + factory: F, + ) -> Result + where + F: Fn(&ModuleRef) -> Result + Send + Sync + 'static, + H: Fn(TransportMessage) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + R: IntoTransportReply + Send + 'static, + { + let factory = move |module_ref: &ModuleRef| { + let handler = factory(module_ref)?; + Ok(Arc::new(TransportHandlerAdapter { handler }) as Arc) + }; + Self::from_handler( + pattern, + kind, + TransportHandlerDefinition::Scoped(Arc::new(factory)), + ) + } + + fn from_handler( + pattern: impl Into, + kind: MessagePatternKind, + handler: TransportHandlerDefinition, + ) -> Result { + let pattern = pattern.into(); + validate_pattern(&pattern)?; + Ok(Self { + pattern, + kind, + handler, + pipes: Vec::new(), + guards: Vec::new(), + interceptors: Vec::new(), + filters: Vec::new(), + validators: Vec::new(), + validation_enabled: false, + validation_disabled: false, + validation_options: ValidationOptions::default(), + metadata: BTreeMap::new(), + module_name: None, + module_ref: None, + }) + } + + pub fn pattern(&self) -> &str { + &self.pattern + } + + pub fn kind(&self) -> MessagePatternKind { + self.kind + } + + /// Return whether this pattern constructs its handler from each dispatch scope. + pub fn is_scoped(&self) -> bool { + self.handler.is_scoped() + } + + pub fn module_name(&self) -> Option<&str> { + self.module_name.as_deref() + } + + pub fn metadata(&self) -> &BTreeMap { + &self.metadata + } + + pub fn metadata_value(&self, key: &str) -> Option<&Value> { + self.metadata.get(key) + } + + pub fn with_metadata(self, key: impl Into, value: V) -> Result + where + V: Serialize, + { + let key = key.into(); + let value = serde_json::to_value(value).map_err(|error| { + BootError::Internal(format!( + "failed to serialize message pattern metadata `{key}`: {error}" + )) + })?; + Ok(self.with_metadata_value(key, value)) + } + + pub fn with_metadata_value(mut self, key: impl Into, value: Value) -> Self { + self.metadata.insert(key.into(), value); + self + } + + pub fn with_pipe

(mut self, pipe: P) -> Self + where + P: TransportPipe, + { + self.pipes + .push(PipelineComponent::::new(pipe)); + self + } + + pub fn with_guard(mut self, guard: G) -> Self + where + G: TransportGuard, + { + self.guards + .push(PipelineComponent::::new(guard)); + self + } + + pub fn with_execution_guard(mut self, guard: G) -> Self + where + G: Guard, + { + self.guards + .push(PipelineComponent::::new( + ExecutionTransportGuard { inner: guard }, + )); + self + } + + pub(crate) fn with_execution_pipeline_prefix( + mut self, + guards: &[Arc], + interceptors: &[Arc], + ) -> Self { + self.guards = prepend_execution_guards(guards, self.guards); + self.interceptors = prepend_execution_interceptors(interceptors, self.interceptors); + self + } + + pub(crate) fn with_guard_prefix(mut self, guards: &[Arc]) -> Self { + let mut merged = guards + .iter() + .cloned() + .map(PipelineComponent::::from_arc) + .collect::>(); + merged.extend(self.guards); + self.guards = merged; + self + } + + pub(crate) fn with_interceptor_prefix( + mut self, + interceptors: &[Arc], + ) -> Self { + let mut merged = interceptors + .iter() + .cloned() + .map(PipelineComponent::::from_arc) + .collect::>(); + merged.extend(self.interceptors); + self.interceptors = merged; + self + } + + pub(crate) fn with_pipe_prefix(mut self, pipes: &[Arc]) -> Self { + let mut merged = pipes + .iter() + .cloned() + .map(PipelineComponent::::from_arc) + .collect::>(); + merged.extend(self.pipes); + self.pipes = merged; + self + } + + pub(crate) fn with_filter_prefix( + mut self, + filters: &[Arc], + ) -> Self { + let mut merged = filters + .iter() + .cloned() + .map(PipelineComponent::::from_arc) + .collect::>(); + merged.extend(self.filters); + self.filters = merged; + self + } + + pub fn with_interceptor(mut self, interceptor: I) -> Self + where + I: TransportInterceptor, + { + self.interceptors + .push(PipelineComponent::::new( + interceptor, + )); + self + } + + pub fn with_execution_interceptor(mut self, interceptor: I) -> Self + where + I: ExecutionInterceptor, + { + self.interceptors + .push(PipelineComponent::::new( + ExecutionTransportInterceptor { inner: interceptor }, + )); + self + } + + pub fn with_filter(mut self, filter: F) -> Self + where + F: TransportExceptionFilter, + { + self.filters + .push(PipelineComponent::::new( + filter, + )); + self + } + + pub fn with_catch_filter(self, kinds: I, filter: F) -> Self + where + I: IntoIterator, + F: TransportExceptionFilter, + { + self.with_filter(catch_errors(kinds, filter)) + } + + pub(crate) fn with_pipeline_overrides(mut self, overrides: &PipelineOverrides) -> Self { + overrides.apply_to_transport_pipes(&mut self.pipes); + overrides.apply_to_transport_guards(&mut self.guards); + overrides.apply_to_transport_interceptors(&mut self.interceptors); + overrides.apply_to_transport_filters(&mut self.filters); + self + } + + pub(crate) fn with_provider_enhancer_prefix( + mut self, + enhancers: &ProviderEnhancerComponents, + ) -> Self { + let mut pipes = enhancers.transport_pipes.clone(); + pipes.extend(self.pipes); + self.pipes = pipes; + + let mut guards = enhancers.transport_guards.clone(); + guards.extend(self.guards); + self.guards = guards; + + let mut interceptors = enhancers.transport_interceptors.clone(); + interceptors.extend(self.interceptors); + self.interceptors = interceptors; + + let mut filters = enhancers.transport_filters.clone(); + filters.extend(self.filters); + self.filters = filters; + self + } + + pub fn with_validation(mut self) -> Self { + self.validation_enabled = true; + self.validation_disabled = false; + self + } + + pub fn with_validation_options(mut self, options: ValidationOptions) -> Self { + self.validation_enabled = true; + self.validation_disabled = false; + self.validation_options = self.validation_options.merge(options); + self + } + + pub fn without_validation(mut self) -> Self { + self.validation_enabled = false; + self.validation_disabled = true; + self + } + + pub(crate) fn with_validation_prefix( + mut self, + validation_enabled: bool, + validation_options: ValidationOptions, + ) -> Self { + if !self.validation_disabled { + self.validation_enabled = validation_enabled || self.validation_enabled; + self.validation_options = validation_options.merge(self.validation_options); + } + self + } + + pub fn with_payload_validation(mut self) -> Self + where + T: DeserializeOwned + Validate + 'static, + { + self.validators.push(Arc::new(|message, _| { + message.validated_data::().map(|_| message) + })); + self.with_validation() + } + + pub fn with_payload_validation_options(mut self, options: ValidationOptions) -> Self + where + T: DeserializeOwned + Serialize + Validate + ValidationSchema + 'static, + { + self.validators + .push(Arc::new(move |mut message, inherited_options| { + let options = inherited_options.merge(options); + let data = validate_json_value_with_options::( + message.data.clone(), + options, + "message property", + )?; + if options.transform || options.whitelist { + message.data = data; + } + Ok(message) + })); + self.with_validation() + } + + pub async fn dispatch(&self, message: TransportMessage) -> Result> { + let state = MessageDispatchState::new(self.module_ref.as_ref()); + let context = TransportContext::new(self, &message.pattern, state.request()); + match self + .dispatch_pipeline(message, context.clone(), state.clone()) + .await + { + Ok(reply) => Ok(reply), + Err(error) => self.handle_error(context, error, &state.context_id).await, + } + } + + async fn dispatch_pipeline( + &self, + message: TransportMessage, + context: TransportContext, + state: MessageDispatchState, + ) -> Result> { + if message.pattern != self.pattern { + return Err(BootError::NotFound(format!( + "message pattern {}", + message.pattern + ))); + } + + for guard in &self.guards { + let can_activate = guard + .resolve(&state.context_id)? + .can_activate(context.clone()) + .await?; + if !can_activate { + return Err(BootError::Forbidden(format!( + "message pattern {}", + message.pattern + ))); + } + } + + let reply = self + .dispatch_interceptor_chain(0, context, message, state) + .await?; + Ok(if self.kind == MessagePatternKind::Event { + None + } else { + reply + }) + } + + fn dispatch_interceptor_chain<'a>( + &'a self, + index: usize, + context: TransportContext, + message: TransportMessage, + state: MessageDispatchState, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + let Some(interceptor) = self.interceptors.get(index) else { + return self.dispatch_handler_pipeline(message, state).await; + }; + let interceptor = interceptor.resolve(&state.context_id)?; + + let next_context = context.clone(); + let next_message = message.clone(); + let next_state = state.clone(); + let next = CallHandler::from_fn(move || { + self.dispatch_interceptor_chain( + index + 1, + next_context.clone(), + next_message.clone(), + next_state.clone(), + ) + }); + interceptor.intercept(context, next).await + }) + } + + async fn dispatch_handler_pipeline( + &self, + mut message: TransportMessage, + state: MessageDispatchState, + ) -> Result> { + for pipe in &self.pipes { + message = pipe.resolve(&state.context_id)?.transform(message).await?; + } + + if self.validation_enabled { + for validator in &self.validators { + message = validator(message, self.validation_options)?; + } + } + + let handler = state.handler.resolve( + &self.handler, + state.module_ref.as_ref(), + self.pattern.as_str(), + )?; + let mut reply = handler.call(message).await?; + if self.kind == MessagePatternKind::Event { + reply = None; + } + Ok(reply) + } + + async fn handle_error( + &self, + context: TransportContext, + error: BootError, + context_id: &ContextId, + ) -> Result> { + for filter in self.filters.iter().rev() { + let filter = filter.resolve(context_id)?; + if let Some(response) = filter + .catch(context.clone(), error.clone_for_filter()) + .await? + { + return Ok(if self.kind == MessagePatternKind::Event { + None + } else { + response.into_reply() + }); + } + } + Err(error) + } + + pub(crate) fn with_module_name(mut self, module_name: &str) -> Self { + self.module_name = Some(module_name.to_string()); + self + } + + pub(crate) fn with_module_ref(mut self, module_ref: ModuleRef) -> Self { + self.module_ref = Some(module_ref); + self + } + + pub(crate) fn with_default_module_ref(mut self, module_ref: ModuleRef) -> Self { + if self.module_ref.is_none() { + self.module_ref = Some(module_ref); + } + self + } +} + +fn prepend_execution_guards( + prefix: &[Arc], + values: Vec>, +) -> Vec> { + let mut merged = prefix + .iter() + .cloned() + .map(|guard| { + PipelineComponent::::new(ExecutionTransportGuard { inner: guard }) + }) + .collect::>(); + merged.extend(values); + merged +} + +fn prepend_execution_interceptors( + prefix: &[Arc], + values: Vec>, +) -> Vec> { + let mut merged = prefix + .iter() + .cloned() + .map(|interceptor| { + PipelineComponent::::new(ExecutionTransportInterceptor { + inner: interceptor, + }) + }) + .collect::>(); + merged.extend(values); + merged +} + +fn validate_pattern(pattern: &str) -> Result<()> { + if pattern.trim().is_empty() { + return Err(BootError::BadRequest( + "message pattern cannot be empty".to_string(), + )); + } + Ok(()) +} diff --git a/src/websocket/connection.rs b/src/websocket/connection.rs index 3b259f4..806104a 100644 --- a/src/websocket/connection.rs +++ b/src/websocket/connection.rs @@ -3,7 +3,10 @@ use super::gateway::WebSocketGatewayDefinition; use super::message::{send_to_outbounds, WebSocketMessage, WebSocketOutbound}; use super::server::WebSocketGatewayServer; use super::state::normalize_room; -use crate::{BootError, BootRequest, BoxFuture, Result}; +use crate::{ + BootError, BootRequest, BoxFuture, CallHandler, ContextId, ContextIdFactory, Result, + WebSocketInterceptor, +}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; @@ -116,17 +119,32 @@ impl WebSocketGatewayConnection { } pub async fn dispatch(&self, message: WebSocketMessage) -> Result> { - let context = WebSocketContext::new(&self.gateway, self.request.clone(), &message.event); - match self.dispatch_pipeline(message, context.clone()).await { + let context_id = ContextIdFactory::create(); + let mut connection = self.clone(); + if let Some(module_ref) = &connection.gateway.module_ref { + connection.request = connection + .request + .with_module_ref(module_ref.context_scope(&context_id)); + } + let context = WebSocketContext::new( + &connection.gateway, + connection.request.clone(), + &message.event, + ); + match connection + .dispatch_pipeline(message, context.clone(), context_id.clone()) + .await + { Ok(reply) => Ok(reply), - Err(error) => self.handle_error(context, error).await, + Err(error) => connection.handle_error(context, error, &context_id).await, } } async fn dispatch_pipeline( &self, - mut message: WebSocketMessage, + message: WebSocketMessage, context: WebSocketContext, + context_id: ContextId, ) -> Result> { let event = message.event.clone(); let handler = self.gateway.handlers.get(&event).cloned().ok_or_else(|| { @@ -134,7 +152,10 @@ impl WebSocketGatewayConnection { })?; for guard in &self.gateway.guards { - let can_activate = guard.inner().can_activate(context.clone()).await?; + let can_activate = guard + .resolve(&context_id)? + .can_activate(context.clone()) + .await?; if !can_activate { return Err(BootError::Forbidden(format!( "websocket event {} {}", @@ -143,7 +164,10 @@ impl WebSocketGatewayConnection { } } for guard in &handler.guards { - let can_activate = guard.inner().can_activate(context.clone()).await?; + let can_activate = guard + .resolve(&context_id)? + .can_activate(context.clone()) + .await?; if !can_activate { return Err(BootError::Forbidden(format!( "websocket event {} {}", @@ -152,45 +176,54 @@ impl WebSocketGatewayConnection { } } - for interceptor in &self.gateway.interceptors { - interceptor.inner().before(context.clone()).await?; - } - for interceptor in &handler.interceptors { - interceptor.inner().before(context.clone()).await?; - } + let interceptors = self + .gateway + .interceptors + .iter() + .chain(handler.interceptors.iter()) + .map(|interceptor| interceptor.resolve(&context_id)) + .collect::>>()?; - for pipe in &self.gateway.pipes { - message = pipe.inner().transform(message).await?; - } - for pipe in &handler.pipes { - message = pipe.inner().transform(message).await?; - } + let connection = self.clone(); + let replay_handler = handler.clone(); + let replay_message = message.clone(); + let replay_context_id = context_id.clone(); + let terminal = CallHandler::from_fn(move || { + let connection = connection.clone(); + let handler = replay_handler.clone(); + let mut message = replay_message.clone(); + let context_id = replay_context_id.clone(); + async move { + for pipe in &connection.gateway.pipes { + message = pipe.resolve(&context_id)?.transform(message).await?; + } + for pipe in &handler.pipes { + message = pipe.resolve(&context_id)?.transform(message).await?; + } + + if handler.validation_enabled { + for validator in &handler.validators { + message = validator(message, handler.validation_options)?; + } + } - if handler.validation_enabled { - for validator in &handler.validators { - message = validator(message, handler.validation_options)?; + handler.handler.call(connection, message).await } - } + }); - let mut reply = handler.handler.call(self.clone(), message).await?; - for interceptor in handler.interceptors.iter().rev() { - reply = interceptor.inner().after(context.clone(), reply).await?; - } - for interceptor in self.gateway.interceptors.iter().rev() { - reply = interceptor.inner().after(context.clone(), reply).await?; - } - Ok(reply) + run_interceptor_chain(&interceptors, context, terminal).await } async fn handle_error( &self, context: WebSocketContext, error: BootError, + context_id: &ContextId, ) -> Result> { if let Some(handler) = self.gateway.handlers.get(&context.event) { for filter in handler.filters.iter().rev() { + let filter = filter.resolve(context_id)?; if let Some(response) = filter - .inner() .catch(context.clone(), error.clone_for_filter()) .await? { @@ -199,8 +232,8 @@ impl WebSocketGatewayConnection { } } for filter in self.gateway.filters.iter().rev() { + let filter = filter.resolve(context_id)?; if let Some(response) = filter - .inner() .catch(context.clone(), error.clone_for_filter()) .await? { @@ -211,6 +244,23 @@ impl WebSocketGatewayConnection { } } +fn run_interceptor_chain<'a>( + interceptors: &'a [Arc], + context: WebSocketContext, + terminal: CallHandler<'a, Option>, +) -> BoxFuture<'a, Result>> { + let Some((interceptor, remaining)) = interceptors.split_first() else { + return terminal.handle(); + }; + + let next_context = context.clone(); + let next_terminal = terminal.clone(); + let next = CallHandler::from_fn(move || { + run_interceptor_chain(remaining, next_context.clone(), next_terminal.clone()) + }); + interceptor.intercept(context, next) +} + impl WebSocketConnection for WebSocketGatewayConnection { fn request(&self) -> &BootRequest { self.request() diff --git a/src/websocket/gateway.rs b/src/websocket/gateway.rs index 9346368..b33f072 100644 --- a/src/websocket/gateway.rs +++ b/src/websocket/gateway.rs @@ -11,13 +11,13 @@ use super::pipeline::{ }; use super::server::WebSocketGatewayServer; use super::state::{normalize_namespace, normalize_room, WebSocketGatewayState}; -use crate::pipeline::{PipelineComponent, PipelineOverrides}; +use crate::pipeline::{PipelineComponent, PipelineOverrides, ProviderEnhancerComponents}; use crate::routing::path::{ join_paths, match_path_params, match_path_shape, route_shape_key, validate_route_path, }; use crate::{ catch_errors, BootError, BootErrorKind, BootRequest, ExecutionInterceptor, Guard, HttpMethod, - Result, ValidationOptions, WebSocketExceptionFilter, + ModuleRef, Result, ValidationOptions, WebSocketExceptionFilter, }; use serde::Serialize; use serde_json::Value; @@ -41,6 +41,7 @@ pub struct WebSocketGatewayDefinition { pub(crate) filters: Vec>, pub(crate) metadata: BTreeMap, pub(crate) module_name: Option, + pub(crate) module_ref: Option, pub(crate) state: Arc, } @@ -61,6 +62,7 @@ impl WebSocketGatewayDefinition { filters: Vec::new(), metadata: BTreeMap::new(), module_name: None, + module_ref: None, state: Arc::new(WebSocketGatewayState::default()), }) } @@ -418,6 +420,28 @@ impl WebSocketGatewayDefinition { self } + pub(crate) fn with_provider_enhancer_prefix( + mut self, + enhancers: &ProviderEnhancerComponents, + ) -> Self { + let mut pipes = enhancers.websocket_pipes.clone(); + pipes.extend(self.pipes); + self.pipes = pipes; + + let mut guards = enhancers.websocket_guards.clone(); + guards.extend(self.guards); + self.guards = guards; + + let mut interceptors = enhancers.websocket_interceptors.clone(); + interceptors.extend(self.interceptors); + self.interceptors = interceptors; + + let mut filters = enhancers.websocket_filters.clone(); + filters.extend(self.filters); + self.filters = filters; + self + } + pub(crate) fn with_validation_prefix( mut self, validation_enabled: bool, @@ -563,4 +587,16 @@ impl WebSocketGatewayDefinition { self.module_name = Some(module_name.to_string()); self } + + pub(crate) fn with_module_ref(mut self, module_ref: ModuleRef) -> Self { + self.module_ref = Some(module_ref); + self + } + + pub(crate) fn with_default_module_ref(mut self, module_ref: ModuleRef) -> Self { + if self.module_ref.is_none() { + self.module_ref = Some(module_ref); + } + self + } } diff --git a/src/websocket/pipeline.rs b/src/websocket/pipeline.rs index e717f44..d0bef07 100644 --- a/src/websocket/pipeline.rs +++ b/src/websocket/pipeline.rs @@ -1,7 +1,7 @@ use super::context::WebSocketContext; use super::message::WebSocketMessage; use crate::pipeline::PipelineComponent; -use crate::{BoxFuture, ExecutionInterceptor, Guard, Result}; +use crate::{BoxFuture, CallHandler, ExecutionInterceptor, Guard, Result}; use std::future::Future; use std::sync::Arc; @@ -50,6 +50,23 @@ where /// Around-handler hook for WebSocket gateway messages. pub trait WebSocketInterceptor: Send + Sync + 'static { + /// Run around the remaining WebSocket pipeline. + /// + /// Override this method to recover downstream errors, retry the handler, + /// or return a reply without calling `next`. The default implementation + /// preserves the legacy `before` and `after` hook behavior. + fn intercept<'a>( + &'a self, + context: WebSocketContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.before(context.clone()).await?; + let reply = next.handle().await?; + self.after(context, reply).await + }) + } + fn before(&self, _context: WebSocketContext) -> BoxFuture<'static, Result<()>> { Box::pin(async { Ok(()) }) } diff --git a/tests/app_enhancer_protocols.rs b/tests/app_enhancer_protocols.rs new file mode 100644 index 0000000..e5a26c4 --- /dev/null +++ b/tests/app_enhancer_protocols.rs @@ -0,0 +1,448 @@ +use a3s_boot::{ + BootApplication, BootError, BootRequest, BoxFuture, CallHandler, FromModuleRef, HttpMethod, + MessagePatternDefinition, Module, ModuleRef, ProviderDefinition, ProviderDependency, Result, + TransportContext, TransportExceptionFilter, TransportExceptionResponse, TransportGuard, + TransportInterceptor, TransportMessage, TransportPipe, TransportReply, WebSocketContext, + WebSocketExceptionFilter, WebSocketExceptionResponse, WebSocketGatewayConnection, + WebSocketGatewayDefinition, WebSocketGuard, WebSocketInterceptor, WebSocketMessage, + WebSocketPipe, +}; +use serde_json::{json, Value}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +const PROVIDER_ID_FIELD: &str = "providerId"; +const PROVIDER_CONTEXT_ID_FIELD: &str = "providerContextId"; + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ProtocolEvent { + stage: &'static str, + provider_id: usize, + provider_context_id: u64, + execution_context_id: u64, +} + +#[derive(Debug)] +struct ProtocolRequestTrace { + provider_id: usize, + context_id: u64, + events: Arc>>, +} + +impl ProtocolRequestTrace { + fn record(&self, stage: &'static str, request: Option<&BootRequest>) { + let execution_context_id = request + .and_then(BootRequest::context_id) + .map(|context_id| context_id.id()) + .unwrap_or(self.context_id); + self.events.lock().unwrap().push(ProtocolEvent { + stage, + provider_id: self.provider_id, + provider_context_id: self.context_id, + execution_context_id, + }); + } +} + +fn protocol_trace_provider( + calls: &Arc, + events: &Arc>>, +) -> ProviderDefinition { + let calls = Arc::clone(calls); + let events = Arc::clone(events); + ProviderDefinition::request_scoped::(move |module_ref| { + let context_id = module_ref + .context_id() + .ok_or_else(|| { + BootError::Internal("protocol trace was resolved without a ContextId".to_string()) + })? + .id(); + Ok(ProtocolRequestTrace { + provider_id: calls.fetch_add(1, Ordering::SeqCst) + 1, + context_id, + events: Arc::clone(&events), + }) + }) +} + +fn attach_trace(data: &mut Value, trace: &ProtocolRequestTrace) -> Result<()> { + let fields = data.as_object_mut().ok_or_else(|| { + BootError::BadRequest("provider enhancer test payload must be an object".to_string()) + })?; + fields.insert(PROVIDER_ID_FIELD.to_string(), json!(trace.provider_id)); + fields.insert( + PROVIDER_CONTEXT_ID_FIELD.to_string(), + json!(trace.context_id), + ); + Ok(()) +} + +macro_rules! trace_enhancer_providers { + ($($provider:ident),+ $(,)?) => { + $( + #[derive(Debug)] + struct $provider(Arc); + + impl FromModuleRef for $provider { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self(module_ref.get::()?)) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } + } + )+ + }; +} + +trace_enhancer_providers!( + ProviderWebSocketGuard, + ProviderWebSocketPipe, + ProviderWebSocketInterceptor, + ProviderWebSocketFilter, + ProviderTransportGuard, + ProviderTransportPipe, + ProviderTransportInterceptor, + ProviderTransportFilter, +); + +impl WebSocketGuard for ProviderWebSocketGuard { + fn can_activate(&self, context: WebSocketContext) -> BoxFuture<'static, Result> { + self.0.record("guard", Some(&context.request)); + Box::pin(async { Ok(true) }) + } +} + +impl WebSocketPipe for ProviderWebSocketPipe { + fn transform( + &self, + mut message: WebSocketMessage, + ) -> BoxFuture<'static, Result> { + self.0.record("pipe", None); + let result = attach_trace(&mut message.data, &self.0).map(|()| message); + Box::pin(async move { result }) + } +} + +impl WebSocketInterceptor for ProviderWebSocketInterceptor { + fn intercept<'a>( + &'a self, + context: WebSocketContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.0.record("interceptor-before", Some(&context.request)); + let reply = next.handle().await?; + self.0.record("interceptor-after", Some(&context.request)); + Ok(reply) + }) + } +} + +impl WebSocketExceptionFilter for ProviderWebSocketFilter { + fn catch( + &self, + context: WebSocketContext, + error: BootError, + ) -> BoxFuture<'static, Result>> { + self.0.record("filter", Some(&context.request)); + let response = WebSocketExceptionResponse::message(WebSocketMessage::text( + "provider.filtered", + format!( + "ws-filtered:{}:{}:{error}", + self.0.provider_id, self.0.context_id + ), + )); + Box::pin(async move { Ok(Some(response)) }) + } +} + +impl TransportGuard for ProviderTransportGuard { + fn can_activate(&self, context: TransportContext) -> BoxFuture<'static, Result> { + self.0 + .record("guard", Some(&context.execution_context().request)); + Box::pin(async { Ok(true) }) + } +} + +impl TransportPipe for ProviderTransportPipe { + fn transform( + &self, + mut message: TransportMessage, + ) -> BoxFuture<'static, Result> { + self.0.record("pipe", None); + let result = attach_trace(&mut message.data, &self.0).map(|()| message); + Box::pin(async move { result }) + } +} + +impl TransportInterceptor for ProviderTransportInterceptor { + fn intercept<'a>( + &'a self, + context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.0.record( + "interceptor-before", + Some(&context.execution_context().request), + ); + let reply = next.handle().await?; + self.0.record( + "interceptor-after", + Some(&context.execution_context().request), + ); + Ok(reply) + }) + } +} + +impl TransportExceptionFilter for ProviderTransportFilter { + fn catch( + &self, + context: TransportContext, + error: BootError, + ) -> BoxFuture<'static, Result>> { + self.0 + .record("filter", Some(&context.execution_context().request)); + let response = TransportExceptionResponse::reply(TransportReply::text(format!( + "transport-filtered:{}:{}:{error}", + self.0.provider_id, self.0.context_id + ))); + Box::pin(async move { Ok(Some(response)) }) + } +} + +#[derive(Debug)] +struct ProtocolAppEnhancerModule { + calls: Arc, + events: Arc>>, +} + +impl Module for ProtocolAppEnhancerModule { + fn name(&self) -> &'static str { + "protocol-app-enhancers" + } + + fn providers(&self) -> Result> { + Ok(vec![ + protocol_trace_provider(&self.calls, &self.events), + ProviderDefinition::app_websocket_guard::(), + ProviderDefinition::app_websocket_pipe::(), + ProviderDefinition::app_websocket_interceptor::(), + ProviderDefinition::app_websocket_filter::(), + ProviderDefinition::app_transport_guard::(), + ProviderDefinition::app_transport_pipe::(), + ProviderDefinition::app_transport_interceptor::(), + ProviderDefinition::app_transport_filter::(), + ]) + } +} + +#[derive(Debug)] +struct ProtocolTargetModule; + +impl Module for ProtocolTargetModule { + fn name(&self) -> &'static str { + "protocol-app-enhancer-target" + } + + fn gateways(&self, module_ref: &ModuleRef) -> Result> { + ensure_trace_is_private(module_ref)?; + Ok(vec![WebSocketGatewayDefinition::new("/provider-app/ws")? + .subscribe_with_connection( + "trace", + |connection: WebSocketGatewayConnection, message: WebSocketMessage| async move { + let provider_id = message.data_field_as::(PROVIDER_ID_FIELD)?; + let provider_context_id = + message.data_field_as::(PROVIDER_CONTEXT_ID_FIELD)?; + let handler_context_id = connection + .request() + .context_id() + .ok_or_else(|| { + BootError::Internal( + "websocket handler request is missing its ContextId".to_string(), + ) + })? + .id(); + if handler_context_id != provider_context_id { + return Err(BootError::Internal(format!( + "websocket handler ContextId {handler_context_id} differs from provider ContextId {provider_context_id}" + ))); + } + if message.data_field_as::("fail")? { + return Err(BootError::BadRequest(format!( + "websocket boom:{provider_id}:{provider_context_id}" + ))); + } + Ok(WebSocketMessage::text( + "provider.reply", + format!("{provider_id}:{provider_context_id}"), + )) + }, + )?]) + } + + fn message_patterns(&self, module_ref: &ModuleRef) -> Result> { + ensure_trace_is_private(module_ref)?; + Ok(vec![MessagePatternDefinition::request( + "provider.transport", + |message: TransportMessage| async move { + let provider_id = message.data_field_as::(PROVIDER_ID_FIELD)?; + let provider_context_id = + message.data_field_as::(PROVIDER_CONTEXT_ID_FIELD)?; + if message.data_field_as::("fail")? { + return Err(BootError::BadRequest(format!( + "transport boom:{provider_id}:{provider_context_id}" + ))); + } + Ok(TransportReply::text(format!( + "{provider_id}:{provider_context_id}" + ))) + }, + )?]) + } +} + +fn ensure_trace_is_private(module_ref: &ModuleRef) -> Result<()> { + if module_ref.contains_provider::()? { + return Err(BootError::Internal( + "the target module unexpectedly sees the private protocol trace provider".to_string(), + )); + } + Ok(()) +} + +fn take_events(events: &Arc>>) -> Vec { + std::mem::take(&mut *events.lock().unwrap()) +} + +fn assert_dispatch_events( + events: &[ProtocolEvent], + expected_stages: &[&'static str], + expected_provider_id: usize, +) -> u64 { + assert_eq!( + events.iter().map(|event| event.stage).collect::>(), + expected_stages + ); + assert!(events + .iter() + .all(|event| event.provider_id == expected_provider_id)); + let context_id = events.first().unwrap().provider_context_id; + assert!(events.iter().all(|event| { + event.provider_context_id == context_id && event.execution_context_id == context_id + })); + context_id +} + +fn protocol_app( + calls: &Arc, + events: &Arc>>, +) -> BootApplication { + BootApplication::builder() + .import(ProtocolAppEnhancerModule { + calls: Arc::clone(calls), + events: Arc::clone(events), + }) + .import(ProtocolTargetModule) + .build() + .unwrap() +} + +#[tokio::test] +async fn provider_backed_websocket_app_enhancers_use_one_fresh_context_per_message() { + let calls = Arc::new(AtomicUsize::new(0)); + let events = Arc::new(Mutex::new(Vec::new())); + let app = protocol_app(&calls, &events); + let connection = app + .gateway_for("/provider-app/ws") + .unwrap() + .connect(BootRequest::new(HttpMethod::Get, "/provider-app/ws")) + .unwrap(); + + let first_reply = connection + .dispatch(WebSocketMessage::new("trace", json!({ "fail": false }))) + .await + .unwrap() + .unwrap(); + let first_context_id = assert_dispatch_events( + &take_events(&events), + &["guard", "interceptor-before", "pipe", "interceptor-after"], + 1, + ); + assert_eq!( + first_reply, + WebSocketMessage::text("provider.reply", format!("1:{first_context_id}")) + ); + + let handled_error = connection + .dispatch(WebSocketMessage::new("trace", json!({ "fail": true }))) + .await + .unwrap() + .unwrap(); + let second_context_id = assert_dispatch_events( + &take_events(&events), + &["guard", "interceptor-before", "pipe", "filter"], + 2, + ); + assert_eq!( + handled_error, + WebSocketMessage::text( + "provider.filtered", + format!( + "ws-filtered:2:{second_context_id}:bad request: websocket boom:2:{second_context_id}" + ), + ) + ); + assert_ne!(first_context_id, second_context_id); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn provider_backed_transport_app_enhancers_use_one_fresh_context_per_dispatch() { + let calls = Arc::new(AtomicUsize::new(0)); + let events = Arc::new(Mutex::new(Vec::new())); + let app = protocol_app(&calls, &events); + + let first_reply = app + .dispatch_message(TransportMessage::new( + "provider.transport", + json!({ "fail": false }), + )) + .await + .unwrap() + .unwrap(); + let first_context_id = assert_dispatch_events( + &take_events(&events), + &["guard", "interceptor-before", "pipe", "interceptor-after"], + 1, + ); + assert_eq!( + first_reply, + TransportReply::text(format!("1:{first_context_id}")) + ); + + let handled_error = app + .dispatch_message(TransportMessage::new( + "provider.transport", + json!({ "fail": true }), + )) + .await + .unwrap() + .unwrap(); + let second_context_id = assert_dispatch_events( + &take_events(&events), + &["guard", "interceptor-before", "pipe", "filter"], + 2, + ); + assert_eq!( + handled_error, + TransportReply::text(format!( + "transport-filtered:2:{second_context_id}:bad request: transport boom:2:{second_context_id}" + )) + ); + assert_ne!(first_context_id, second_context_id); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} diff --git a/tests/app_enhancer_resolution.rs b/tests/app_enhancer_resolution.rs new file mode 100644 index 0000000..9622e76 --- /dev/null +++ b/tests/app_enhancer_resolution.rs @@ -0,0 +1,449 @@ +use a3s_boot::{ + BootApplication, BootError, BootRequest, BootResponse, BoxFuture, ExceptionFilter, + ExecutionContext, FromModuleRef, Guard, HttpMethod, Module, ModuleRef, ProviderDefinition, + ProviderDependency, ProviderOnModuleInit, Result, RouteDefinition, +}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; + +#[derive(Debug)] +struct NotFoundRequestProbe { + id: usize, +} + +#[derive(Debug)] +struct ProviderNotFoundFilter { + probe: Arc, +} + +impl FromModuleRef for ProviderNotFoundFilter { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + probe: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl ExceptionFilter for ProviderNotFoundFilter { + fn catch( + &self, + context: ExecutionContext, + error: BootError, + ) -> BoxFuture<'static, Result>> { + let has_context = context.request.context_id().is_some(); + let response = + BootResponse::text(format!("not-found:{}:{has_context}:{error}", self.probe.id)); + Box::pin(async move { Ok(Some(response)) }) + } +} + +#[derive(Debug)] +struct ProviderNotFoundFilterModule { + probe_calls: Arc, +} + +impl Module for ProviderNotFoundFilterModule { + fn name(&self) -> &'static str { + "provider-not-found-filter" + } + + fn providers(&self) -> Result> { + let probe_calls = Arc::clone(&self.probe_calls); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + Ok(NotFoundRequestProbe { + id: probe_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + ProviderDefinition::app_filter::(), + ]) + } +} + +#[tokio::test] +async fn provider_app_filter_handles_unmatched_application_404_with_request_context() { + let probe_calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(ProviderNotFoundFilterModule { + probe_calls: Arc::clone(&probe_calls), + }) + .build() + .unwrap(); + + let first = app + .call(BootRequest::new(HttpMethod::Get, "/missing")) + .await + .unwrap(); + let second = app + .call(BootRequest::new(HttpMethod::Get, "/still-missing")) + .await + .unwrap(); + + assert!(first.body_text().unwrap().starts_with("not-found:1:true:")); + assert!(second.body_text().unwrap().starts_with("not-found:2:true:")); + assert_eq!(probe_calls.load(Ordering::SeqCst), 2); +} + +#[derive(Debug)] +struct FailingRequestGuard; + +impl Guard for FailingRequestGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct ResolutionFilterProbe { + id: usize, +} + +#[derive(Debug)] +struct ResolutionFailureFilter { + probe: Arc, +} + +impl FromModuleRef for ResolutionFailureFilter { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + probe: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl ExceptionFilter for ResolutionFailureFilter { + fn catch( + &self, + _context: ExecutionContext, + error: BootError, + ) -> BoxFuture<'static, Result>> { + let response = BootResponse::text(format!("resolved:{}:{error}", self.probe.id)); + Box::pin(async move { Ok(Some(response)) }) + } +} + +#[derive(Debug)] +struct ResolutionFailureModule { + guard_factory_calls: Arc, + probe_calls: Arc, +} + +impl Module for ResolutionFailureModule { + fn name(&self) -> &'static str { + "app-enhancer-resolution-failure" + } + + fn providers(&self) -> Result> { + let guard_factory_calls = Arc::clone(&self.guard_factory_calls); + let probe_calls = Arc::clone(&self.probe_calls); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + guard_factory_calls.fetch_add(1, Ordering::SeqCst); + Err(BootError::Internal( + "request guard construction failed".to_string(), + )) + }) + .with_app_guard::(), + ProviderDefinition::request_scoped::(move |_| { + Ok(ResolutionFilterProbe { + id: probe_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + ProviderDefinition::app_filter::(), + ]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/resolution-failure", + |_| async { Ok(BootResponse::text("unreachable")) }, + )?]) + } +} + +#[tokio::test] +async fn provider_filter_handles_request_scoped_app_enhancer_resolution_failure() { + let guard_factory_calls = Arc::new(AtomicUsize::new(0)); + let probe_calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(ResolutionFailureModule { + guard_factory_calls: Arc::clone(&guard_factory_calls), + probe_calls: Arc::clone(&probe_calls), + }) + .build() + .unwrap(); + + let first = app + .call(BootRequest::new(HttpMethod::Get, "/resolution-failure")) + .await + .unwrap(); + let second = app + .call(BootRequest::new(HttpMethod::Get, "/resolution-failure")) + .await + .unwrap(); + + assert!(first.body_text().unwrap().starts_with("resolved:1:")); + assert!(second.body_text().unwrap().starts_with("resolved:2:")); + assert_eq!(guard_factory_calls.load(Ordering::SeqCst), 2); + assert_eq!(probe_calls.load(Ordering::SeqCst), 2); +} + +#[derive(Debug)] +struct WrongEnhancerValue; + +#[derive(Debug)] +struct ExpectedEnhancerGuard; + +impl Guard for ExpectedEnhancerGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct MismatchedEnhancerMarkerModule; + +impl Module for MismatchedEnhancerMarkerModule { + fn name(&self) -> &'static str { + "mismatched-app-enhancer-marker" + } + + fn providers(&self) -> Result> { + Ok(vec![ProviderDefinition::named_singleton( + "mismatched-app-guard", + WrongEnhancerValue, + ) + .with_app_guard::()]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/mismatched-app-guard", + |_| async { Ok(BootResponse::text("unreachable")) }, + )?]) + } +} + +#[tokio::test] +async fn mismatched_app_enhancer_marker_reports_provider_type_at_resolution() { + let app = BootApplication::builder() + .import(MismatchedEnhancerMarkerModule) + .build() + .unwrap(); + + let error = app + .call(BootRequest::new(HttpMethod::Get, "/mismatched-app-guard")) + .await + .unwrap_err(); + + assert!(matches!( + error, + BootError::ProviderTypeMismatch(provider) + if provider == "mismatched-app-guard" + )); +} + +#[derive(Debug)] +struct AsyncSingletonAppGuard { + activations: Arc, +} + +impl Guard for AsyncSingletonAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.activations.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct AsyncSingletonAppGuardModule { + factory_calls: Arc, + activations: Arc, +} + +impl Module for AsyncSingletonAppGuardModule { + fn name(&self) -> &'static str { + "async-singleton-app-guard" + } + + fn providers(&self) -> Result> { + let factory_calls = Arc::clone(&self.factory_calls); + let activations = Arc::clone(&self.activations); + Ok(vec![ProviderDefinition::async_factory::< + AsyncSingletonAppGuard, + _, + _, + >(move |_| { + let factory_calls = Arc::clone(&factory_calls); + let activations = Arc::clone(&activations); + async move { + factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(AsyncSingletonAppGuard { activations }) + } + }) + .with_dependencies(Vec::new()) + .with_app_guard::()]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get("/async-app-guard", |_| async { + Ok(BootResponse::text("allowed")) + })?]) + } +} + +#[tokio::test] +async fn build_async_supports_static_async_provider_app_enhancers() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let activations = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(AsyncSingletonAppGuardModule { + factory_calls: Arc::clone(&factory_calls), + activations: Arc::clone(&activations), + }) + .build_async() + .await + .unwrap(); + + assert_eq!(factory_calls.load(Ordering::SeqCst), 1); + for _ in 0..2 { + let response = app + .call(BootRequest::new(HttpMethod::Get, "/async-app-guard")) + .await + .unwrap(); + assert_eq!(response.body_text().unwrap(), "allowed"); + } + assert_eq!(factory_calls.load(Ordering::SeqCst), 1); + assert_eq!(activations.load(Ordering::SeqCst), 2); +} + +#[derive(Debug)] +struct ContextualDependency; + +#[derive(Debug)] +struct ContextualAsyncAppGuard { + _dependency: Arc, +} + +impl Guard for ContextualAsyncAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct ContextualAsyncAppGuardModule { + factory_calls: Arc, +} + +impl Module for ContextualAsyncAppGuardModule { + fn name(&self) -> &'static str { + "contextual-async-app-guard" + } + + fn providers(&self) -> Result> { + let factory_calls = Arc::clone(&self.factory_calls); + Ok(vec![ + ProviderDefinition::request_scoped::(|_| { + Ok(ContextualDependency) + }), + ProviderDefinition::async_factory::(move |module_ref| { + factory_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok(ContextualAsyncAppGuard { + _dependency: module_ref.get::()?, + }) + } + }) + .depends_on::() + .with_app_guard::(), + ]) + } +} + +#[tokio::test] +async fn contextual_async_app_enhancer_is_rejected_before_factory_invocation() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let result = BootApplication::builder() + .import(ContextualAsyncAppGuardModule { + factory_calls: Arc::clone(&factory_calls), + }) + .build_async() + .await; + + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message.contains("ContextualAsyncAppGuard") + && message.contains("cannot depend on a request-context provider") + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} + +#[derive(Debug)] +struct ContextualLifecycleAppGuard { + _dependency: Arc, +} + +impl Guard for ContextualLifecycleAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(true) }) + } +} + +impl ProviderOnModuleInit for ContextualLifecycleAppGuard {} + +#[derive(Debug)] +struct ContextualLifecycleAppGuardModule { + factory_calls: Arc, +} + +impl Module for ContextualLifecycleAppGuardModule { + fn name(&self) -> &'static str { + "contextual-lifecycle-app-guard" + } + + fn providers(&self) -> Result> { + let factory_calls = Arc::clone(&self.factory_calls); + Ok(vec![ + ProviderDefinition::request_scoped::(|_| { + Ok(ContextualDependency) + }), + ProviderDefinition::factory::(move |module_ref| { + factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(ContextualLifecycleAppGuard { + _dependency: module_ref.get::()?, + }) + }) + .depends_on::() + .with_on_module_init::() + .with_app_guard::(), + ]) + } +} + +#[test] +fn contextual_lifecycle_app_enhancer_is_rejected_before_factory_invocation() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let result = BootApplication::builder() + .import(ContextualLifecycleAppGuardModule { + factory_calls: Arc::clone(&factory_calls), + }) + .build(); + + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message.contains("ContextualLifecycleAppGuard") + && message.contains("singleton lifecycle hooks") + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} diff --git a/tests/app_enhancer_retries.rs b/tests/app_enhancer_retries.rs new file mode 100644 index 0000000..6dde5ff --- /dev/null +++ b/tests/app_enhancer_retries.rs @@ -0,0 +1,326 @@ +use a3s_boot::{ + BootApplication, BootError, BootRequest, BootResponse, BoxFuture, CallHandler, + ExecutionContext, FromModuleRef, HttpMethod, Interceptor, MessagePatternDefinition, Module, + ModuleRef, ProviderDefinition, ProviderDependency, Result, RouteDefinition, TransportContext, + TransportInterceptor, TransportMessage, TransportPipe, TransportReply, WebSocketContext, + WebSocketGatewayDefinition, WebSocketInterceptor, WebSocketMessage, WebSocketPipe, +}; +use serde_json::json; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug)] +struct RetryTrace { + context_id: u64, + events: Arc>>, +} + +impl RetryTrace { + fn record(&self, stage: &'static str) { + self.events.lock().unwrap().push((stage, self.context_id)); + } + + fn ensure_request_context(&self, request: &BootRequest) -> Result<()> { + let execution_context_id = request + .context_id() + .ok_or_else(|| BootError::Internal("retry request has no ContextId".to_string()))? + .id(); + if execution_context_id != self.context_id { + return Err(BootError::Internal(format!( + "retry execution ContextId {execution_context_id} differs from provider ContextId {}", + self.context_id + ))); + } + Ok(()) + } +} + +macro_rules! retry_enhancer_providers { + ($($provider:ident),+ $(,)?) => { + $( + #[derive(Debug)] + struct $provider(Arc); + + impl FromModuleRef for $provider { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self(module_ref.get::()?)) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } + } + )+ + }; +} + +retry_enhancer_providers!( + HttpRetryInterceptor, + HttpRetryPipe, + WebSocketRetryInterceptor, + WebSocketRetryPipe, + TransportRetryInterceptor, + TransportRetryPipe, +); + +impl Interceptor for HttpRetryInterceptor { + fn intercept<'a>( + &'a self, + context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + self.0.ensure_request_context(&context.request)?; + self.0.record("http-interceptor"); + match next.handle().await { + Err(BootError::ServiceUnavailable(_)) => next.handle().await, + result => result, + } + }) + } +} + +impl a3s_boot::Pipe for HttpRetryPipe { + fn transform(&self, request: BootRequest) -> BoxFuture<'static, Result> { + let result = self.0.ensure_request_context(&request).map(|()| { + self.0.record("http-pipe"); + request + }); + Box::pin(async move { result }) + } +} + +impl WebSocketInterceptor for WebSocketRetryInterceptor { + fn intercept<'a>( + &'a self, + context: WebSocketContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.0.ensure_request_context(&context.request)?; + self.0.record("websocket-interceptor"); + match next.handle().await { + Err(BootError::ServiceUnavailable(_)) => next.handle().await, + result => result, + } + }) + } +} + +impl WebSocketPipe for WebSocketRetryPipe { + fn transform(&self, message: WebSocketMessage) -> BoxFuture<'static, Result> { + self.0.record("websocket-pipe"); + Box::pin(async move { Ok(message) }) + } +} + +impl TransportInterceptor for TransportRetryInterceptor { + fn intercept<'a>( + &'a self, + context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + self.0 + .ensure_request_context(&context.execution_context().request)?; + self.0.record("transport-interceptor"); + match next.handle().await { + Err(BootError::ServiceUnavailable(_)) => next.handle().await, + result => result, + } + }) + } +} + +impl TransportPipe for TransportRetryPipe { + fn transform(&self, message: TransportMessage) -> BoxFuture<'static, Result> { + self.0.record("transport-pipe"); + Box::pin(async move { Ok(message) }) + } +} + +#[derive(Debug)] +struct RetryAppEnhancerModule { + trace_factory_calls: Arc, + events: Arc>>, + http_attempts: Arc, + websocket_attempts: Arc, + transport_attempts: Arc, +} + +impl Module for RetryAppEnhancerModule { + fn name(&self) -> &'static str { + "retry-app-enhancers" + } + + fn providers(&self) -> Result> { + let trace_factory_calls = Arc::clone(&self.trace_factory_calls); + let events = Arc::clone(&self.events); + Ok(vec![ + ProviderDefinition::request_scoped::(move |module_ref| { + let context_id = module_ref + .context_id() + .ok_or_else(|| BootError::Internal("retry trace has no ContextId".to_string()))? + .id(); + trace_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(RetryTrace { + context_id, + events: Arc::clone(&events), + }) + }), + ProviderDefinition::app_interceptor::(), + ProviderDefinition::app_pipe::(), + ProviderDefinition::app_websocket_interceptor::(), + ProviderDefinition::app_websocket_pipe::(), + ProviderDefinition::app_transport_interceptor::(), + ProviderDefinition::app_transport_pipe::(), + ]) + } + + fn routes(&self) -> Result> { + let attempts = Arc::clone(&self.http_attempts); + Ok(vec![RouteDefinition::get("/retry-context", move |_| { + let attempts = Arc::clone(&attempts); + async move { + if attempts.fetch_add(1, Ordering::SeqCst) == 0 { + return Err(BootError::ServiceUnavailable("retry HTTP".to_string())); + } + Ok(BootResponse::text("http-retried")) + } + })?]) + } + + fn gateways(&self, _module_ref: &ModuleRef) -> Result> { + let attempts = Arc::clone(&self.websocket_attempts); + Ok(vec![WebSocketGatewayDefinition::new("/retry-context/ws")? + .subscribe("retry", move |_| { + let attempts = Arc::clone(&attempts); + async move { + if attempts.fetch_add(1, Ordering::SeqCst) == 0 { + return Err(BootError::ServiceUnavailable("retry WebSocket".to_string())); + } + Ok(WebSocketMessage::text("retried", "websocket-retried")) + } + })?]) + } + + fn message_patterns(&self, _module_ref: &ModuleRef) -> Result> { + let attempts = Arc::clone(&self.transport_attempts); + Ok(vec![MessagePatternDefinition::request( + "retry.context", + move |_| { + let attempts = Arc::clone(&attempts); + async move { + if attempts.fetch_add(1, Ordering::SeqCst) == 0 { + return Err(BootError::ServiceUnavailable("retry transport".to_string())); + } + Ok(TransportReply::text("transport-retried")) + } + }, + )?]) + } +} + +struct RetryTestHarness { + app: BootApplication, + trace_factory_calls: Arc, + events: Arc>>, + http_attempts: Arc, + websocket_attempts: Arc, + transport_attempts: Arc, +} + +fn retry_app() -> RetryTestHarness { + let trace_factory_calls = Arc::new(AtomicUsize::new(0)); + let events = Arc::new(Mutex::new(Vec::new())); + let http_attempts = Arc::new(AtomicUsize::new(0)); + let websocket_attempts = Arc::new(AtomicUsize::new(0)); + let transport_attempts = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(RetryAppEnhancerModule { + trace_factory_calls: Arc::clone(&trace_factory_calls), + events: Arc::clone(&events), + http_attempts: Arc::clone(&http_attempts), + websocket_attempts: Arc::clone(&websocket_attempts), + transport_attempts: Arc::clone(&transport_attempts), + }) + .build() + .unwrap(); + RetryTestHarness { + app, + trace_factory_calls, + events, + http_attempts, + websocket_attempts, + transport_attempts, + } +} + +fn assert_retry_context(harness: &RetryTestHarness, stages: &[&'static str]) { + let events = harness.events.lock().unwrap(); + assert_eq!( + events.iter().map(|(stage, _)| *stage).collect::>(), + stages + ); + let context_id = events.first().unwrap().1; + assert!(events.iter().all(|(_, id)| *id == context_id)); + assert_eq!(harness.trace_factory_calls.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn http_app_enhancer_retry_reuses_the_invocation_context() { + let harness = retry_app(); + let response = harness + .app + .call(BootRequest::new(HttpMethod::Get, "/retry-context")) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "http-retried"); + assert_eq!(harness.http_attempts.load(Ordering::SeqCst), 2); + assert_retry_context(&harness, &["http-interceptor", "http-pipe", "http-pipe"]); +} + +#[tokio::test] +async fn websocket_app_enhancer_retry_reuses_the_message_context() { + let harness = retry_app(); + let connection = harness + .app + .gateway_for("/retry-context/ws") + .unwrap() + .connect(BootRequest::new(HttpMethod::Get, "/retry-context/ws")) + .unwrap(); + let reply = connection + .dispatch(WebSocketMessage::new("retry", json!({}))) + .await + .unwrap() + .unwrap(); + + assert_eq!( + reply, + WebSocketMessage::text("retried", "websocket-retried") + ); + assert_eq!(harness.websocket_attempts.load(Ordering::SeqCst), 2); + assert_retry_context( + &harness, + &["websocket-interceptor", "websocket-pipe", "websocket-pipe"], + ); +} + +#[tokio::test] +async fn transport_app_enhancer_retry_reuses_the_message_context() { + let harness = retry_app(); + let reply = harness + .app + .dispatch_message(TransportMessage::new("retry.context", json!({}))) + .await + .unwrap() + .unwrap(); + + assert_eq!(reply, TransportReply::text("transport-retried")); + assert_eq!(harness.transport_attempts.load(Ordering::SeqCst), 2); + assert_retry_context( + &harness, + &["transport-interceptor", "transport-pipe", "transport-pipe"], + ); +} diff --git a/tests/app_enhancer_topology.rs b/tests/app_enhancer_topology.rs new file mode 100644 index 0000000..0c86640 --- /dev/null +++ b/tests/app_enhancer_topology.rs @@ -0,0 +1,251 @@ +use a3s_boot::{ + BootApplication, BootRequest, BootResponse, BoxFuture, ExecutionContext, FromModuleRef, Guard, + HttpMethod, Module, ModuleRef, OpenApiInfo, ProviderDefinition, ProviderDependency, + ProviderToken, Result, RouteDefinition, +}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug)] +struct AliasedRequestGuard { + activations: Arc, +} + +impl Guard for AliasedRequestGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.activations.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct AliasedRequestGuardModule { + factory_calls: Arc, + activations: Arc, +} + +impl Module for AliasedRequestGuardModule { + fn name(&self) -> &'static str { + "aliased-request-app-guard" + } + + fn providers(&self) -> Result> { + let factory_calls = Arc::clone(&self.factory_calls); + let activations = Arc::clone(&self.activations); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(AliasedRequestGuard { + activations: Arc::clone(&activations), + }) + }), + ProviderDefinition::named_alias( + "application-guard-alias", + ProviderToken::of::(), + ) + .with_app_guard::(), + ]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/aliased-app-guard", + |_| async { Ok(BootResponse::text("aliased")) }, + )?]) + } +} + +#[tokio::test] +async fn alias_marker_preserves_request_scoped_app_enhancer_resolution() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let activations = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(AliasedRequestGuardModule { + factory_calls: Arc::clone(&factory_calls), + activations: Arc::clone(&activations), + }) + .build() + .unwrap(); + + for _ in 0..2 { + let response = app + .call(BootRequest::new(HttpMethod::Get, "/aliased-app-guard")) + .await + .unwrap(); + assert_eq!(response.body_text().unwrap(), "aliased"); + } + + assert_eq!(factory_calls.load(Ordering::SeqCst), 2); + assert_eq!(activations.load(Ordering::SeqCst), 2); +} + +#[derive(Debug)] +struct ValueMarkedAppGuard { + activations: Arc, +} + +impl Guard for ValueMarkedAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.activations.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct ValueMarkedAppGuardModule { + activations: Arc, +} + +impl Module for ValueMarkedAppGuardModule { + fn name(&self) -> &'static str { + "value-marked-app-guard" + } + + fn providers(&self) -> Result> { + Ok(vec![ProviderDefinition::singleton(ValueMarkedAppGuard { + activations: Arc::clone(&self.activations), + }) + .with_app_guard::()]) + } +} + +#[tokio::test] +async fn value_marked_app_enhancer_applies_to_framework_provided_routes() { + let activations = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .serve_openapi( + "/provider-openapi.json", + OpenApiInfo::new("Provider Enhancers", "1.0.0"), + ) + .import(ValueMarkedAppGuardModule { + activations: Arc::clone(&activations), + }) + .build() + .unwrap(); + + let response = app + .call(BootRequest::new(HttpMethod::Get, "/provider-openapi.json")) + .await + .unwrap(); + + assert_eq!(response.status, 200); + assert_eq!(activations.load(Ordering::SeqCst), 1); +} + +#[derive(Debug)] +struct PrivateEnhancerDependency { + events: Arc>>, +} + +#[derive(Debug)] +struct PrivateDependencyAppGuard { + dependency: Arc, +} + +impl FromModuleRef for PrivateDependencyAppGuard { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + dependency: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ + ProviderDependency::typed::(), + ]) + } +} + +impl Guard for PrivateDependencyAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.dependency.events.lock().unwrap().push("provider"); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct PrivateEnhancerOwnerModule { + events: Arc>>, +} + +impl Module for PrivateEnhancerOwnerModule { + fn name(&self) -> &'static str { + "private-enhancer-owner" + } + + fn providers(&self) -> Result> { + Ok(vec![ + ProviderDefinition::singleton(PrivateEnhancerDependency { + events: Arc::clone(&self.events), + }), + ProviderDefinition::app_guard::(), + ]) + } +} + +#[derive(Debug)] +struct IsolatedTargetModule { + events: Arc>>, +} + +impl Module for IsolatedTargetModule { + fn name(&self) -> &'static str { + "isolated-app-enhancer-target" + } + + fn routes(&self) -> Result> { + let local_events = Arc::clone(&self.events); + let route = RouteDefinition::get("/isolated-target", |request: BootRequest| async move { + let visibility = if request + .get_optional::()? + .is_some() + { + "visible" + } else { + "hidden" + }; + Ok(BootResponse::text(visibility)) + })? + .with_guard(move |_| { + let local_events = Arc::clone(&local_events); + async move { + local_events.lock().unwrap().push("local"); + Ok(true) + } + }); + Ok(vec![route]) + } +} + +#[tokio::test] +async fn provider_enhancer_uses_declaring_module_without_widening_target_visibility() { + let events = Arc::new(Mutex::new(Vec::new())); + let builder_events = Arc::clone(&events); + let app = BootApplication::builder() + .use_global_guard(move |_| { + let builder_events = Arc::clone(&builder_events); + async move { + builder_events.lock().unwrap().push("builder"); + Ok(true) + } + }) + .import(PrivateEnhancerOwnerModule { + events: Arc::clone(&events), + }) + .import(IsolatedTargetModule { + events: Arc::clone(&events), + }) + .build() + .unwrap(); + + let response = app + .call(BootRequest::new(HttpMethod::Get, "/isolated-target")) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "hidden"); + assert_eq!( + events.lock().unwrap().as_slice(), + ["builder", "provider", "local"] + ); +} diff --git a/tests/app_enhancers.rs b/tests/app_enhancers.rs new file mode 100644 index 0000000..d2f60dd --- /dev/null +++ b/tests/app_enhancers.rs @@ -0,0 +1,702 @@ +use a3s_boot::{ + BootApplication, BootError, BootRequest, BootResponse, BoxFuture, CallHandler, ExceptionFilter, + ExecutionContext, FromModuleRef, Guard, HttpMethod, Interceptor, Module, ModuleRef, Pipe, + ProviderDefinition, ProviderDependency, Result, RouteDefinition, TestingModule, +}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug)] +struct RequestTrace { + id: usize, + events: Arc>>, +} + +impl RequestTrace { + fn record(&self, stage: &str) { + self.events + .lock() + .unwrap() + .push(format!("{stage}:{}", self.id)); + } +} + +#[derive(Debug)] +struct ProviderAppGuard { + trace: Arc, +} + +impl FromModuleRef for ProviderAppGuard { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + trace: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl Guard for ProviderAppGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.trace.record("guard"); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct ProviderAppPipe { + trace: Arc, +} + +impl FromModuleRef for ProviderAppPipe { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + trace: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl Pipe for ProviderAppPipe { + fn transform(&self, request: BootRequest) -> BoxFuture<'static, Result> { + self.trace.record("pipe"); + Box::pin(async move { Ok(request) }) + } +} + +#[derive(Debug)] +struct ProviderAppInterceptor { + trace: Arc, +} + +impl FromModuleRef for ProviderAppInterceptor { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + trace: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl Interceptor for ProviderAppInterceptor { + fn intercept<'a>( + &'a self, + context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + self.trace.record("interceptor-before"); + let response = next.handle().await?; + self.trace.record("interceptor-after"); + let _ = context; + Ok(response) + }) + } +} + +#[derive(Debug)] +struct ProviderAppFilter { + trace: Arc, +} + +impl FromModuleRef for ProviderAppFilter { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + trace: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl ExceptionFilter for ProviderAppFilter { + fn catch( + &self, + _context: ExecutionContext, + error: BootError, + ) -> BoxFuture<'static, Result>> { + self.trace.record("filter"); + let response = BootResponse::text(format!("filtered:{}:{error}", self.trace.id)); + Box::pin(async move { Ok(Some(response)) }) + } +} + +#[derive(Debug)] +struct AppEnhancerModule { + calls: Arc, + events: Arc>>, +} + +impl Module for AppEnhancerModule { + fn name(&self) -> &'static str { + "app-enhancers" + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + let events = Arc::clone(&self.events); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + Ok(RequestTrace { + id: calls.fetch_add(1, Ordering::SeqCst) + 1, + events: Arc::clone(&events), + }) + }), + ProviderDefinition::app_guard::(), + ProviderDefinition::app_interceptor::(), + ProviderDefinition::app_pipe::(), + ProviderDefinition::app_filter::(), + ]) + } + + fn routes(&self) -> Result> { + Ok(vec![ + RouteDefinition::get("/module", |request: BootRequest| async move { + let trace = request.get::()?; + trace.record("handler"); + Ok(BootResponse::text(format!("module:{}", trace.id))) + })?, + RouteDefinition::get("/boom", |request: BootRequest| async move { + let trace = request.get::()?; + trace.record("handler"); + Err(BootError::Internal("boom".to_string())) + })?, + ]) + } +} + +#[derive(Debug)] +struct SiblingModule; + +impl Module for SiblingModule { + fn name(&self) -> &'static str { + "app-enhancer-sibling" + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get("/sibling", |_| async { + Ok(BootResponse::text("sibling")) + })?]) + } +} + +fn take_events(events: &Arc>>) -> Vec { + std::mem::take(&mut *events.lock().unwrap()) +} + +#[tokio::test] +async fn provider_backed_http_app_enhancers_share_one_request_context() { + let calls = Arc::new(AtomicUsize::new(0)); + let events = Arc::new(Mutex::new(Vec::new())); + let direct = RouteDefinition::get("/direct", |request: BootRequest| async move { + let trace = request.get::()?; + trace.record("handler"); + Ok(BootResponse::text(format!("direct:{}", trace.id))) + }) + .unwrap(); + let app = BootApplication::builder() + .route(direct) + .import(AppEnhancerModule { + calls: Arc::clone(&calls), + events: Arc::clone(&events), + }) + .import(SiblingModule) + .build() + .unwrap(); + + let module = app + .call(BootRequest::new(HttpMethod::Get, "/module")) + .await + .unwrap(); + assert_eq!(module.body_text().unwrap(), "module:1"); + assert_eq!( + take_events(&events), + [ + "guard:1", + "interceptor-before:1", + "pipe:1", + "handler:1", + "interceptor-after:1", + ] + ); + + let direct = app + .call(BootRequest::new(HttpMethod::Get, "/direct")) + .await + .unwrap(); + assert_eq!(direct.body_text().unwrap(), "direct:2"); + assert_eq!( + take_events(&events), + [ + "guard:2", + "interceptor-before:2", + "pipe:2", + "handler:2", + "interceptor-after:2", + ] + ); + + let sibling = app + .call(BootRequest::new(HttpMethod::Get, "/sibling")) + .await + .unwrap(); + assert_eq!(sibling.body_text().unwrap(), "sibling"); + assert_eq!( + take_events(&events), + [ + "guard:3", + "interceptor-before:3", + "pipe:3", + "interceptor-after:3", + ] + ); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} + +#[tokio::test] +async fn provider_backed_app_filters_handle_pipeline_and_early_route_errors() { + let calls = Arc::new(AtomicUsize::new(0)); + let events = Arc::new(Mutex::new(Vec::new())); + let app = BootApplication::builder() + .import(AppEnhancerModule { + calls: Arc::clone(&calls), + events: Arc::clone(&events), + }) + .build() + .unwrap(); + + let handler_error = app + .call(BootRequest::new(HttpMethod::Get, "/boom")) + .await + .unwrap(); + assert!(handler_error + .body_text() + .unwrap() + .starts_with("filtered:1:")); + assert_eq!( + take_events(&events), + [ + "guard:1", + "interceptor-before:1", + "pipe:1", + "handler:1", + "filter:1", + ] + ); + + let method_error = app + .call(BootRequest::new(HttpMethod::Post, "/module")) + .await + .unwrap(); + assert!(method_error.body_text().unwrap().starts_with("filtered:2:")); + assert_eq!(take_events(&events), ["filter:2"]); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[derive(Debug)] +struct EnhancerOrderLog(Arc>>); + +#[derive(Debug)] +struct FirstDeclaredProviderGuard { + log: Arc, +} + +impl FromModuleRef for FirstDeclaredProviderGuard { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + log: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl Guard for FirstDeclaredProviderGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.log.0.lock().unwrap().push("provider-first"); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct SecondDeclaredProviderGuard { + log: Arc, +} + +impl FromModuleRef for SecondDeclaredProviderGuard { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + log: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +impl Guard for SecondDeclaredProviderGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.log.0.lock().unwrap().push("provider-second"); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct EnhancerDeclarationOrderModule { + log: Arc>>, +} + +impl Module for EnhancerDeclarationOrderModule { + fn name(&self) -> &'static str { + "enhancer-declaration-order" + } + + fn providers(&self) -> Result> { + Ok(vec![ + ProviderDefinition::singleton(EnhancerOrderLog(Arc::clone(&self.log))), + ProviderDefinition::app_guard::(), + ProviderDefinition::app_guard::(), + ]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get("/enhancer-order", |_| async { + Ok(BootResponse::text("ordered")) + })?]) + } +} + +#[tokio::test] +async fn provider_app_enhancers_follow_builder_globals_in_declaration_order() { + let log = Arc::new(Mutex::new(Vec::new())); + let builder_log = Arc::clone(&log); + let app = BootApplication::builder() + .use_global_guard(move |_| { + let builder_log = Arc::clone(&builder_log); + async move { + builder_log.lock().unwrap().push("builder-global"); + Ok(true) + } + }) + .import(EnhancerDeclarationOrderModule { + log: Arc::clone(&log), + }) + .build() + .unwrap(); + + let response = app + .call(BootRequest::new(HttpMethod::Get, "/enhancer-order")) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "ordered"); + assert_eq!( + log.lock().unwrap().as_slice(), + ["builder-global", "provider-first", "provider-second"] + ); +} + +#[derive(Debug)] +struct NamedFactoryMarkerGuard { + activations: Arc, +} + +impl Guard for NamedFactoryMarkerGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.activations.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct NamedFactoryMarkerModule { + factory_calls: Arc, + activations: Arc, +} + +impl Module for NamedFactoryMarkerModule { + fn name(&self) -> &'static str { + "named-factory-marker" + } + + fn providers(&self) -> Result> { + let factory_calls = Arc::clone(&self.factory_calls); + let activations = Arc::clone(&self.activations); + Ok(vec![ProviderDefinition::named_factory::< + NamedFactoryMarkerGuard, + _, + >("named-factory-app-guard", move |_| { + factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(NamedFactoryMarkerGuard { + activations: Arc::clone(&activations), + }) + }) + .with_app_guard::()]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/named-factory-marker", + |_| async { Ok(BootResponse::text("named")) }, + )?]) + } +} + +#[tokio::test] +async fn named_custom_factory_marker_resolves_provider_backed_app_guard() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let activations = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(NamedFactoryMarkerModule { + factory_calls: Arc::clone(&factory_calls), + activations: Arc::clone(&activations), + }) + .build() + .unwrap(); + + assert_eq!(factory_calls.load(Ordering::SeqCst), 1); + + let response = app + .call(BootRequest::new(HttpMethod::Get, "/named-factory-marker")) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "named"); + assert_eq!(factory_calls.load(Ordering::SeqCst), 1); + assert_eq!(activations.load(Ordering::SeqCst), 1); +} + +#[derive(Debug)] +struct TestingProviderBackedDenyGuard; + +impl FromModuleRef for TestingProviderBackedDenyGuard { + fn from_module_ref(_module_ref: &ModuleRef) -> Result { + Ok(Self) + } + + fn provider_dependencies() -> Option> { + Some(Vec::new()) + } +} + +impl Guard for TestingProviderBackedDenyGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(false) }) + } +} + +#[derive(Debug)] +struct TestingProviderBackedAllowGuard { + calls: Arc, +} + +impl Guard for TestingProviderBackedAllowGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.calls.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct TestingProviderBackedGuardModule; + +impl Module for TestingProviderBackedGuardModule { + fn name(&self) -> &'static str { + "testing-provider-backed-guard" + } + + fn providers(&self) -> Result> { + Ok(vec![ProviderDefinition::app_guard::< + TestingProviderBackedDenyGuard, + >()]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/testing-provider-guard", + |_| async { Ok(BootResponse::text("allowed")) }, + )?]) + } +} + +#[tokio::test] +async fn testing_module_override_guard_replaces_provider_backed_component() { + let replacement_calls = Arc::new(AtomicUsize::new(0)); + let testing = TestingModule::builder() + .import(TestingProviderBackedGuardModule) + .override_guard::(TestingProviderBackedAllowGuard { + calls: Arc::clone(&replacement_calls), + }) + .compile() + .unwrap(); + + let response = testing + .call(BootRequest::new(HttpMethod::Get, "/testing-provider-guard")) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "allowed"); + assert_eq!(replacement_calls.load(Ordering::SeqCst), 1); +} + +#[derive(Debug)] +struct MarkerPreservingOverrideGuard { + allow: bool, + activations: Arc, +} + +impl Guard for MarkerPreservingOverrideGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + self.activations.fetch_add(1, Ordering::SeqCst); + let allow = self.allow; + Box::pin(async move { Ok(allow) }) + } +} + +#[derive(Debug)] +struct MarkerPreservingOverrideModule { + original_factory_calls: Arc, + original_activations: Arc, +} + +impl Module for MarkerPreservingOverrideModule { + fn name(&self) -> &'static str { + "marker-preserving-provider-override" + } + + fn providers(&self) -> Result> { + let original_factory_calls = Arc::clone(&self.original_factory_calls); + let original_activations = Arc::clone(&self.original_activations); + Ok(vec![ProviderDefinition::factory::< + MarkerPreservingOverrideGuard, + _, + >(move |_| { + original_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(MarkerPreservingOverrideGuard { + allow: false, + activations: Arc::clone(&original_activations), + }) + }) + .with_app_guard::()]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/provider-marker-override", + |_| async { Ok(BootResponse::text("overridden")) }, + )?]) + } +} + +#[tokio::test] +async fn override_provider_retains_original_app_enhancer_marker() { + let original_factory_calls = Arc::new(AtomicUsize::new(0)); + let original_activations = Arc::new(AtomicUsize::new(0)); + let replacement_activations = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(MarkerPreservingOverrideModule { + original_factory_calls: Arc::clone(&original_factory_calls), + original_activations: Arc::clone(&original_activations), + }) + .override_provider(ProviderDefinition::singleton( + MarkerPreservingOverrideGuard { + allow: true, + activations: Arc::clone(&replacement_activations), + }, + )) + .build() + .unwrap(); + + assert_eq!(original_factory_calls.load(Ordering::SeqCst), 0); + + let response = app + .call(BootRequest::new( + HttpMethod::Get, + "/provider-marker-override", + )) + .await + .unwrap(); + + assert_eq!(response.body_text().unwrap(), "overridden"); + assert_eq!(original_factory_calls.load(Ordering::SeqCst), 0); + assert_eq!(original_activations.load(Ordering::SeqCst), 0); + assert_eq!(replacement_activations.load(Ordering::SeqCst), 1); +} + +#[derive(Debug)] +struct LazyEnhancerPreflightProbe; + +#[derive(Debug)] +struct LazyRejectedProviderGuard; + +impl Guard for LazyRejectedProviderGuard { + fn can_activate(&self, _context: ExecutionContext) -> BoxFuture<'static, Result> { + Box::pin(async { Ok(true) }) + } +} + +#[derive(Debug)] +struct LazyRejectedEnhancerModule { + probe_factory_calls: Arc, + enhancer_factory_calls: Arc, +} + +impl Module for LazyRejectedEnhancerModule { + fn name(&self) -> &'static str { + "lazy-app-enhancer-preflight" + } + + fn providers(&self) -> Result> { + let probe_factory_calls = Arc::clone(&self.probe_factory_calls); + let enhancer_factory_calls = Arc::clone(&self.enhancer_factory_calls); + Ok(vec![ + ProviderDefinition::factory::(move |_| { + probe_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(LazyEnhancerPreflightProbe) + }), + ProviderDefinition::factory::(move |_| { + enhancer_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(LazyRejectedProviderGuard) + }) + .with_app_guard::(), + ]) + } +} + +#[test] +fn lazy_module_loader_rejects_app_enhancer_before_provider_factories_run() { + let probe_factory_calls = Arc::new(AtomicUsize::new(0)); + let enhancer_factory_calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder().build().unwrap(); + let loader = app.lazy_module_loader().unwrap(); + + let error = loader + .load(LazyRejectedEnhancerModule { + probe_factory_calls: Arc::clone(&probe_factory_calls), + enhancer_factory_calls: Arc::clone(&enhancer_factory_calls), + }) + .unwrap_err(); + + match error { + BootError::Internal(message) => { + assert!(message.contains("lazy-app-enhancer-preflight")); + assert!(message.contains("APP_*")); + } + other => panic!("expected a lazy application-enhancer error, got {other}"), + } + assert_eq!(probe_factory_calls.load(Ordering::SeqCst), 0); + assert_eq!(enhancer_factory_calls.load(Ordering::SeqCst), 0); +} diff --git a/tests/lazy_modules.rs b/tests/lazy_modules.rs index 7e77e11..3264c71 100644 --- a/tests/lazy_modules.rs +++ b/tests/lazy_modules.rs @@ -169,6 +169,193 @@ struct LazyLifecycleModule { controller_calls: Arc, } +#[derive(Debug)] +struct LazyForwardRequest { + id: usize, +} + +#[derive(Debug)] +struct LazyForwardContext { + request: Arc, +} + +#[derive(Debug)] +struct LazyForwardContextRoot { + calls: Arc, +} + +impl Module for LazyForwardContextRoot { + fn name(&self) -> &'static str { + "lazy-forward-context-root" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(LazyForwardContextFeature { + calls: Arc::clone(&self.calls), + })] + } + + fn providers(&self) -> Result> { + Ok(vec![ProviderDefinition::factory::( + |module_ref| { + Ok(LazyForwardContext { + request: module_ref.get::()?, + }) + }, + ) + .depends_on::()]) + } +} + +#[derive(Debug)] +struct LazyForwardContextFeature { + calls: Arc, +} + +impl Module for LazyForwardContextFeature { + fn name(&self) -> &'static str { + "lazy-forward-context-feature" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(LazyForwardContextRoot { + calls: Arc::clone(&self.calls), + })] + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ProviderDefinition::request_scoped::< + LazyForwardRequest, + _, + >(move |_| { + Ok(LazyForwardRequest { + id: calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + })]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } +} + +#[derive(Debug)] +struct LazyForwardAsyncConfig { + value: &'static str, +} + +#[derive(Debug)] +struct LazyForwardAsyncService { + config: Arc, +} + +#[derive(Debug)] +struct LazyForwardAsyncRoot { + log: Arc>>, +} + +impl Module for LazyForwardAsyncRoot { + fn name(&self) -> &'static str { + "lazy-forward-async-root" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(LazyForwardAsyncFeature { + log: Arc::clone(&self.log), + })] + } + + fn providers(&self) -> Result> { + let log = Arc::clone(&self.log); + Ok(vec![ProviderDefinition::async_factory::< + LazyForwardAsyncConfig, + _, + _, + >(move |_| { + let log = Arc::clone(&log); + async move { + log.lock().unwrap().push("config"); + Ok(LazyForwardAsyncConfig { value: "ready" }) + } + })]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } +} + +#[derive(Debug)] +struct LazyForwardAsyncFeature { + log: Arc>>, +} + +impl Module for LazyForwardAsyncFeature { + fn name(&self) -> &'static str { + "lazy-forward-async-feature" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(LazyForwardAsyncRoot { + log: Arc::clone(&self.log), + })] + } + + fn providers(&self) -> Result> { + let log = Arc::clone(&self.log); + Ok(vec![ProviderDefinition::async_factory::< + LazyForwardAsyncService, + _, + _, + >(move |module_ref| { + let log = Arc::clone(&log); + async move { + let config = module_ref.get::()?; + log.lock().unwrap().push("service"); + Ok(LazyForwardAsyncService { config }) + } + }) + .depends_on::()]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } +} + +#[derive(Debug)] +struct LateLazyGlobalValue; + +#[derive(Debug)] +struct LateLazyGlobalModule { + calls: Arc, +} + +impl Module for LateLazyGlobalModule { + fn name(&self) -> &'static str { + "late-lazy-global" + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ProviderDefinition::factory::( + move |_| { + calls.fetch_add(1, Ordering::SeqCst); + Ok(LateLazyGlobalValue) + }, + )]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } + + fn is_global(&self) -> bool { + true + } +} + impl Module for LazyLifecycleModule { fn name(&self) -> &'static str { "lazy-lifecycle" @@ -345,3 +532,99 @@ fn lazy_modules_do_not_register_controllers_or_lifecycle_hooks() { assert_eq!(provider_calls.load(Ordering::SeqCst), 0); assert_eq!(controller_calls.load(Ordering::SeqCst), 0); } + +#[test] +fn lazy_forward_graph_is_complete_before_contextual_scope_is_planned() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder().build().unwrap(); + let loaded = app + .lazy_module_loader() + .unwrap() + .load(LazyForwardContextRoot { + calls: Arc::clone(&calls), + }) + .unwrap(); + + assert!(loaded + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert!( + matches!(loaded.get::(), Err(BootError::Internal(message)) if message.contains("requires an active request scope")) + ); + + let first = loaded.module_ref().resolve::().unwrap(); + let second = loaded.module_ref().resolve::().unwrap(); + + assert_eq!(first.request.id, 1); + assert_eq!(second.request.id, 2); + assert!(!Arc::ptr_eq(&first, &second)); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn lazy_async_forward_dependencies_seed_after_the_full_graph_is_registered() { + let log = Arc::new(std::sync::Mutex::new(Vec::new())); + let app = BootApplication::builder().build().unwrap(); + let loaded = app + .lazy_module_loader() + .unwrap() + .load_async(LazyForwardAsyncRoot { + log: Arc::clone(&log), + }) + .await + .unwrap(); + + let config = loaded.get::().unwrap(); + let service = loaded.get::().unwrap(); + + assert_eq!(config.value, "ready"); + assert!(Arc::ptr_eq(&config, &service.config)); + assert_eq!(log.lock().unwrap().as_slice(), ["config", "service"]); +} + +#[test] +fn lazy_loader_rejects_late_global_modules_before_factories_run() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder().build().unwrap(); + let result = app + .lazy_module_loader() + .unwrap() + .load(LateLazyGlobalModule { + calls: Arc::clone(&calls), + }); + + assert!( + matches!(result, Err(BootError::Internal(message)) if message.contains("register global modules eagerly")) + ); + assert_eq!(calls.load(Ordering::SeqCst), 0); + assert!(matches!( + app.get::(), + Err(BootError::MissingProvider(_)) + )); +} + +#[test] +fn lazy_loader_can_reuse_an_eagerly_registered_global_module() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(LateLazyGlobalModule { + calls: Arc::clone(&calls), + }) + .build() + .unwrap(); + let eager = app.get::().unwrap(); + + let loaded = app + .lazy_module_loader() + .unwrap() + .load(LateLazyGlobalModule { + calls: Arc::clone(&calls), + }) + .unwrap() + .get::() + .unwrap(); + + assert!(Arc::ptr_eq(&eager, &loaded)); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} diff --git a/tests/macros.rs b/tests/macros.rs index 4f6319d..5bd4643 100644 --- a/tests/macros.rs +++ b/tests/macros.rs @@ -2,18 +2,20 @@ use std::collections::BTreeMap; use std::str::FromStr; -use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; use a3s_boot::{ controller, injectable, ApiVersioning, BootApplication, BootError, BootErrorKind, BootRequest, - BootResponse, BoxFuture, ControllerDefinition, ExceptionFilter, ExecutionContext, Guard, - Interceptor, Module, ModuleRef, OpenApiInfo, ParseArrayPipe, ParseBoolPipe, ParseEnumPipe, - ParseFloatPipe, ParseIntPipe, ParseUuidPipe, Pipe, ProviderDefinition, ProviderRef, - ProviderToken, Result, SseEvent, SseStream, StringTemplateViewEngine, TransportContext, - TransportExceptionFilter, TransportExceptionResponse, TransportMessage, TransportReply, - UuidVersion, Validate, ViewModule, WebSocketContext, WebSocketExceptionFilter, - WebSocketExceptionResponse, WebSocketGatewayConnection, WebSocketGatewayInitContext, - WebSocketGatewayServer, WebSocketMessage, + BootResponse, BoxFuture, CallHandler, ControllerDefinition, ExceptionFilter, ExecutionContext, + FromModuleRef, Guard, Interceptor, Module, ModuleRef, OpenApiInfo, ParseArrayPipe, + ParseBoolPipe, ParseEnumPipe, ParseFloatPipe, ParseIntPipe, ParseUuidPipe, Pipe, + ProviderDefinition, ProviderRef, ProviderScope, ProviderToken, Result, SseEvent, SseStream, + StringTemplateViewEngine, TransportContext, TransportExceptionFilter, + TransportExceptionResponse, TransportMessage, TransportReply, UuidVersion, Validate, + ViewModule, WebSocketContext, WebSocketExceptionFilter, WebSocketExceptionResponse, + WebSocketGatewayConnection, WebSocketGatewayInitContext, WebSocketGatewayServer, + WebSocketMessage, }; use futures_util::StreamExt; use serde::{Deserialize, Serialize}; @@ -720,6 +722,26 @@ struct MacroCatDetailsDto { #[derive(Debug)] struct MacroCatsModule; +#[a3s_boot::module( + name = "macro-provider-flavors", + imports = [ + ViewModule::new( + "macro-provider-views", + StringTemplateViewEngine::new() + .with_template("cats/card", "

{{ id }}:{{ name }}
"), + ) + ], + providers = [ + MacroCatsService, + MacroCatsService.into_named_provider("readonly-cats"), + MacroAutoCatsReader, + MacroCatsController::request_scoped_provider(), + ], + controllers = [MacroCatsController], +)] +#[derive(Debug)] +struct MacroProviderFlavorModule; + #[derive(Debug)] struct MacroHiddenOpenApiController; @@ -748,6 +770,7 @@ impl Module for MacroHiddenOpenApiModule { } } +#[injectable] #[derive(Debug)] struct MacroValidationController; @@ -832,23 +855,7 @@ fn macro_default_kind() -> String { "cat".to_string() } -#[derive(Debug)] -struct MacroValidationModule; - -impl Module for MacroValidationModule { - fn name(&self) -> &'static str { - "macro-validation" - } - - fn controllers(&self, _module_ref: &ModuleRef) -> Result> { - Ok(vec![ - Arc::new(MacroValidationController).controller()?, - Arc::new(MacroRouteValidationController).controller()?, - Arc::new(MacroWhitelistValidationController).controller()?, - ]) - } -} - +#[injectable] #[derive(Debug)] struct MacroRouteValidationController; @@ -864,6 +871,7 @@ impl MacroRouteValidationController { } } +#[injectable] #[derive(Debug)] struct MacroWhitelistValidationController; @@ -899,6 +907,22 @@ impl MacroWhitelistValidationController { } } +#[a3s_boot::module( + name = "macro-validation", + providers = [ + MacroValidationController::request_scoped_provider(), + MacroRouteValidationController::request_scoped_provider(), + MacroWhitelistValidationController::request_scoped_provider(), + ], + controllers = [ + MacroValidationController, + MacroRouteValidationController, + MacroWhitelistValidationController, + ], +)] +#[derive(Debug)] +struct MacroValidationModule; + struct MacroControllerHeaderInterceptor; impl Interceptor for MacroControllerHeaderInterceptor { @@ -1041,7 +1065,7 @@ impl MacroPipelineController { #[a3s_boot::module( name = "macro-pipeline", - providers = [MacroPipelineController], + providers = [MacroPipelineController::request_scoped_provider()], controllers = [MacroPipelineController], )] #[derive(Debug)] @@ -1229,6 +1253,504 @@ struct MacroForwardRootModule; #[derive(Debug)] struct MacroForwardFeatureModule; +static MACRO_BUBBLED_STATE_CALLS: AtomicUsize = AtomicUsize::new(0); +static MACRO_REQUEST_STATE_CALLS: AtomicUsize = AtomicUsize::new(0); +static MACRO_RETRY_CONTROLLER_CALLS: AtomicUsize = AtomicUsize::new(0); +static MACRO_RETRY_CONTROLLER_ATTEMPTS: Mutex> = Mutex::new(Vec::new()); + +#[derive(Debug)] +struct MacroBubbledState { + id: usize, +} + +#[injectable] +#[derive(Debug)] +struct MacroBubbledController { + first: Arc, + second: Arc, +} + +#[controller("/macro-bubbled-controller")] +#[a3s_boot::metadata("controller-scope", "bubbled")] +#[a3s_boot::use_interceptor(MacroControllerHeaderInterceptor)] +impl MacroBubbledController { + #[a3s_boot::get("/", raw)] + #[a3s_boot::metadata("route-scope", "bubbled")] + async fn current(&self) -> Result { + Ok(BootResponse::text(format!( + "{}:{}:{}", + self.first.id, + self.second.id, + Arc::ptr_eq(&self.first, &self.second), + ))) + } +} + +#[derive(Debug)] +struct MacroRequestState { + id: usize, +} + +#[injectable] +#[derive(Debug)] +struct MacroRequestController { + state: Arc, +} + +#[controller("/macro-request-controller")] +impl MacroRequestController { + #[a3s_boot::get("/", raw)] + async fn current(&self) -> Result { + Ok(BootResponse::text(self.state.id.to_string())) + } +} + +#[derive(Debug)] +struct MacroRetryRequestController { + id: usize, + attempts: AtomicUsize, +} + +impl FromModuleRef for MacroRetryRequestController { + fn from_module_ref(_module_ref: &ModuleRef) -> Result { + Ok(Self { + id: MACRO_RETRY_CONTROLLER_CALLS.fetch_add(1, Ordering::SeqCst) + 1, + attempts: AtomicUsize::new(0), + }) + } + + fn provider_dependencies() -> Option> { + Some(Vec::new()) + } +} + +struct MacroRetryOnceInterceptor; + +impl Interceptor for MacroRetryOnceInterceptor { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + match next.handle().await { + Ok(response) => Ok(response), + Err(_) => next.handle().await, + } + }) + } +} + +#[controller("/macro-retry-request-controller")] +#[a3s_boot::use_interceptor(MacroRetryOnceInterceptor)] +impl MacroRetryRequestController { + #[a3s_boot::get("/", raw)] + async fn current(&self) -> Result { + MACRO_RETRY_CONTROLLER_ATTEMPTS + .lock() + .unwrap() + .push(self.id); + if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 { + return Err(BootError::Internal("retry controller".to_string())); + } + + Ok(BootResponse::text(self.id.to_string())) + } +} + +#[a3s_boot::module( + name = "macro-provider-controllers", + providers = [ + ProviderDefinition::request_scoped::(|_| { + Ok(MacroBubbledState { + id: MACRO_BUBBLED_STATE_CALLS.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + ProviderDefinition::request_scoped::(|_| { + Ok(MacroRequestState { + id: MACRO_REQUEST_STATE_CALLS.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + MacroBubbledController, + MacroRequestController::request_scoped_provider(), + ProviderDefinition::request_scoped_injectable::(), + ], + controllers = [ + MacroBubbledController, + MacroRequestController, + MacroRetryRequestController, + ], +)] +#[derive(Debug)] +struct MacroProviderControllerModule; + +#[test] +fn injectable_macros_describe_provider_dependencies() { + let unit = MacroCatsService::provider(); + assert_eq!(unit.dependencies(), Some(&[][..])); + + let reader = MacroAutoCatsReader::provider(); + let dependencies = reader.dependencies().unwrap(); + let actual = dependencies + .iter() + .map(|dependency| { + ( + dependency.token().clone(), + dependency.is_optional(), + dependency.is_lazy(), + ) + }) + .collect::>(); + + assert_eq!( + actual, + vec![ + (ProviderToken::of::(), false, false), + (ProviderToken::named("readonly-cats"), false, false), + (ProviderToken::of::(), true, false,), + (ProviderToken::named("missing-cats"), true, false), + (ProviderToken::of::(), false, true), + (ProviderToken::named("readonly-cats"), false, true), + (ProviderToken::of::(), true, true,), + ] + ); + + let named = MacroAutoCatsReader::named_provider("reader"); + assert_eq!(named.dependencies(), reader.dependencies()); +} + +#[tokio::test] +async fn module_macros_use_provider_backed_contextual_controllers() { + MACRO_BUBBLED_STATE_CALLS.store(0, Ordering::SeqCst); + MACRO_REQUEST_STATE_CALLS.store(0, Ordering::SeqCst); + let app = BootApplication::builder() + .import(MacroProviderControllerModule) + .build() + .unwrap(); + + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert_eq!( + app.module_ref() + .provider_scope::() + .unwrap(), + ProviderScope::Singleton, + ); + assert_eq!( + app.module_ref() + .provider_scope::() + .unwrap(), + ProviderScope::Request, + ); + let first_bubbled = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-bubbled-controller", + )) + .await + .unwrap(); + let second_bubbled = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-bubbled-controller", + )) + .await + .unwrap(); + let first_request = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-request-controller", + )) + .await + .unwrap(); + let second_request = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-request-controller", + )) + .await + .unwrap(); + + assert_eq!(first_bubbled.body_text().unwrap(), "1:1:true"); + assert_eq!(first_bubbled.header("x-macro-controller"), Some("yes")); + assert_eq!(second_bubbled.body_text().unwrap(), "2:2:true"); + assert_eq!(first_request.body_text().unwrap(), "1"); + assert_eq!(second_request.body_text().unwrap(), "2"); + assert_eq!(MACRO_BUBBLED_STATE_CALLS.load(Ordering::SeqCst), 2); + assert_eq!(MACRO_REQUEST_STATE_CALLS.load(Ordering::SeqCst), 2); + + let first_scope = app.module_ref().request_scope(); + let first_bubbled_controller = first_scope.get::().unwrap(); + let repeated_bubbled_controller = first_scope.get::().unwrap(); + let first_request_controller = first_scope.get::().unwrap(); + let repeated_request_controller = first_scope.get::().unwrap(); + assert!(Arc::ptr_eq( + &first_bubbled_controller, + &repeated_bubbled_controller, + )); + assert!(Arc::ptr_eq( + &first_request_controller, + &repeated_request_controller, + )); + + let second_scope = app.module_ref().request_scope(); + let second_bubbled_controller = second_scope.get::().unwrap(); + let second_request_controller = second_scope.get::().unwrap(); + assert!(!Arc::ptr_eq( + &first_bubbled_controller, + &second_bubbled_controller, + )); + assert!(!Arc::ptr_eq( + &first_request_controller, + &second_request_controller, + )); + + assert_eq!( + app.reflector().unwrap().metadata_value( + a3s_boot::HttpMethod::Get, + "/macro-bubbled-controller", + "controller-scope", + ), + Some(&json!("bubbled")), + ); + assert_eq!( + app.reflector().unwrap().metadata_value( + a3s_boot::HttpMethod::Get, + "/macro-bubbled-controller", + "route-scope", + ), + Some(&json!("bubbled")), + ); +} + +#[tokio::test] +async fn interceptor_retry_reuses_the_request_scoped_controller_instance() { + MACRO_RETRY_CONTROLLER_CALLS.store(0, Ordering::SeqCst); + MACRO_RETRY_CONTROLLER_ATTEMPTS.lock().unwrap().clear(); + let app = BootApplication::builder() + .import(MacroProviderControllerModule) + .build() + .unwrap(); + + assert_eq!( + app.module_ref() + .provider_scope::() + .unwrap(), + ProviderScope::Request, + ); + + let first = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-retry-request-controller", + )) + .await + .unwrap(); + let second = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-retry-request-controller", + )) + .await + .unwrap(); + + assert_eq!(first.body_text().unwrap(), "1"); + assert_eq!(second.body_text().unwrap(), "2"); + assert_eq!(MACRO_RETRY_CONTROLLER_CALLS.load(Ordering::SeqCst), 2); + assert_eq!( + *MACRO_RETRY_CONTROLLER_ATTEMPTS.lock().unwrap(), + [1, 1, 2, 2], + ); +} + +#[tokio::test] +async fn provider_backed_controllers_support_all_http_handler_flavors() { + let app = BootApplication::builder() + .import(MacroProviderFlavorModule) + .build() + .unwrap(); + + let raw = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/raw-id", + )) + .await + .unwrap(); + assert_eq!(raw.body_text().unwrap(), "raw-id:Milo"); + + let json_response = app + .call( + BootRequest::new(a3s_boot::HttpMethod::Get, "/macro-cats/json-id/json") + .with_header("accept", "application/json"), + ) + .await + .unwrap(); + assert_eq!( + json_response.body_json::().unwrap(), + MacroCatDto { + id: "json-id".to_string(), + name: "Milo".to_string(), + }, + ); + + let json_body = BootRequest::new(a3s_boot::HttpMethod::Post, "/macro-cats") + .with_json(&MacroCreateCatDto { + name: "Nori".to_string(), + }) + .unwrap(); + let created = app.call(json_body).await.unwrap(); + assert_eq!(created.status(), 201); + assert_eq!( + created.body_json::().unwrap(), + MacroCatDto { + id: "generated".to_string(), + name: "Nori".to_string(), + }, + ); + + let extracted = app + .call( + BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/pipe/cat_extracted?page=3", + ) + .with_header("x-cat-kind", " TABBY "), + ) + .await + .unwrap(); + assert_eq!( + extracted.body_json::().unwrap(), + MacroCatDto { + id: "cat_extracted".to_string(), + name: "TABBY:3".to_string(), + }, + ); + + let extracted_raw = app + .call( + BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/extracted-raw/raw-details", + ) + .with_header("x-request-id", "provider-request") + .with_header("user-agent", "provider-controller-test"), + ) + .await + .unwrap(); + assert_eq!( + extracted_raw.body_text().unwrap(), + "extracted-raw:provider-request:provider-controller-test:/macro-cats/extracted-raw/raw-details", + ); + + let extracted_body = + BootRequest::new(a3s_boot::HttpMethod::Post, "/macro-cats/adopted/adoptions") + .with_json(&MacroCreateCatDto { + name: "Nori".to_string(), + }) + .unwrap(); + let adopted = app.call(extracted_body).await.unwrap(); + assert_eq!(adopted.status(), 201); + assert_eq!( + adopted.body_json::().unwrap(), + MacroCatDto { + id: "adopted".to_string(), + name: "Nori".to_string(), + }, + ); + + let catch_all = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Patch, + "/macro-cats/catch", + )) + .await + .unwrap(); + assert_eq!( + catch_all.body_json::().unwrap(), + MacroCatDto { + id: "PATCH".to_string(), + name: "Catch".to_string(), + }, + ); + + let rendered = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/rendered/card", + )) + .await + .unwrap(); + assert_eq!( + rendered.body_text().unwrap(), + "
rendered:Milo
", + ); + + let cached = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/cache", + )) + .await + .unwrap(); + assert_eq!(cached.header("cache-control"), Some("max-age=60")); + + let redirected = app + .call(BootRequest::new( + a3s_boot::HttpMethod::Get, + "/macro-cats/legacy", + )) + .await + .unwrap(); + assert_eq!(redirected.status(), 301); + assert_eq!(redirected.location(), Some("/macro-cats/42")); + + let events = app + .call( + BootRequest::new(a3s_boot::HttpMethod::Get, "/macro-cats/events") + .with_header("accept", "text/event-stream"), + ) + .await + .unwrap(); + let mut stream = events.into_sse_stream().unwrap(); + assert_eq!( + String::from_utf8(stream.next().await.unwrap().unwrap().encode()).unwrap(), + "event: cat.found\ndata: Milo\n\n", + ); + assert!(stream.next().await.is_none()); + + let extracted_events = app + .call( + BootRequest::new(a3s_boot::HttpMethod::Get, "/macro-cats/provider-sse/events") + .with_header("accept", "text/event-stream"), + ) + .await + .unwrap(); + let mut extracted_stream = extracted_events.into_sse_stream().unwrap(); + assert_eq!( + String::from_utf8(extracted_stream.next().await.unwrap().unwrap().encode(),).unwrap(), + "event: cat.selected\ndata: provider-sse\n\n", + ); + assert!(extracted_stream.next().await.is_none()); + + assert_eq!( + app.reflector().unwrap().metadata_value( + a3s_boot::HttpMethod::Get, + "/macro-cats/{id}/details", + "roles", + ), + Some(&json!(["admin"])), + ); + let document = + serde_json::to_value(app.openapi(OpenApiInfo::new("Provider", "1.0.0"))).unwrap(); + assert_eq!( + document["paths"]["/macro-cats/{id}/details"]["get"]["operationId"], + json!("findMacroCatDetails"), + ); +} + #[tokio::test] async fn macros_register_injectable_services_and_controller_routes() { let app = BootApplication::builder() @@ -1245,6 +1767,32 @@ async fn macros_register_injectable_services_and_controller_routes() { assert!(exports.contains(&ProviderToken::named("readonly-cats"))); let reader = app.get::().unwrap(); let controller = app.get::().unwrap(); + let instance_controller = Arc::clone(&controller).controller().unwrap(); + let provider_controller = MacroCatsController::provider_controller().unwrap(); + assert_eq!( + instance_controller.routes().len(), + provider_controller.routes().len() + ); + for (instance_route, provider_route) in instance_controller + .routes() + .iter() + .zip(provider_controller.routes()) + { + assert_eq!(instance_route.method(), provider_route.method()); + assert_eq!(instance_route.path(), provider_route.path()); + assert_eq!(instance_route.host(), provider_route.host()); + assert_eq!(instance_route.openapi(), provider_route.openapi()); + assert_eq!(instance_route.versioning(), provider_route.versioning()); + assert_eq!( + instance_route.serialization(), + provider_route.serialization() + ); + assert_eq!(instance_route.metadata(), provider_route.metadata()); + assert_eq!( + instance_route.validation_enabled(), + provider_route.validation_enabled(), + ); + } assert_eq!( reader.summary(), "auto:readonly:true:true:lazy:lazy-readonly:true" @@ -2125,6 +2673,12 @@ async fn macro_pipeline_decorators_register_controller_and_route_hooks() { .import(MacroPipelineModule) .build() .unwrap(); + assert_eq!( + app.module_ref() + .provider_scope::() + .unwrap(), + ProviderScope::Request, + ); let guarded = app .call(BootRequest::new( @@ -2408,6 +2962,12 @@ async fn validate_macro_enables_body_and_query_dto_validation() { .import(MacroValidationModule) .build() .unwrap(); + assert_eq!( + app.module_ref() + .provider_scope::() + .unwrap(), + ProviderScope::Request, + ); let body_error = app .call( diff --git a/tests/message_controller_contexts.rs b/tests/message_controller_contexts.rs new file mode 100644 index 0000000..2769a12 --- /dev/null +++ b/tests/message_controller_contexts.rs @@ -0,0 +1,348 @@ +#![cfg(feature = "macros")] + +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +use a3s_boot::{ + injectable, BootApplication, BootError, BoxFuture, CallHandler, ProviderDefinition, Result, + TransportContext, TransportInterceptor, TransportMessage, TransportReply, +}; +use serde::{Deserialize, Serialize}; +use serde_json::json; + +static REQUEST_CONSTRUCTIONS: AtomicUsize = AtomicUsize::new(0); +static REQUEST_EVENTS: Mutex> = Mutex::new(Vec::new()); +static TRANSIENT_CONSTRUCTIONS: AtomicUsize = AtomicUsize::new(0); +static TRANSIENT_ATTEMPTS: Mutex> = Mutex::new(Vec::new()); +static CONTEXT_STATE_CONSTRUCTIONS: AtomicUsize = AtomicUsize::new(0); +static SINGLETON_CONSTRUCTIONS: AtomicUsize = AtomicUsize::new(0); + +#[derive(Debug, Deserialize)] +struct TypedPayload { + value: String, +} + +#[derive(Debug, Deserialize, PartialEq, Eq, Serialize)] +struct ControllerReply { + controller_id: usize, + value: String, +} + +#[derive(Debug)] +struct RequestMessageController { + id: usize, +} + +impl RequestMessageController { + fn provider() -> ProviderDefinition { + ProviderDefinition::request_scoped::(|_| { + Ok(Self { + id: REQUEST_CONSTRUCTIONS.fetch_add(1, Ordering::SeqCst) + 1, + }) + }) + } +} + +#[a3s_boot::message_controller] +#[a3s_boot::metadata("controller", "request")] +impl RequestMessageController { + #[a3s_boot::message_pattern("macro.context.request.typed")] + #[a3s_boot::metadata("binding", "typed")] + async fn typed(&self, payload: TypedPayload) -> Result { + Ok(ControllerReply { + controller_id: self.id, + value: payload.value, + }) + } + + #[a3s_boot::event_pattern("macro.context.request.event")] + async fn event(&self, #[a3s_boot::payload("value")] value: String) -> Result<()> { + REQUEST_EVENTS.lock().unwrap().push((self.id, value)); + Ok(()) + } +} + +#[derive(Debug)] +struct TransientMessageController { + id: usize, + attempts: AtomicUsize, +} + +impl TransientMessageController { + fn provider() -> ProviderDefinition { + ProviderDefinition::transient::(|_| { + Ok(Self { + id: TRANSIENT_CONSTRUCTIONS.fetch_add(1, Ordering::SeqCst) + 1, + attempts: AtomicUsize::new(0), + }) + }) + } +} + +struct RetryOnce; + +impl TransportInterceptor for RetryOnce { + fn intercept<'a>( + &'a self, + _context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + match next.handle().await { + Ok(reply) => Ok(reply), + Err(_) => next.handle().await, + } + }) + } +} + +#[a3s_boot::message_controller] +impl TransientMessageController { + #[a3s_boot::message_pattern("macro.context.transient.retry")] + #[a3s_boot::use_interceptor(RetryOnce)] + async fn retry(&self, #[a3s_boot::payload("value")] value: String) -> Result { + TRANSIENT_ATTEMPTS.lock().unwrap().push(self.id); + if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 { + return Err(BootError::Internal( + "retry transient controller".to_string(), + )); + } + + Ok(ControllerReply { + controller_id: self.id, + value, + }) + } +} + +#[derive(Debug)] +struct ContextState { + id: usize, +} + +#[injectable] +#[derive(Debug)] +struct ContextualSingletonMessageController { + state: Arc, +} + +#[a3s_boot::message_controller] +impl ContextualSingletonMessageController { + #[a3s_boot::message_pattern("macro.context.singleton.bubbled")] + async fn current(&self, message: TransportMessage) -> Result { + Ok(ControllerReply { + controller_id: self.state.id, + value: message.pattern().to_string(), + }) + } +} + +#[derive(Debug)] +struct PureSingletonMessageController { + id: usize, +} + +impl PureSingletonMessageController { + fn provider() -> ProviderDefinition { + ProviderDefinition::factory::(|_| { + Ok(Self { + id: SINGLETON_CONSTRUCTIONS.fetch_add(1, Ordering::SeqCst) + 1, + }) + }) + } +} + +#[a3s_boot::message_controller] +impl PureSingletonMessageController { + #[a3s_boot::message_pattern("macro.context.singleton.pure")] + async fn current(&self, payload: TypedPayload) -> Result { + Ok(ControllerReply { + controller_id: self.id, + value: payload.value, + }) + } +} + +#[a3s_boot::module( + name = "macro-context-message-controllers", + providers = [ + ProviderDefinition::request_scoped::(|_| { + Ok(ContextState { + id: CONTEXT_STATE_CONSTRUCTIONS.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + RequestMessageController, + TransientMessageController, + ContextualSingletonMessageController, + PureSingletonMessageController, + ], + message_controllers = [ + RequestMessageController, + TransientMessageController, + ContextualSingletonMessageController, + PureSingletonMessageController, + ], +)] +#[derive(Debug)] +struct ContextMessageControllerModule; + +async fn dispatch_reply( + app: &BootApplication, + pattern: &str, + data: serde_json::Value, +) -> ControllerReply { + app.dispatch_message(TransportMessage::new(pattern, data)) + .await + .unwrap() + .unwrap() + .data_as::() + .unwrap() +} + +#[tokio::test] +async fn module_macros_scope_message_controllers_per_dispatch() { + REQUEST_CONSTRUCTIONS.store(0, Ordering::SeqCst); + REQUEST_EVENTS.lock().unwrap().clear(); + TRANSIENT_CONSTRUCTIONS.store(0, Ordering::SeqCst); + TRANSIENT_ATTEMPTS.lock().unwrap().clear(); + CONTEXT_STATE_CONSTRUCTIONS.store(0, Ordering::SeqCst); + SINGLETON_CONSTRUCTIONS.store(0, Ordering::SeqCst); + + let app = BootApplication::builder() + .import(ContextMessageControllerModule) + .build() + .unwrap(); + + for pattern in [ + "macro.context.request.typed", + "macro.context.request.event", + "macro.context.transient.retry", + "macro.context.singleton.bubbled", + ] { + assert!(app.message_pattern_for(pattern).unwrap().is_scoped()); + } + assert!(!app + .message_pattern_for("macro.context.singleton.pure") + .unwrap() + .is_scoped()); + + assert_eq!(REQUEST_CONSTRUCTIONS.load(Ordering::SeqCst), 0); + assert_eq!(TRANSIENT_CONSTRUCTIONS.load(Ordering::SeqCst), 0); + assert_eq!(CONTEXT_STATE_CONSTRUCTIONS.load(Ordering::SeqCst), 0); + assert_eq!(SINGLETON_CONSTRUCTIONS.load(Ordering::SeqCst), 1); + + let request_one = dispatch_reply( + &app, + "macro.context.request.typed", + json!({ "value": "first" }), + ) + .await; + let request_two = dispatch_reply( + &app, + "macro.context.request.typed", + json!({ "value": "second" }), + ) + .await; + assert_eq!( + (request_one.controller_id, request_one.value.as_str()), + (1, "first") + ); + assert_eq!( + (request_two.controller_id, request_two.value.as_str()), + (2, "second") + ); + + app.emit_message(TransportMessage::new( + "macro.context.request.event", + json!({ "value": "observed" }), + )) + .await + .unwrap(); + assert_eq!( + REQUEST_EVENTS.lock().unwrap().as_slice(), + &[(3, "observed".to_string())] + ); + assert_eq!(REQUEST_CONSTRUCTIONS.load(Ordering::SeqCst), 3); + + let transient_one = dispatch_reply( + &app, + "macro.context.transient.retry", + json!({ "value": "retry-one" }), + ) + .await; + let transient_two = dispatch_reply( + &app, + "macro.context.transient.retry", + json!({ "value": "retry-two" }), + ) + .await; + assert_eq!(transient_one.controller_id, 1); + assert_eq!(transient_two.controller_id, 2); + assert_eq!(TRANSIENT_ATTEMPTS.lock().unwrap().as_slice(), &[1, 1, 2, 2]); + assert_eq!(TRANSIENT_CONSTRUCTIONS.load(Ordering::SeqCst), 2); + + let contextual_one = dispatch_reply( + &app, + "macro.context.singleton.bubbled", + json!({ "ignored": true }), + ) + .await; + let contextual_two = dispatch_reply( + &app, + "macro.context.singleton.bubbled", + json!({ "ignored": true }), + ) + .await; + assert_eq!(contextual_one.controller_id, 1); + assert_eq!(contextual_two.controller_id, 2); + assert_eq!(CONTEXT_STATE_CONSTRUCTIONS.load(Ordering::SeqCst), 2); + + let singleton_one = dispatch_reply( + &app, + "macro.context.singleton.pure", + json!({ "value": "captured-one" }), + ) + .await; + let singleton_two = dispatch_reply( + &app, + "macro.context.singleton.pure", + json!({ "value": "captured-two" }), + ) + .await; + assert_eq!(singleton_one.controller_id, 1); + assert_eq!(singleton_two.controller_id, 1); + assert_eq!(SINGLETON_CONSTRUCTIONS.load(Ordering::SeqCst), 1); + + let request_pattern = app + .message_pattern_for("macro.context.request.typed") + .unwrap(); + assert_eq!( + request_pattern.metadata_value("controller"), + Some(&json!("request")) + ); + assert_eq!( + request_pattern.metadata_value("binding"), + Some(&json!("typed")) + ); +} + +#[test] +fn generated_message_controller_metadata_exposes_both_handler_modes() { + let instance = Arc::new(RequestMessageController { id: 7 }) + .message_patterns() + .unwrap(); + let provider = RequestMessageController::provider_message_patterns().unwrap(); + + assert_eq!(instance.len(), provider.len()); + for (instance, provider) in instance.iter().zip(&provider) { + assert!(!instance.is_scoped()); + assert!(provider.is_scoped()); + assert_eq!(instance.pattern(), provider.pattern()); + assert_eq!(instance.metadata(), provider.metadata()); + } + assert_eq!( + provider[0].metadata_value("controller"), + Some(&json!("request")) + ); + assert_eq!(provider[0].metadata_value("binding"), Some(&json!("typed"))); +} diff --git a/tests/module_registration.rs b/tests/module_registration.rs new file mode 100644 index 0000000..245a99d --- /dev/null +++ b/tests/module_registration.rs @@ -0,0 +1,191 @@ +use a3s_boot::{ + BootApplication, BootError, Module, ModuleRef, ProviderDefinition, ProviderToken, Result, +}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug)] +struct LaterAsyncGlobalConfig { + value: &'static str, +} + +#[derive(Debug)] +struct EarlyAsyncConsumer { + config: Arc, +} + +#[derive(Debug)] +struct EarlyAsyncConsumerModule { + log: Arc>>, +} + +impl Module for EarlyAsyncConsumerModule { + fn name(&self) -> &'static str { + "early-async-consumer" + } + + fn providers(&self) -> Result> { + let log = Arc::clone(&self.log); + Ok(vec![ProviderDefinition::async_factory::< + EarlyAsyncConsumer, + _, + _, + >(move |module_ref| { + let log = Arc::clone(&log); + async move { + let config = module_ref.get::()?; + log.lock().unwrap().push("consumer-factory"); + Ok(EarlyAsyncConsumer { config }) + } + }) + .depends_on::()]) + } + + fn on_module_init(&self, _module_ref: &ModuleRef) -> Result<()> { + self.log.lock().unwrap().push("consumer-init"); + Ok(()) + } +} + +#[derive(Debug)] +struct LaterAsyncGlobalModule { + log: Arc>>, +} + +impl Module for LaterAsyncGlobalModule { + fn name(&self) -> &'static str { + "later-async-global" + } + + fn providers(&self) -> Result> { + let log = Arc::clone(&self.log); + Ok(vec![ProviderDefinition::async_factory::< + LaterAsyncGlobalConfig, + _, + _, + >(move |_| { + let log = Arc::clone(&log); + async move { + log.lock().unwrap().push("config-factory"); + Ok(LaterAsyncGlobalConfig { value: "ready" }) + } + })]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } + + fn is_global(&self) -> bool { + true + } + + fn on_module_init(&self, _module_ref: &ModuleRef) -> Result<()> { + self.log.lock().unwrap().push("global-init"); + Ok(()) + } +} + +#[derive(Debug)] +struct LaterRequestGlobal; + +#[derive(Debug)] +struct InvalidAsyncContextConsumer; + +#[derive(Debug)] +struct InvalidAsyncContextConsumerModule { + calls: Arc, +} + +impl Module for InvalidAsyncContextConsumerModule { + fn name(&self) -> &'static str { + "invalid-async-context-consumer" + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ProviderDefinition::async_factory::< + InvalidAsyncContextConsumer, + _, + _, + >(move |module_ref| { + let calls = Arc::clone(&calls); + async move { + calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + let _ = module_ref.get::()?; + Ok(InvalidAsyncContextConsumer) + } + }) + .depends_on::()]) + } +} + +#[derive(Debug)] +struct LaterRequestGlobalModule; + +impl Module for LaterRequestGlobalModule { + fn name(&self) -> &'static str { + "later-request-global" + } + + fn providers(&self) -> Result> { + Ok(vec![ProviderDefinition::request_scoped::< + LaterRequestGlobal, + _, + >(|_| Ok(LaterRequestGlobal))]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } + + fn is_global(&self) -> bool { + true + } +} + +#[tokio::test] +async fn async_finalization_seeds_later_global_dependencies_before_consumers() { + let log = Arc::new(Mutex::new(Vec::new())); + let app = BootApplication::builder() + .import(EarlyAsyncConsumerModule { + log: Arc::clone(&log), + }) + .import(LaterAsyncGlobalModule { + log: Arc::clone(&log), + }) + .build_async() + .await + .unwrap(); + + let config = app.get::().unwrap(); + let consumer = app.get::().unwrap(); + + assert_eq!(config.value, "ready"); + assert!(Arc::ptr_eq(&config, &consumer.config)); + assert_eq!( + log.lock().unwrap().as_slice(), + [ + "config-factory", + "consumer-factory", + "consumer-init", + "global-init", + ] + ); +} + +#[tokio::test] +async fn full_graph_validation_rejects_contextual_async_providers_before_factories_run() { + let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let result = BootApplication::builder() + .import(InvalidAsyncContextConsumerModule { + calls: Arc::clone(&calls), + }) + .import(LaterRequestGlobalModule) + .build_async() + .await; + + assert!( + matches!(result, Err(BootError::Internal(message)) if message.contains("cannot depend on a request-context provider")) + ); + assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0); +} diff --git a/tests/pipeline.rs b/tests/pipeline.rs index fcaf69f..93c6eec 100644 --- a/tests/pipeline.rs +++ b/tests/pipeline.rs @@ -1,5 +1,5 @@ use a3s_boot::{ - BootApplication, BootError, BootErrorKind, BootRequest, BootResponse, BoxFuture, + BootApplication, BootError, BootErrorKind, BootRequest, BootResponse, BoxFuture, CallHandler, ControllerDefinition, ExecutionContext, ExecutionInterceptor, Guard, HttpMethod, Interceptor, MessagePatternDefinition, Middleware, MiddlewareConsumer, MiddlewareOutcome, MiddlewareRoute, Module, ModuleRef, Result, RouteDefinition, TransportMessage, TransportReply, @@ -396,6 +396,100 @@ async fn route_filters_handle_pipe_errors() { assert_eq!(response.body, b"/: bad request: invalid input"); } +#[tokio::test] +async fn route_filters_see_the_request_snapshot_before_the_failing_pipe() { + let route = RouteDefinition::get("/", |_| async { Ok(BootResponse::text("unreachable")) }) + .unwrap() + .with_pipe(|request: BootRequest| async move { + Ok(request.with_header("x-pipeline-stage", "first")) + }) + .with_pipe(|_| async { Err(BootError::BadRequest("second pipe failed".to_string())) }) + .with_filter(|context: ExecutionContext, _| async move { + Ok(Some(BootResponse::text( + context + .request + .header("x-pipeline-stage") + .unwrap_or("missing"), + ))) + }); + + let response = route + .call(BootRequest::new(HttpMethod::Get, "/")) + .await + .unwrap(); + + assert_eq!(response.body, b"first"); +} + +#[tokio::test] +async fn retried_interceptor_errors_reset_stale_pipe_filter_context() { + struct RetryOnce; + + impl Interceptor for RetryOnce { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + match next.handle().await { + Ok(response) => Ok(response), + Err(_) => next.handle().await, + } + }) + } + } + + struct FailSecondBefore { + calls: Arc, + } + + impl Interceptor for FailSecondBefore { + fn before(&self, _context: ExecutionContext) -> BoxFuture<'static, Result<()>> { + let calls = Arc::clone(&self.calls); + Box::pin(async move { + if calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst) == 0 { + Ok(()) + } else { + Err(BootError::Internal("second before failed".to_string())) + } + }) + } + } + + let before_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let route = RouteDefinition::get("/", |_| async { Ok(BootResponse::text("unreachable")) }) + .unwrap() + .with_interceptor(RetryOnce) + .with_interceptor(FailSecondBefore { + calls: Arc::clone(&before_calls), + }) + .with_pipe(|request: BootRequest| async move { + Ok(request.with_header("x-pipeline-stage", "stale")) + }) + .with_pipe(|_| async { Err(BootError::BadRequest("retry".to_string())) }) + .with_filter(|context: ExecutionContext, error: BootError| async move { + Ok(Some(BootResponse::text(format!( + "{}:{error}", + context + .request + .header("x-pipeline-stage") + .unwrap_or("missing") + )))) + }); + + let response = route + .call(BootRequest::new(HttpMethod::Get, "/")) + .await + .unwrap(); + + assert_eq!( + response.body, + b"missing:internal error: second before failed" + ); + assert_eq!(before_calls.load(std::sync::atomic::Ordering::SeqCst), 2); +} + #[tokio::test] async fn route_filters_handle_interceptor_errors() { struct FailingAfterInterceptor; @@ -428,6 +522,279 @@ async fn route_filters_handle_interceptor_errors() { assert_eq!(response.body, b"/: internal error: response failed"); } +#[tokio::test] +async fn around_interceptors_can_recover_pipe_errors_before_filters() { + struct RecoverBadRequest; + + impl Interceptor for RecoverBadRequest { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + match next.handle().await { + Err(BootError::BadRequest(message)) => { + Ok(BootResponse::text(format!("recovered: {message}")).with_status(422)) + } + result => result, + } + }) + } + } + + let filter_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let filter_counter = Arc::clone(&filter_calls); + let route = RouteDefinition::post("/", |_| async { Ok(BootResponse::text("unreachable")) }) + .unwrap() + .with_interceptor(RecoverBadRequest) + .with_pipe(|_| async { Err(BootError::BadRequest("invalid input".to_string())) }) + .with_filter(move |_, _| { + let filter_counter = Arc::clone(&filter_counter); + async move { + filter_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(Some(BootResponse::text("filtered").with_status(400))) + } + }); + + let response = route + .call(BootRequest::new(HttpMethod::Post, "/")) + .await + .unwrap(); + + assert_eq!(response.status(), 422); + assert_eq!(response.body, b"recovered: invalid input"); + assert_eq!(filter_calls.load(std::sync::atomic::Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn around_interceptors_can_retry_the_downstream_pipeline() { + struct RetryOnce; + + impl Interceptor for RetryOnce { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + match next.handle().await { + Ok(response) => Ok(response), + Err(_) => next.handle().await, + } + }) + } + } + + struct CountDownstreamInterceptor { + before_calls: Arc, + after_calls: Arc, + } + + impl Interceptor for CountDownstreamInterceptor { + fn before(&self, _context: ExecutionContext) -> BoxFuture<'static, Result<()>> { + let calls = Arc::clone(&self.before_calls); + Box::pin(async move { + calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(()) + }) + } + + fn after( + &self, + _context: ExecutionContext, + response: BootResponse, + ) -> BoxFuture<'static, Result> { + let calls = Arc::clone(&self.after_calls); + Box::pin(async move { + calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(response) + }) + } + } + + let pipe_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let pipe_counter = Arc::clone(&pipe_calls); + let handler_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let handler_counter = Arc::clone(&handler_calls); + let inner_before_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let inner_after_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let route = RouteDefinition::get("/", move |_| { + let handler_counter = Arc::clone(&handler_counter); + async move { + let attempt = handler_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if attempt == 0 { + return Err(BootError::Internal("try again".to_string())); + } + Ok(BootResponse::text("ok")) + } + }) + .unwrap() + .with_interceptor(RetryOnce) + .with_interceptor(CountDownstreamInterceptor { + before_calls: Arc::clone(&inner_before_calls), + after_calls: Arc::clone(&inner_after_calls), + }) + .with_pipe(move |request: BootRequest| { + let pipe_counter = Arc::clone(&pipe_counter); + async move { + pipe_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(request) + } + }); + + let response = route + .call(BootRequest::new(HttpMethod::Get, "/")) + .await + .unwrap(); + + assert_eq!(response.body, b"ok"); + assert_eq!(pipe_calls.load(std::sync::atomic::Ordering::SeqCst), 2); + assert_eq!(handler_calls.load(std::sync::atomic::Ordering::SeqCst), 2); + assert_eq!( + inner_before_calls.load(std::sync::atomic::Ordering::SeqCst), + 2 + ); + assert_eq!( + inner_after_calls.load(std::sync::atomic::Ordering::SeqCst), + 1 + ); +} + +#[tokio::test] +async fn call_handler_rejects_concurrent_calls_and_resets_after_cancellation() { + fn assert_send_sync() {} + assert_send_sync::>(); + + let attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let observed = Arc::clone(&attempts); + let handler = CallHandler::from_fn(move || { + let observed = Arc::clone(&observed); + async move { + let attempt = observed.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if attempt == 0 { + std::future::pending::<()>().await; + } + Ok(BootResponse::text("completed")) + } + }); + + let running_handler = handler.clone(); + let running = tokio::spawn(async move { running_handler.handle().await }); + while attempts.load(std::sync::atomic::Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + + let error = handler.handle().await.unwrap_err(); + assert!(matches!( + error, + BootError::Internal(message) if message == "call handler is already running" + )); + + running.abort(); + assert!(running.await.unwrap_err().is_cancelled()); + + let response = handler.handle().await.unwrap(); + assert_eq!(response.body, b"completed"); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn around_interceptor_short_circuits_still_unwind_outer_legacy_hooks() { + struct OuterHeader; + + impl Interceptor for OuterHeader { + fn after( + &self, + _context: ExecutionContext, + response: BootResponse, + ) -> BoxFuture<'static, Result> { + Box::pin(async move { Ok(response.with_header("x-outer", "yes")) }) + } + } + + struct ShortCircuit; + + impl Interceptor for ShortCircuit { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + _next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async { Ok(BootResponse::text("cached").with_status(202)) }) + } + } + + let handler_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let handler_counter = Arc::clone(&handler_calls); + let route = RouteDefinition::get("/", move |_| { + let handler_counter = Arc::clone(&handler_counter); + async move { + handler_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(BootResponse::text("handler")) + } + }) + .unwrap() + .with_interceptor(OuterHeader) + .with_interceptor(ShortCircuit); + + let response = route + .call(BootRequest::new(HttpMethod::Get, "/")) + .await + .unwrap(); + + assert_eq!(response.status(), 202); + assert_eq!(response.body, b"cached"); + assert_eq!(response.header("x-outer"), Some("yes")); + assert_eq!(handler_calls.load(std::sync::atomic::Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn unrecovered_around_interceptor_errors_reach_filters_once() { + struct MapToConflict; + + impl Interceptor for MapToConflict { + fn intercept<'a>( + &'a self, + _context: ExecutionContext, + next: CallHandler<'a>, + ) -> BoxFuture<'a, Result> { + Box::pin(async move { + match next.handle().await { + Err(BootError::BadRequest(message)) => { + Err(BootError::Conflict(format!("mapped: {message}"))) + } + result => result, + } + }) + } + } + + let filter_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let filter_counter = Arc::clone(&filter_calls); + let route = RouteDefinition::get("/", |_| async { + Err(BootError::BadRequest("invalid".to_string())) + }) + .unwrap() + .with_interceptor(MapToConflict) + .with_filter(move |_, error: BootError| { + let filter_counter = Arc::clone(&filter_counter); + async move { + filter_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(Some(BootResponse::text(error.to_string()).with_status(409))) + } + }); + + let response = route + .call(BootRequest::new(HttpMethod::Get, "/")) + .await + .unwrap(); + + assert_eq!(response.status(), 409); + assert_eq!(response.body, b"resource conflict: mapped: invalid"); + assert_eq!(filter_calls.load(std::sync::atomic::Ordering::SeqCst), 1); +} + #[tokio::test] async fn route_filters_can_decline_and_next_filter_sees_original_error() { let log = Arc::new(std::sync::Mutex::new(Vec::new())); diff --git a/tests/provider_contexts.rs b/tests/provider_contexts.rs new file mode 100644 index 0000000..640f6f5 --- /dev/null +++ b/tests/provider_contexts.rs @@ -0,0 +1,990 @@ +use a3s_boot::{ + BootApplication, BootError, BootRequest, BootResponse, ContextId, ContextIdFactory, + FromModuleRef, HttpMethod, Module, ModuleRef, ProviderDefinition, ProviderDependency, + ProviderRef, ProviderToken, Result, RouteDefinition, +}; +use std::panic::{catch_unwind, AssertUnwindSafe}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{mpsc, Arc, Barrier, Condvar, Mutex}; +use std::thread; +use std::time::Duration; + +fn assert_send_sync() {} + +#[derive(Debug)] +struct RequestValue { + id: usize, +} + +#[derive(Debug)] +struct TransientValue { + id: usize, +} + +#[derive(Debug)] +struct FirstConsumer { + first: Arc, + second: Arc, +} + +#[derive(Debug)] +struct SecondConsumer { + value: Arc, +} + +#[derive(Debug)] +struct LazyConsumer { + eager: Arc, + lazy: ProviderRef, +} + +#[derive(Debug)] +struct PublishedLazyRoot; + +#[derive(Debug)] +struct PublishedLazyDependency { + root: Arc, +} + +#[derive(Debug)] +struct StaticLazyRequestConsumer { + request: ProviderRef, +} + +#[derive(Debug)] +struct ScopedLazyRequestConsumer { + request: ProviderRef, + drops: Arc, +} + +impl Drop for ScopedLazyRequestConsumer { + fn drop(&mut self) { + self.drops.fetch_add(1, Ordering::SeqCst); + } +} + +#[derive(Debug)] +struct CreatedConsumer { + first: Arc, + second: Arc, +} + +#[derive(Debug)] +struct SharedDependency; + +#[derive(Debug)] +struct ConcurrentRootA { + dependency: Arc, +} + +#[derive(Debug)] +struct ConcurrentRootB { + dependency: Arc, +} + +#[derive(Debug)] +struct PanicRecoveryConsumer { + dependency: Arc, +} + +#[derive(Debug)] +struct ParallelDependency; + +#[derive(Debug)] +struct ParallelConsumer { + dependency: Arc, +} + +#[derive(Debug)] +struct ParallelRoot { + consumer: Arc, + dependency: Arc, +} + +impl FromModuleRef for CreatedConsumer { + fn from_module_ref(module_ref: &ModuleRef) -> Result { + Ok(Self { + first: module_ref.get::()?, + second: module_ref.get::()?, + }) + } + + fn provider_dependencies() -> Option> { + Some(vec![ProviderDependency::typed::()]) + } +} + +#[test] +fn context_ids_are_unique_shareable_and_exposed_by_scoped_module_refs() { + assert_send_sync::(); + + let first = ContextIdFactory::create(); + let first_clone = first.clone(); + let second = ContextIdFactory::create(); + let module_ref = ModuleRef::new(); + let scoped = module_ref.context_scope(&first); + assert_eq!(first, first_clone); + assert_ne!(first, second); + assert_ne!(first.id(), second.id()); + assert_eq!(scoped.context_id(), Some(&first)); + assert!(format!("{first:?}").contains(&first.id().to_string())); +} + +#[test] +fn explicit_context_ids_reuse_request_providers_across_resolve_calls() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(RequestValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register(ProviderDefinition::named_alias( + "request-value", + ProviderToken::of::(), + )) + .unwrap(); + + let first_context = ContextIdFactory::create(); + let second_context = ContextIdFactory::create(); + let first = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let same = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let same_named = module_ref + .resolve_named_with_context::("request-value", &first_context) + .unwrap(); + let same_scoped = module_ref + .context_scope(&first_context) + .get::() + .unwrap(); + let second = module_ref + .resolve_with_context::(&second_context) + .unwrap(); + let fresh = module_ref.resolve::().unwrap(); + + assert_eq!(first.id, 1); + assert!(Arc::ptr_eq(&first, &same)); + assert!(Arc::ptr_eq(&first, &same_named)); + assert!(Arc::ptr_eq(&first, &same_scoped)); + assert_eq!(second.id, 2); + assert_eq!(fresh.id, 3); + assert!(!Arc::ptr_eq(&first, &second)); + assert!(!Arc::ptr_eq(&second, &fresh)); + assert!(module_ref + .resolve_optional_with_context::(&first_context) + .unwrap() + .is_none()); + assert!(module_ref + .resolve_optional_named_with_context::("missing", &first_context) + .unwrap() + .is_none()); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} + +#[test] +fn provider_refs_resolve_in_an_explicit_context() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(RequestValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + + let provider_ref = module_ref.provider_ref::(); + let first_context = ContextIdFactory::create(); + let second_context = ContextIdFactory::create(); + let first = provider_ref.resolve_with_context(&first_context).unwrap(); + let same = provider_ref.resolve_with_context(&first_context).unwrap(); + let second = provider_ref.resolve_with_context(&second_context).unwrap(); + let fresh = provider_ref.resolve().unwrap(); + + assert!(Arc::ptr_eq(&first, &same)); + assert!(!Arc::ptr_eq(&first, &second)); + assert!(!Arc::ptr_eq(&second, &fresh)); + assert_eq!((first.id, second.id, fresh.id), (1, 2, 3)); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} + +#[test] +fn provider_refs_detach_the_active_construction_path() { + let module_ref = ModuleRef::new(); + module_ref + .register( + ProviderDefinition::request_scoped::(|module_ref| { + Ok(PublishedLazyDependency { + root: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + let child_ready = Arc::new(Barrier::new(2)); + let (sender, receiver) = mpsc::channel(); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + let dependency = module_ref.provider_ref::(); + let thread_ready = Arc::clone(&child_ready); + let sender = sender.clone(); + thread::spawn(move || { + thread_ready.wait(); + sender.send(dependency.get()).unwrap(); + }); + child_ready.wait(); + thread::sleep(Duration::from_millis(25)); + Ok(PublishedLazyRoot) + }) + .with_dependency(ProviderDependency::typed::().lazy()), + ) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let root = module_ref + .resolve_with_context::(&context_id) + .unwrap(); + let dependency = receiver + .recv_timeout(Duration::from_secs(2)) + .unwrap() + .unwrap(); + + assert!(Arc::ptr_eq(&root, &dependency.root)); +} + +#[test] +fn transient_providers_are_reused_per_inquirer_without_changing_root_get_compatibility() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::transient::( + move |_| { + Ok(TransientValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(FirstConsumer { + first: module_ref.get::()?, + second: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(SecondConsumer { + value: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(LazyConsumer { + eager: module_ref.get::()?, + lazy: module_ref.provider_ref::(), + }) + }) + .with_dependencies([ + ProviderDependency::typed::(), + ProviderDependency::typed::().lazy(), + ]), + ) + .unwrap(); + + let first = module_ref.get::().unwrap(); + let second = module_ref.get::().unwrap(); + let lazy = module_ref.get::().unwrap(); + let lazy_value = lazy.lazy.get().unwrap(); + let direct_first = module_ref.get::().unwrap(); + let direct_second = module_ref.get::().unwrap(); + + assert!(Arc::ptr_eq(&first.first, &first.second)); + assert!(!Arc::ptr_eq(&first.first, &second.value)); + assert!(Arc::ptr_eq(&lazy.eager, &lazy_value)); + assert!(!Arc::ptr_eq(&direct_first, &direct_second)); + assert_ne!(first.first.id, second.value.id); + assert_ne!(direct_first.id, direct_second.id); + assert_eq!(calls.load(Ordering::SeqCst), 5); +} + +#[test] +fn static_singletons_do_not_capture_the_context_that_first_resolves_them() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(RequestValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(StaticLazyRequestConsumer { + request: module_ref.provider_ref::(), + }) + }) + .with_dependency(ProviderDependency::typed::().lazy()), + ) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let consumer = module_ref + .resolve_with_context::(&context_id) + .unwrap(); + + assert!(consumer.request.module_ref().context_id().is_none()); + assert!(matches!( + consumer.request.get(), + Err(BootError::Internal(message)) + if message.contains("requires an active request scope") + )); + assert_eq!(consumer.request.resolve().unwrap().id, 1); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn request_contexts_release_scoped_providers_that_hold_lazy_provider_refs() { + let drops = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + module_ref + .register(ProviderDefinition::request_scoped::( + |_| Ok(RequestValue { id: 1 }), + )) + .unwrap(); + let provider_drops = Arc::clone(&drops); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + Ok(ScopedLazyRequestConsumer { + request: module_ref.provider_ref::(), + drops: Arc::clone(&provider_drops), + }) + }) + .with_dependency(ProviderDependency::typed::().lazy()), + ) + .unwrap(); + + let weak_consumer = { + let context_id = ContextIdFactory::create(); + let consumer = module_ref + .resolve_with_context::(&context_id) + .unwrap(); + assert_eq!(consumer.request.get().unwrap().id, 1); + let weak_consumer = Arc::downgrade(&consumer); + + drop(consumer); + assert!(weak_consumer.upgrade().is_some()); + weak_consumer + }; + + assert!( + weak_consumer.upgrade().is_none(), + "the released ContextId must not be retained by its cached provider" + ); + assert_eq!(drops.load(Ordering::SeqCst), 1); +} + +#[test] +fn transient_instances_are_scoped_by_context_and_consumer() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::transient::( + move |_| { + Ok(TransientValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::request_scoped::(|module_ref| { + Ok(FirstConsumer { + first: module_ref.get::()?, + second: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::request_scoped::(|module_ref| { + Ok(SecondConsumer { + value: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + let first_context = ContextIdFactory::create(); + let second_context = ContextIdFactory::create(); + let first = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let first_again = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let other_consumer = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let other_context = module_ref + .resolve_with_context::(&second_context) + .unwrap(); + let root_first = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + let root_again = module_ref + .resolve_with_context::(&first_context) + .unwrap(); + + assert!(Arc::ptr_eq(&first, &first_again)); + assert!(Arc::ptr_eq(&first.first, &first.second)); + assert!(!Arc::ptr_eq(&first.first, &other_consumer.value)); + assert!(!Arc::ptr_eq(&first.first, &other_context.first)); + assert!(Arc::ptr_eq(&root_first, &root_again)); + assert!(!Arc::ptr_eq(&root_first, &first.first)); + assert_eq!(calls.load(Ordering::SeqCst), 4); +} + +#[test] +fn module_ref_create_assigns_a_distinct_synthetic_inquirer() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::transient::( + move |_| { + Ok(TransientValue { + id: provider_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + + let first = module_ref.create::().unwrap(); + let second = module_ref.create::().unwrap(); + + assert!(Arc::ptr_eq(&first.first, &first.second)); + assert!(Arc::ptr_eq(&second.first, &second.second)); + assert!(!Arc::ptr_eq(&first.first, &second.first)); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn concurrent_resolution_in_one_context_runs_a_factory_once() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + provider_calls.fetch_add(1, Ordering::SeqCst); + thread::sleep(Duration::from_millis(25)); + Ok(RequestValue { id: 1 }) + }, + )) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let barrier = Arc::new(Barrier::new(3)); + let handles = (0..2) + .map(|_| { + let module_ref = module_ref.clone(); + let context_id = context_id.clone(); + let barrier = Arc::clone(&barrier); + thread::spawn(move || { + barrier.wait(); + module_ref.resolve_with_context::(&context_id) + }) + }) + .collect::>(); + barrier.wait(); + + let mut values = handles + .into_iter() + .map(|handle| handle.join().unwrap().unwrap()) + .collect::>(); + let second = values.pop().unwrap(); + let first_value = values.pop().unwrap(); + + assert!(Arc::ptr_eq(&first_value, &second)); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn concurrent_roots_can_share_a_dependency_without_a_false_cycle() { + let module_ref = ModuleRef::new(); + let shared_calls = Arc::new(AtomicUsize::new(0)); + let root_calls = Arc::new(AtomicUsize::new(0)); + let provider_calls = Arc::clone(&shared_calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + provider_calls.fetch_add(1, Ordering::SeqCst); + thread::sleep(Duration::from_millis(25)); + Ok(SharedDependency) + }, + )) + .unwrap(); + + let roots_ready = Arc::new(Barrier::new(2)); + let a_ready = Arc::clone(&roots_ready); + let a_calls = Arc::clone(&root_calls); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + a_calls.fetch_add(1, Ordering::SeqCst); + a_ready.wait(); + Ok(ConcurrentRootA { + dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + let b_ready = Arc::clone(&roots_ready); + let b_calls = Arc::clone(&root_calls); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + b_calls.fetch_add(1, Ordering::SeqCst); + b_ready.wait(); + Ok(ConcurrentRootB { + dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let first_ref = module_ref.clone(); + let first_context = context_id.clone(); + let first = + thread::spawn(move || first_ref.resolve_with_context::(&first_context)); + let second = + thread::spawn(move || module_ref.resolve_with_context::(&context_id)); + + let first = first.join().unwrap().unwrap(); + let second = second.join().unwrap().unwrap(); + + assert!(Arc::ptr_eq(&first.dependency, &second.dependency)); + assert_eq!(shared_calls.load(Ordering::SeqCst), 1); + assert_eq!(root_calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn failed_context_construction_releases_the_single_flight_slot() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + let call = provider_calls.fetch_add(1, Ordering::SeqCst) + 1; + if call == 1 { + return Err(BootError::Internal("first construction failed".to_string())); + } + Ok(RequestValue { id: call }) + }, + )) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let first = module_ref.resolve_with_context::(&context_id); + let second = module_ref + .resolve_with_context::(&context_id) + .unwrap(); + let same = module_ref + .resolve_with_context::(&context_id) + .unwrap(); + + assert!(matches!( + first, + Err(BootError::Internal(message)) if message == "first construction failed" + )); + assert_eq!(second.id, 2); + assert!(Arc::ptr_eq(&second, &same)); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn panicking_context_construction_releases_waiters_and_allows_retry() { + let calls = Arc::new(AtomicUsize::new(0)); + let first_started = Arc::new(Barrier::new(2)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + let provider_started = Arc::clone(&first_started); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + let call = provider_calls.fetch_add(1, Ordering::SeqCst) + 1; + if call == 1 { + provider_started.wait(); + thread::sleep(Duration::from_millis(25)); + panic!("first construction panicked"); + } + Ok(RequestValue { id: call }) + }, + )) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let (panic_sender, panic_receiver) = mpsc::channel(); + let (result_sender, result_receiver) = mpsc::channel(); + let first_ref = module_ref.clone(); + let first_context = context_id.clone(); + thread::spawn(move || { + let panicked = catch_unwind(AssertUnwindSafe(|| { + first_ref.resolve_with_context::(&first_context) + })) + .is_err(); + panic_sender.send(panicked).unwrap(); + }); + first_started.wait(); + + thread::spawn(move || { + let result = module_ref + .resolve_with_context::(&context_id) + .map(|value| value.id); + result_sender.send(result).unwrap(); + }); + + assert!(panic_receiver.recv_timeout(Duration::from_secs(2)).unwrap()); + assert_eq!( + result_receiver + .recv_timeout(Duration::from_secs(2)) + .unwrap() + .unwrap(), + 2 + ); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn caught_dependency_panics_do_not_leave_stale_resolution_frames() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let provider_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + let call = provider_calls.fetch_add(1, Ordering::SeqCst) + 1; + if call == 1 { + panic!("first dependency construction panicked"); + } + Ok(RequestValue { id: call }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::request_scoped::(|module_ref| { + assert!( + catch_unwind(AssertUnwindSafe(|| module_ref.get::())).is_err() + ); + Ok(PanicRecoveryConsumer { + dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + let consumer = module_ref + .resolve_with_context::(&ContextIdFactory::create()) + .unwrap(); + + assert_eq!(consumer.dependency.id, 2); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn parallel_resolution_branches_do_not_share_sibling_frames() { + let module_ref = ModuleRef::new(); + let factories_ready = Arc::new(Barrier::new(2)); + let consumer_started = Arc::new((Mutex::new(false), Condvar::new())); + let dependency_calls = Arc::new(AtomicUsize::new(0)); + + let dependency_ready = Arc::clone(&factories_ready); + let dependency_signal = Arc::clone(&consumer_started); + let provider_calls = Arc::clone(&dependency_calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + provider_calls.fetch_add(1, Ordering::SeqCst); + dependency_ready.wait(); + let (started, ready) = &*dependency_signal; + let mut started = started.lock().unwrap(); + while !*started { + started = ready.wait(started).unwrap(); + } + drop(started); + thread::sleep(Duration::from_millis(25)); + Ok(ParallelDependency) + }, + )) + .unwrap(); + + let consumer_ready = Arc::clone(&factories_ready); + let consumer_signal = Arc::clone(&consumer_started); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + consumer_ready.wait(); + let (started, ready) = &*consumer_signal; + *started.lock().unwrap() = true; + ready.notify_all(); + Ok(ParallelConsumer { + dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::request_scoped::(|module_ref| { + let consumer_ref = module_ref.clone(); + let dependency_ref = module_ref.clone(); + let (consumer, dependency) = thread::scope(|scope| { + let consumer = scope.spawn(move || consumer_ref.get::()); + let dependency = + scope.spawn(move || dependency_ref.get::()); + (consumer.join().unwrap(), dependency.join().unwrap()) + }); + Ok(ParallelRoot { + consumer: consumer?, + dependency: dependency?, + }) + }) + .with_dependencies([ + ProviderDependency::typed::(), + ProviderDependency::typed::(), + ]), + ) + .unwrap(); + + let root = module_ref + .resolve_with_context::(&ContextIdFactory::create()) + .unwrap(); + + assert!(Arc::ptr_eq(&root.consumer.dependency, &root.dependency)); + assert_eq!(dependency_calls.load(Ordering::SeqCst), 1); +} + +#[derive(Debug)] +struct SelfCycle; + +#[test] +fn context_cache_waiting_does_not_hide_provider_cycles() { + let module_ref = ModuleRef::new(); + module_ref + .register( + ProviderDefinition::transient::(|module_ref| { + let _ = module_ref.get::()?; + Ok(SelfCycle) + }) + .depends_on::(), + ) + .unwrap(); + + let error = module_ref + .resolve_with_context::(&ContextIdFactory::create()) + .unwrap_err(); + + assert!(matches!( + error, + BootError::Internal(message) + if message.contains("cyclic provider dependency detected") + && message.contains("SelfCycle") + )); +} + +#[derive(Debug)] +struct ConcurrentCycleA { + _dependency: Arc, +} + +#[derive(Debug)] +struct ConcurrentCycleB { + _dependency: Arc, +} + +#[test] +fn concurrent_context_roots_detect_cross_thread_dependency_cycles() { + let module_ref = ModuleRef::new(); + let barrier = Arc::new(Barrier::new(2)); + let a_calls = Arc::new(AtomicUsize::new(0)); + let b_calls = Arc::new(AtomicUsize::new(0)); + let a_barrier = Arc::clone(&barrier); + let a_factory_calls = Arc::clone(&a_calls); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + if a_factory_calls.fetch_add(1, Ordering::SeqCst) == 0 { + a_barrier.wait(); + } + Ok(ConcurrentCycleA { + _dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + let b_barrier = Arc::clone(&barrier); + let b_factory_calls = Arc::clone(&b_calls); + module_ref + .register( + ProviderDefinition::request_scoped::(move |module_ref| { + if b_factory_calls.fetch_add(1, Ordering::SeqCst) == 0 { + b_barrier.wait(); + } + Ok(ConcurrentCycleB { + _dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + let context_id = ContextIdFactory::create(); + let (sender, receiver) = mpsc::channel(); + let first_ref = module_ref.clone(); + let first_context = context_id.clone(); + let first_sender = sender.clone(); + thread::spawn(move || { + let result = first_ref + .resolve_with_context::(&first_context) + .map(|_| ()); + first_sender.send(result).unwrap(); + }); + let second_ref = module_ref.clone(); + thread::spawn(move || { + let result = second_ref + .resolve_with_context::(&context_id) + .map(|_| ()); + sender.send(result).unwrap(); + }); + + let first = receiver.recv_timeout(Duration::from_secs(2)).unwrap(); + let second = receiver.recv_timeout(Duration::from_secs(2)).unwrap(); + + for result in [first, second] { + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message.contains("cyclic") && message.contains("provider dependency") + )); + } +} + +#[derive(Debug)] +struct ContextRouteModule { + calls: Arc, +} + +impl Module for ContextRouteModule { + fn name(&self) -> &'static str { + "context-route" + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ProviderDefinition::request_scoped::( + move |_| { + Ok(RequestValue { + id: calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )]) + } + + fn routes(&self) -> Result> { + Ok(vec![RouteDefinition::get( + "/context-id", + |request: BootRequest| async move { + let context_id = request.context_id().ok_or_else(|| { + BootError::Internal("route request is missing a context id".to_string()) + })?; + let discovered = ContextIdFactory::get_by_request(&request); + if &discovered != context_id { + return Err(BootError::Internal( + "request context discovery returned a different id".to_string(), + )); + } + let first = request.get::()?; + let second = request.get::()?; + Ok(BootResponse::text(format!( + "{}:{}:{}", + context_id.id(), + first.id, + Arc::ptr_eq(&first, &second) + ))) + }, + )?]) + } +} + +#[tokio::test] +async fn http_routes_expose_one_context_id_per_request() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(ContextRouteModule { + calls: Arc::clone(&calls), + }) + .build() + .unwrap(); + + let first = app + .call(BootRequest::new(HttpMethod::Get, "/context-id")) + .await + .unwrap(); + let second = app + .call(BootRequest::new(HttpMethod::Get, "/context-id")) + .await + .unwrap(); + let first_body = first.body_text().unwrap(); + let second_body = second.body_text().unwrap(); + let first_parts = first_body.split(':').collect::>(); + let second_parts = second_body.split(':').collect::>(); + + assert_ne!(first_parts[0], second_parts[0]); + assert_eq!(&first_parts[1..], ["1", "true"]); + assert_eq!(&second_parts[1..], ["2", "true"]); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} diff --git a/tests/providers.rs b/tests/providers.rs index dde7a7f..a08ef5a 100644 --- a/tests/providers.rs +++ b/tests/providers.rs @@ -1,7 +1,8 @@ use a3s_boot::{ BootApplication, BootError, BootFactory, BootRequest, BootResponse, ControllerDefinition, DynamicModule, ExecutionContext, FromModuleRef, HttpMethod, Module, ModuleRef, - ProviderDefinition, ProviderRef, ProviderScope, ProviderToken, Result, TestingModule, + ProviderDefinition, ProviderDependency, ProviderOnModuleInit, ProviderRef, ProviderScope, + ProviderToken, Result, TestingModule, }; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; @@ -274,6 +275,83 @@ struct ScopedConsumer { second: Arc, } +#[derive(Debug)] +struct BubbledConsumer { + counter: Arc, +} + +#[derive(Debug)] +struct ContextualTransient { + counter: Arc, +} + +#[derive(Debug)] +struct TransitiveBubbledConsumer { + dependency: Arc, +} + +#[derive(Debug)] +struct LazyCounterConsumer { + counter: ProviderRef, +} + +#[derive(Debug)] +struct NamedBubbledConsumer { + counter: Arc, +} + +#[derive(Debug)] +struct OptionalExistingConsumer { + counter: Option>, +} + +#[derive(Debug)] +struct OptionalMissingConsumer { + counter: Option>, +} + +#[derive(Debug)] +struct MissingLazyConsumer { + _counter: ProviderRef, +} + +#[derive(Debug)] +struct LifecycleBubbledConsumer { + _counter: Arc, +} + +impl ProviderOnModuleInit for LifecycleBubbledConsumer {} + +#[derive(Debug)] +struct AsyncContextualConsumer { + _counter: Arc, +} + +#[derive(Debug)] +struct AsyncDependency { + value: &'static str, +} + +#[derive(Debug)] +struct AsyncDependent { + dependency: Arc, +} + +#[derive(Debug)] +struct AsyncCycleA { + _dependency: Arc, +} + +#[derive(Debug)] +struct AsyncCycleB { + _dependency: Arc, +} + +#[derive(Debug)] +struct ForwardContextConsumer { + counter: Arc, +} + #[derive(Debug)] struct ArcFactoryModule; @@ -430,6 +508,119 @@ impl Module for RequestScopedDependencyModule { } } +#[derive(Debug)] +struct BubbledProviderModule { + calls: Arc, +} + +impl Module for BubbledProviderModule { + fn name(&self) -> &'static str { + "bubbled-provider" + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + Ok(ScopedCounter { + id: calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + ProviderDefinition::factory::(|module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ]) + } + + fn controllers(&self, _module_ref: &ModuleRef) -> Result> { + Ok(vec![ControllerDefinition::new("/bubbled")?.get( + "/", + |request: BootRequest| async move { + let consumer = request.get::()?; + let same = request.get::()?; + let direct = request.get::()?; + Ok(BootResponse::text(format!( + "{}:{}:{}:{}", + consumer.counter.id, + same.counter.id, + direct.id, + Arc::ptr_eq(&consumer, &same) + ))) + }, + )?]) + } +} + +#[derive(Debug)] +struct ForwardContextRootModule { + calls: Arc, +} + +impl Module for ForwardContextRootModule { + fn name(&self) -> &'static str { + "forward-context-root" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(ForwardContextFeatureModule { + calls: Arc::clone(&self.calls), + })] + } + + fn providers(&self) -> Result> { + let calls = Arc::clone(&self.calls); + Ok(vec![ + ProviderDefinition::request_scoped::(move |_| { + Ok(ScopedCounter { + id: calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }), + ]) + } + + fn exports(&self) -> Result> { + Ok(vec![ + ProviderToken::of::(), + ProviderToken::of::(), + ]) + } +} + +#[derive(Debug)] +struct ForwardContextFeatureModule { + calls: Arc, +} + +impl Module for ForwardContextFeatureModule { + fn name(&self) -> &'static str { + "forward-context-feature" + } + + fn forward_imports(&self) -> Vec> { + vec![Arc::new(ForwardContextRootModule { + calls: Arc::clone(&self.calls), + })] + } + + fn providers(&self) -> Result> { + Ok(vec![ + ProviderDefinition::factory::(|module_ref| { + Ok(ForwardContextConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } +} + #[derive(Debug)] struct RequestScopedControllerModule { calls: Arc, @@ -796,6 +987,33 @@ impl Module for UsesRuntimeConfigModule { } } +#[derive(Debug)] +struct AsyncGlobalRuntimeConfigModule; + +impl Module for AsyncGlobalRuntimeConfigModule { + fn name(&self) -> &'static str { + "async-global-runtime-config" + } + + fn providers(&self) -> Result> { + Ok(vec![ + ProviderDefinition::async_factory::(|_| async { + Ok(RuntimeConfig { + value: "async-global".to_string(), + }) + }), + ]) + } + + fn exports(&self) -> Result> { + Ok(vec![ProviderToken::of::()]) + } + + fn is_global(&self) -> bool { + true + } +} + #[derive(Debug)] struct AutoProviderModule; @@ -945,6 +1163,506 @@ fn module_ref_resolve_uses_a_fresh_request_resolution_context() { assert_eq!(calls.load(Ordering::SeqCst), 2); } +#[test] +fn request_scope_propagates_to_declared_singleton_dependencies() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let counter_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: counter_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + + assert!(module_ref + .provider_is_contextual::() + .unwrap()); + let error = module_ref.get::().unwrap_err(); + assert!(matches!( + error, + BootError::Internal(message) + if message.contains("requires an active request scope") + )); + + let first_scope = module_ref.request_scope(); + let first = first_scope.get::().unwrap(); + let same = first_scope.get::().unwrap(); + let direct = first_scope.get::().unwrap(); + let second = module_ref.request_scope().get::().unwrap(); + + assert!(Arc::ptr_eq(&first, &same)); + assert!(Arc::ptr_eq(&first.counter, &direct)); + assert_eq!(first.counter.id, 1); + assert_eq!(second.counter.id, 2); + assert!(!Arc::ptr_eq(&first, &second)); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn request_scope_propagates_through_transient_and_alias_dependencies() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let counter_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: counter_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::transient::(|module_ref| { + Ok(ContextualTransient { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(TransitiveBubbledConsumer { + dependency: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .unwrap(); + module_ref + .register(ProviderDefinition::named_alias( + "request-counter", + ProviderToken::of::(), + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::named_factory::( + "alias-consumer", + |module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get_named::("request-counter")?, + }) + }, + ) + .depends_on_named("request-counter"), + ) + .unwrap(); + + assert!(module_ref + .provider_is_contextual::() + .unwrap()); + assert!(module_ref + .provider_is_contextual::() + .unwrap()); + assert!(module_ref + .named_provider_is_contextual("alias-consumer") + .unwrap()); + + let scope = module_ref.request_scope(); + let transitive = scope.get::().unwrap(); + let alias = scope + .get_named::("alias-consumer") + .unwrap(); + assert!(Arc::ptr_eq(&transitive.dependency.counter, &alias.counter)); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn lazy_dependencies_do_not_bubble_but_require_an_explicit_resolution_context() { + let module_ref = ModuleRef::new(); + module_ref + .register(ProviderDefinition::request_scoped::( + |_| Ok(ScopedCounter { id: 1 }), + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(LazyCounterConsumer { + counter: module_ref.provider_ref::(), + }) + }) + .with_dependency(ProviderDependency::typed::().lazy()), + ) + .unwrap(); + + assert!(!module_ref + .provider_is_contextual::() + .unwrap()); + let consumer = module_ref.get::().unwrap(); + assert!(consumer.counter.get().is_err()); + assert_eq!(consumer.counter.resolve().unwrap().id, 1); +} + +#[test] +fn opaque_singleton_factories_cannot_capture_request_scoped_providers() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + module_ref + .register(ProviderDefinition::request_scoped::( + |_| Ok(ScopedCounter { id: 1 }), + )) + .unwrap(); + let factory_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::factory::( + move |module_ref| { + factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }, + )) + .unwrap(); + + let error = module_ref.get::().unwrap_err(); + assert!(matches!( + error, + BootError::Internal(message) + if message.contains("BubbledConsumer") + && message.contains("ScopedCounter") + && message.contains(" -> ") + && message.contains("declare factory dependencies") + )); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn dependency_plans_validate_required_contextual_and_lazy_edges() { + let contextual_ref = ModuleRef::new(); + contextual_ref + .register( + ProviderDefinition::named_request_scoped::( + "request-with-missing-dependency", + |_| Ok(ItemsService), + ) + .depends_on_named("missing-eager-dependency"), + ) + .unwrap(); + + let contextual_error = contextual_ref + .named_provider_is_contextual("request-with-missing-dependency") + .unwrap_err(); + assert!(matches!( + contextual_error, + BootError::MissingProvider(token) if token == "missing-eager-dependency" + )); + + let lazy_ref = ModuleRef::new(); + lazy_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(MissingLazyConsumer { + _counter: module_ref.provider_ref::(), + }) + }) + .with_dependency(ProviderDependency::typed::().lazy()), + ) + .unwrap(); + + let lazy_error = lazy_ref + .provider_is_contextual::() + .unwrap_err(); + assert!(matches!( + lazy_error, + BootError::MissingProvider(token) + if token == ProviderToken::of::().to_string() + )); +} + +#[test] +fn optional_dependencies_bubble_only_when_the_provider_exists() { + let calls = Arc::new(AtomicUsize::new(0)); + let module_ref = ModuleRef::new(); + let counter_calls = Arc::clone(&calls); + module_ref + .register(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: counter_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(OptionalExistingConsumer { + counter: module_ref.get_optional::()?, + }) + }) + .with_dependency(ProviderDependency::typed::().optional()), + ) + .unwrap(); + module_ref + .register( + ProviderDefinition::factory::(|module_ref| { + Ok(OptionalMissingConsumer { + counter: module_ref.get_optional_named::("missing-counter")?, + }) + }) + .with_dependency(ProviderDependency::named("missing-counter").optional()), + ) + .unwrap(); + + assert!(module_ref + .provider_is_contextual::() + .unwrap()); + assert!(!module_ref + .provider_is_contextual::() + .unwrap()); + assert!(module_ref + .get::() + .unwrap() + .counter + .is_none()); + + let scope = module_ref.request_scope(); + let consumer = scope.get::().unwrap(); + let direct = scope.get::().unwrap(); + assert!(Arc::ptr_eq(consumer.counter.as_ref().unwrap(), &direct)); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn contextual_plans_use_the_declaring_module_for_imports_and_named_aliases() { + let calls = Arc::new(AtomicUsize::new(0)); + let request_calls = Arc::clone(&calls); + let child = DynamicModule::new("contextual-owner-child") + .provider(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: request_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .provider(ProviderDefinition::named_alias( + "contextual-owner-alias", + ProviderToken::of::(), + )) + .export::() + .export_named("contextual-owner-alias"); + let parent = DynamicModule::new("contextual-owner-parent") + .import(child) + .provider( + ProviderDefinition::factory::(|module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ) + .provider( + ProviderDefinition::factory::(|module_ref| { + Ok(NamedBubbledConsumer { + counter: module_ref.get_named::("contextual-owner-alias")?, + }) + }) + .depends_on_named("contextual-owner-alias"), + ); + let sibling = DynamicModule::new("contextual-owner-sibling") + .provider(ProviderDefinition::singleton(ScopedCounter { id: 99 })); + + let app = BootApplication::builder() + .import(sibling) + .import(parent) + .build() + .unwrap(); + + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert_eq!(app.get::().unwrap().id, 99); + + let scope = app.module_ref().request_scope(); + let typed = scope.get::().unwrap(); + let named = scope.get::().unwrap(); + assert_eq!(typed.counter.id, 1); + assert!(Arc::ptr_eq(&typed.counter, &named.counter)); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn request_scope_bubbles_from_a_late_global_provider() { + let calls = Arc::new(AtomicUsize::new(0)); + let consumer = DynamicModule::new("global-context-consumer").provider( + ProviderDefinition::factory::(|module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ); + let request_calls = Arc::clone(&calls); + let global = DynamicModule::new("global-context-provider") + .provider(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: request_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .export::() + .global(); + + let app = BootApplication::builder() + .import(consumer) + .import(global) + .build() + .unwrap(); + + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + let first = app + .module_ref() + .request_scope() + .get::() + .unwrap(); + let second = app + .module_ref() + .request_scope() + .get::() + .unwrap(); + assert_eq!(first.counter.id, 1); + assert_eq!(second.counter.id, 2); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn request_scope_bubbles_across_forward_imports_after_graph_registration() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(ForwardContextRootModule { + calls: Arc::clone(&calls), + }) + .build() + .unwrap(); + + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + let first_scope = app.module_ref().request_scope(); + let first = first_scope.get::().unwrap(); + let direct = first_scope.get::().unwrap(); + let second = app + .module_ref() + .request_scope() + .get::() + .unwrap(); + + assert!(Arc::ptr_eq(&first.counter, &direct)); + assert_eq!(first.counter.id, 1); + assert_eq!(second.counter.id, 2); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn provider_overrides_recalculate_contextual_dependency_plans() { + let calls = Arc::new(AtomicUsize::new(0)); + let module = DynamicModule::new("contextual-provider-override") + .provider(ProviderDefinition::singleton(ScopedCounter { id: 0 })) + .provider( + ProviderDefinition::factory::(|module_ref| { + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ); + let request_calls = Arc::clone(&calls); + let app = BootApplication::builder() + .import(module) + .override_provider(ProviderDefinition::request_scoped::( + move |_| { + Ok(ScopedCounter { + id: request_calls.fetch_add(1, Ordering::SeqCst) + 1, + }) + }, + )) + .build() + .unwrap(); + + assert!(app + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert!(app.get::().is_err()); + let first = app + .module_ref() + .request_scope() + .get::() + .unwrap(); + let second = app + .module_ref() + .request_scope() + .get::() + .unwrap(); + assert_eq!(first.counter.id, 1); + assert_eq!(second.counter.id, 2); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn declared_static_singleton_graphs_are_initialized_only_once() { + let dependency_calls = Arc::new(AtomicUsize::new(0)); + let consumer_calls = Arc::new(AtomicUsize::new(0)); + let dependency_factory_calls = Arc::clone(&dependency_calls); + let consumer_factory_calls = Arc::clone(&consumer_calls); + let module = DynamicModule::new("declared-static-singletons") + .provider(ProviderDefinition::factory::(move |_| { + dependency_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(ScopedCounter { id: 7 }) + })) + .provider( + ProviderDefinition::factory::(move |module_ref| { + consumer_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(BubbledConsumer { + counter: module_ref.get::()?, + }) + }) + .depends_on::(), + ); + + let app = BootApplication::builder().import(module).build().unwrap(); + let first = app.get::().unwrap(); + let second = app.get::().unwrap(); + + assert!(!app + .module_ref() + .provider_is_contextual::() + .unwrap()); + assert!(Arc::ptr_eq(&first, &second)); + assert!(Arc::ptr_eq(&first.counter, &second.counter)); + assert_eq!(dependency_calls.load(Ordering::SeqCst), 1); + assert_eq!(consumer_calls.load(Ordering::SeqCst), 1); +} + #[test] fn module_ref_resolve_supports_named_and_optional_lookup() { let module_ref = ModuleRef::new(); @@ -1234,6 +1952,31 @@ async fn request_scoped_provider_dependencies_share_the_request_scope() { assert_eq!(calls.load(Ordering::SeqCst), 2); } +#[tokio::test] +async fn request_scope_bubbles_into_declared_singletons_per_request() { + let calls = Arc::new(AtomicUsize::new(0)); + let app = BootApplication::builder() + .import(BubbledProviderModule { + calls: Arc::clone(&calls), + }) + .build() + .unwrap(); + + assert!(app.get::().is_err()); + let first = app + .call(BootRequest::new(HttpMethod::Get, "/bubbled")) + .await + .unwrap(); + let second = app + .call(BootRequest::new(HttpMethod::Get, "/bubbled")) + .await + .unwrap(); + + assert_eq!(first.body_text().unwrap(), "1:1:1:true"); + assert_eq!(second.body_text().unwrap(), "2:2:2:true"); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + #[tokio::test] async fn scoped_controller_handlers_are_built_for_each_request_scope() { let calls = Arc::new(AtomicUsize::new(0)); @@ -1320,6 +2063,190 @@ fn singleton_provider_factories_can_depend_on_later_module_providers() { assert!(Arc::ptr_eq(&repository.config, &config)); } +#[test] +fn contextual_singletons_with_lifecycle_hooks_are_rejected_before_construction() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let provider_factory_calls = Arc::clone(&factory_calls); + let module = DynamicModule::new("contextual-lifecycle-provider") + .provider(ProviderDefinition::request_scoped::( + |_| Ok(ScopedCounter { id: 1 }), + )) + .provider( + ProviderDefinition::factory::(move |module_ref| { + provider_factory_calls.fetch_add(1, Ordering::SeqCst); + Ok(LifecycleBubbledConsumer { + _counter: module_ref.get::()?, + }) + }) + .depends_on::() + .with_on_module_init::(), + ); + + let result = BootApplication::builder().import(module).build(); + + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message.contains("LifecycleBubbledConsumer") + && message.contains("singleton lifecycle hooks") + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn contextual_async_providers_are_rejected_before_factory_invocation() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let provider_factory_calls = Arc::clone(&factory_calls); + let module = DynamicModule::new("contextual-async-provider") + .provider(ProviderDefinition::request_scoped::( + |_| Ok(ScopedCounter { id: 1 }), + )) + .provider( + ProviderDefinition::async_factory::(move |module_ref| { + provider_factory_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok(AsyncContextualConsumer { + _counter: module_ref.get::()?, + }) + } + }) + .depends_on::(), + ); + + let result = BootApplication::builder() + .import(module) + .build_async() + .await; + + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message.contains("AsyncContextualConsumer") + && message.contains("cannot depend on a request-context provider") + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn async_singletons_seed_declared_dependencies_before_their_dependents() { + let dependency_calls = Arc::new(AtomicUsize::new(0)); + let dependent_calls = Arc::new(AtomicUsize::new(0)); + let dependency_factory_calls = Arc::clone(&dependency_calls); + let dependent_factory_calls = Arc::clone(&dependent_calls); + let module = DynamicModule::new("async-provider-topology") + .provider( + ProviderDefinition::async_factory::(move |module_ref| { + dependent_factory_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok(AsyncDependent { + dependency: module_ref.get::()?, + }) + } + }) + .depends_on::(), + ) + .provider(ProviderDefinition::async_factory::( + move |_| { + dependency_factory_calls.fetch_add(1, Ordering::SeqCst); + async { + Ok(AsyncDependency { + value: "dependency-first", + }) + } + }, + )); + + let app = BootApplication::builder() + .import(module) + .build_async() + .await + .unwrap(); + let dependent = app.get::().unwrap(); + let dependency = app.get::().unwrap(); + + assert_eq!(dependent.dependency.value, "dependency-first"); + assert!(Arc::ptr_eq(&dependent.dependency, &dependency)); + assert_eq!(dependency_calls.load(Ordering::SeqCst), 1); + assert_eq!(dependent_calls.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn async_dependency_cycles_are_rejected_without_running_factories() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let first_calls = Arc::clone(&factory_calls); + let second_calls = Arc::clone(&factory_calls); + let module = DynamicModule::new("async-provider-cycle") + .provider( + ProviderDefinition::named_async_factory::( + "async-cycle-a", + move |module_ref| { + first_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok(AsyncCycleA { + _dependency: module_ref.get_named::("async-cycle-b")?, + }) + } + }, + ) + .depends_on_named("async-cycle-b"), + ) + .provider( + ProviderDefinition::named_async_factory::( + "async-cycle-b", + move |module_ref| { + second_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok(AsyncCycleB { + _dependency: module_ref.get_named::("async-cycle-a")?, + }) + } + }, + ) + .depends_on_named("async-cycle-a"), + ); + + let result = BootApplication::builder() + .import(module) + .build_async() + .await; + + assert!(matches!( + result, + Err(BootError::Internal(message)) + if message + == "cyclic async provider dependency detected: async-cycle-a -> async-cycle-b -> async-cycle-a" + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn missing_declared_async_dependencies_fail_before_factory_invocation() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let provider_factory_calls = Arc::clone(&factory_calls); + let module = DynamicModule::new("async-provider-missing-dependency").provider( + ProviderDefinition::async_factory::(move |_| { + provider_factory_calls.fetch_add(1, Ordering::SeqCst); + async { + Ok(AsyncDependency { + value: "unreachable", + }) + } + }) + .depends_on_named("missing-async-dependency"), + ); + + let result = BootApplication::builder() + .import(module) + .build_async() + .await; + + assert!(matches!( + result, + Err(BootError::MissingProvider(token)) if token == "missing-async-dependency" + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} + #[tokio::test] async fn async_provider_factories_are_awaited_before_controllers_build() { let calls = Arc::new(AtomicUsize::new(0)); @@ -1495,8 +2422,8 @@ fn imported_exports_can_be_re_exported_transitively() { #[test] fn global_modules_expose_exported_providers_to_other_modules() { let app = BootApplication::builder() - .import(GlobalConfigModule) .import(UsesGlobalConfigModule) + .import(GlobalConfigModule) .build() .unwrap(); @@ -1507,6 +2434,22 @@ fn global_modules_expose_exported_providers_to_other_modules() { assert_eq!(dependent.config.value, "global"); } +#[tokio::test] +async fn async_global_modules_are_seeded_after_the_full_graph_is_registered() { + let app = BootApplication::builder() + .import(UsesRuntimeConfigModule) + .import(AsyncGlobalRuntimeConfigModule) + .build_async() + .await + .unwrap(); + + let config = app.get::().unwrap(); + let dependent = app.get::().unwrap(); + + assert_eq!(config.value, "async-global"); + assert!(Arc::ptr_eq(&config, &dependent.config)); +} + #[test] fn dynamic_modules_can_provide_exported_runtime_configuration() { let dynamic_config = DynamicModule::new("runtime-config") diff --git a/tests/transport.rs b/tests/transport.rs index 0a90669..7fa9bde 100644 --- a/tests/transport.rs +++ b/tests/transport.rs @@ -1,9 +1,9 @@ use a3s_boot::{ - BootApplication, BootError, BootErrorKind, BoxFuture, ExecutionContext, ExecutionInterceptor, - ExecutionProtocol, ExecutionTransportKind, Guard, InProcessTransport, MessagePatternDefinition, - MessageTransport, Module, ModuleRef, ProviderDefinition, Result, TransportContext, - TransportExceptionResponse, TransportInterceptor, TransportMessage, TransportReply, Validate, - ValidationOptions, ValidationSchema, + BootApplication, BootError, BootErrorKind, BoxFuture, CallHandler, ExecutionContext, + ExecutionInterceptor, ExecutionProtocol, ExecutionTransportKind, Guard, InProcessTransport, + MessagePatternDefinition, MessageTransport, Module, ModuleRef, ProviderDefinition, Result, + TransportContext, TransportExceptionResponse, TransportInterceptor, TransportMessage, + TransportReply, Validate, ValidationOptions, ValidationSchema, }; #[cfg(feature = "grpc-transport")] use a3s_boot::{GrpcTransport, GrpcTransportClient, GrpcTransportOptions}; @@ -21,6 +21,7 @@ use a3s_boot::{RedisTransport, RedisTransportClient, RedisTransportOptions}; use a3s_boot::{TcpTransport, TcpTransportClient}; use serde::{Deserialize, Serialize}; use serde_json::json; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; #[cfg(any( feature = "grpc-transport", @@ -440,6 +441,147 @@ async fn transport_pipeline_runs_in_order() { ); } +struct RecoveringTransportInterceptor; + +impl TransportInterceptor for RecoveringTransportInterceptor { + fn intercept<'a>( + &'a self, + _context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + match next.handle().await { + Ok(reply) => Ok(reply), + Err(error) => Ok(Some(TransportReply::text(format!("recovered: {error}")))), + } + }) + } +} + +#[tokio::test] +async fn transport_interceptors_can_recover_errors_before_filters() { + let filter_calls = Arc::new(AtomicUsize::new(0)); + let filter_log = Arc::clone(&filter_calls); + let pattern = MessagePatternDefinition::request("cats.recover", |_message| async { + Err::(BootError::BadRequest("invalid cat payload".to_string())) + }) + .unwrap() + .with_interceptor(RecoveringTransportInterceptor) + .with_filter(move |_context: TransportContext, _error: BootError| { + let filter_log = Arc::clone(&filter_log); + async move { + filter_log.fetch_add(1, Ordering::SeqCst); + Ok(Some(TransportExceptionResponse::reply( + TransportReply::text("filtered"), + ))) + } + }); + + let reply = pattern + .dispatch(TransportMessage::new("cats.recover", json!({ "id": 1 }))) + .await + .unwrap() + .unwrap(); + + assert_eq!( + reply.data(), + &json!("recovered: bad request: invalid cat payload") + ); + assert_eq!(filter_calls.load(Ordering::SeqCst), 0); +} + +struct RetryUnavailableTransportInterceptor; + +impl TransportInterceptor for RetryUnavailableTransportInterceptor { + fn intercept<'a>( + &'a self, + _context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + match next.handle().await { + Err(BootError::ServiceUnavailable(_)) => next.handle().await, + result => result, + } + }) + } +} + +#[tokio::test] +async fn transport_call_handlers_can_replay_the_downstream_pipeline() { + let pipe_calls = Arc::new(AtomicUsize::new(0)); + let handler_calls = Arc::new(AtomicUsize::new(0)); + let pipe_log = Arc::clone(&pipe_calls); + let handler_log = Arc::clone(&handler_calls); + let pattern = MessagePatternDefinition::request("cats.retry", move |message| { + let handler_log = Arc::clone(&handler_log); + async move { + let attempt = handler_log.fetch_add(1, Ordering::SeqCst) + 1; + if attempt == 1 { + return Err(BootError::ServiceUnavailable( + "temporary failure".to_string(), + )); + } + Ok(TransportReply::new(message.data)) + } + }) + .unwrap() + .with_interceptor(RetryUnavailableTransportInterceptor) + .with_pipe(move |mut message: TransportMessage| { + let pipe_log = Arc::clone(&pipe_log); + async move { + let attempt = pipe_log.fetch_add(1, Ordering::SeqCst) + 1; + message.data = json!({ "pipeAttempt": attempt }); + Ok(message) + } + }); + + let reply = pattern + .dispatch(TransportMessage::new("cats.retry", json!({ "id": 1 }))) + .await + .unwrap() + .unwrap(); + + assert_eq!(reply.data(), &json!({ "pipeAttempt": 2 })); + assert_eq!(pipe_calls.load(Ordering::SeqCst), 2); + assert_eq!(handler_calls.load(Ordering::SeqCst), 2); +} + +struct ShortCircuitTransportInterceptor; + +impl TransportInterceptor for ShortCircuitTransportInterceptor { + fn intercept<'a>( + &'a self, + _context: TransportContext, + _next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async { Ok(Some(TransportReply::text("ignored event reply"))) }) + } +} + +#[tokio::test] +async fn event_patterns_discard_short_circuit_interceptor_replies() { + let handler_calls = Arc::new(AtomicUsize::new(0)); + let handler_log = Arc::clone(&handler_calls); + let pattern = MessagePatternDefinition::event("cats.created", move |_message| { + let handler_log = Arc::clone(&handler_log); + async move { + handler_log.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + }) + .unwrap() + .with_interceptor(ShortCircuitTransportInterceptor); + + let reply = pattern + .dispatch(TransportMessage::new("cats.created", json!({ "id": 1 }))) + .await + .unwrap(); + + assert_eq!(reply, None); + assert_eq!(handler_calls.load(Ordering::SeqCst), 0); +} + #[tokio::test] async fn global_transport_pipes_apply_before_pattern_pipes() { let log = Arc::new(Mutex::new(Vec::new())); diff --git a/tests/transport_contextual_handlers.rs b/tests/transport_contextual_handlers.rs new file mode 100644 index 0000000..e7d57fa --- /dev/null +++ b/tests/transport_contextual_handlers.rs @@ -0,0 +1,239 @@ +use a3s_boot::{ + BootApplication, BootError, BoxFuture, CallHandler, MessagePatternDefinition, Module, + ModuleRef, ProviderDefinition, Result, TransportContext, TransportInterceptor, + TransportMessage, TransportReply, +}; +use serde_json::json; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +#[derive(Debug)] +struct ScopedMessageDependency { + id: usize, + context_id: u64, +} + +#[derive(Debug)] +struct RetryOnce; + +impl TransportInterceptor for RetryOnce { + fn intercept<'a>( + &'a self, + context: TransportContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + let _ = context; + match next.handle().await { + Err(BootError::ServiceUnavailable(_)) => next.handle().await, + result => result, + } + }) + } +} + +#[derive(Debug)] +struct ContextualPatternModule { + dependency_calls: Arc, + request_factory_calls: Arc, + request_events: Arc>>, + event_factory_calls: Arc, + event_events: Arc>>, +} + +impl Module for ContextualPatternModule { + fn name(&self) -> &'static str { + "contextual-message-patterns" + } + + fn providers(&self) -> Result> { + let dependency_calls = Arc::clone(&self.dependency_calls); + Ok(vec![ProviderDefinition::request_scoped::< + ScopedMessageDependency, + _, + >(move |module_ref| { + let context_id = module_ref + .context_id() + .ok_or_else(|| { + BootError::Internal( + "scoped message dependency was built without a ContextId".to_string(), + ) + })? + .id(); + Ok(ScopedMessageDependency { + id: dependency_calls.fetch_add(1, Ordering::SeqCst) + 1, + context_id, + }) + })]) + } + + fn message_patterns(&self, _module_ref: &ModuleRef) -> Result> { + let request_factory_calls = Arc::clone(&self.request_factory_calls); + let request_events = Arc::clone(&self.request_events); + let request = + MessagePatternDefinition::request_scoped("contextual.request", move |module_ref| { + request_factory_calls.fetch_add(1, Ordering::SeqCst); + let dependency = module_ref.get::()?; + let attempts = Arc::new(AtomicUsize::new(0)); + let request_events = Arc::clone(&request_events); + Ok(move |_message: TransportMessage| { + let dependency = Arc::clone(&dependency); + let attempts = Arc::clone(&attempts); + let request_events = Arc::clone(&request_events); + async move { + let attempt = attempts.fetch_add(1, Ordering::SeqCst) + 1; + request_events.lock().unwrap().push(( + dependency.id, + dependency.context_id, + attempt, + )); + if attempt == 1 { + return Err(BootError::ServiceUnavailable( + "retry contextual handler".to_string(), + )); + } + Ok(TransportReply::text(format!( + "{}:{}:{attempt}", + dependency.id, dependency.context_id + ))) + } + }) + })? + .with_interceptor(RetryOnce); + + let event_factory_calls = Arc::clone(&self.event_factory_calls); + let event_events = Arc::clone(&self.event_events); + let event = + MessagePatternDefinition::event_scoped("contextual.event", move |module_ref| { + event_factory_calls.fetch_add(1, Ordering::SeqCst); + let dependency = module_ref.get::()?; + let event_events = Arc::clone(&event_events); + Ok(move |_message: TransportMessage| { + let dependency = Arc::clone(&dependency); + let event_events = Arc::clone(&event_events); + async move { + event_events + .lock() + .unwrap() + .push((dependency.id, dependency.context_id)); + Ok(()) + } + }) + })?; + + Ok(vec![request, event]) + } +} + +struct ContextualPatternHarness { + app: BootApplication, + dependency_calls: Arc, + request_factory_calls: Arc, + request_events: Arc>>, + event_factory_calls: Arc, + event_events: Arc>>, +} + +fn contextual_pattern_app() -> ContextualPatternHarness { + let dependency_calls = Arc::new(AtomicUsize::new(0)); + let request_factory_calls = Arc::new(AtomicUsize::new(0)); + let request_events = Arc::new(Mutex::new(Vec::new())); + let event_factory_calls = Arc::new(AtomicUsize::new(0)); + let event_events = Arc::new(Mutex::new(Vec::new())); + let app = BootApplication::builder() + .import(ContextualPatternModule { + dependency_calls: Arc::clone(&dependency_calls), + request_factory_calls: Arc::clone(&request_factory_calls), + request_events: Arc::clone(&request_events), + event_factory_calls: Arc::clone(&event_factory_calls), + event_events: Arc::clone(&event_events), + }) + .build() + .unwrap(); + ContextualPatternHarness { + app, + dependency_calls, + request_factory_calls, + request_events, + event_factory_calls, + event_events, + } +} + +#[tokio::test] +async fn request_scoped_pattern_reuses_one_handler_and_context_across_retry() { + let harness = contextual_pattern_app(); + + let first = harness + .app + .dispatch_message(TransportMessage::new("contextual.request", json!({}))) + .await + .unwrap() + .unwrap(); + let second = harness + .app + .dispatch_message(TransportMessage::new("contextual.request", json!({}))) + .await + .unwrap() + .unwrap(); + + let events = harness.request_events.lock().unwrap(); + assert_eq!(events.len(), 4); + assert_eq!(events[0], (1, events[0].1, 1)); + assert_eq!(events[1], (1, events[0].1, 2)); + assert_eq!(events[2], (2, events[2].1, 1)); + assert_eq!(events[3], (2, events[2].1, 2)); + assert_ne!(events[0].1, events[2].1); + assert_eq!(first, TransportReply::text(format!("1:{}:2", events[0].1))); + assert_eq!(second, TransportReply::text(format!("2:{}:2", events[2].1))); + assert_eq!(harness.request_factory_calls.load(Ordering::SeqCst), 2); + assert_eq!(harness.dependency_calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn event_scoped_pattern_uses_a_fresh_private_dependency_per_message() { + let harness = contextual_pattern_app(); + + harness + .app + .emit_message(TransportMessage::new("contextual.event", json!({}))) + .await + .unwrap(); + harness + .app + .emit_message(TransportMessage::new("contextual.event", json!({}))) + .await + .unwrap(); + + let events = harness.event_events.lock().unwrap(); + assert_eq!(events.len(), 2); + assert_eq!(events[0].0, 1); + assert_eq!(events[1].0, 2); + assert_ne!(events[0].1, events[1].1); + assert_eq!(harness.event_factory_calls.load(Ordering::SeqCst), 2); + assert_eq!(harness.dependency_calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn standalone_scoped_pattern_requires_a_module_context_before_factory_runs() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let observed_calls = Arc::clone(&factory_calls); + let pattern = + MessagePatternDefinition::request_scoped("standalone.contextual", move |_module_ref| { + observed_calls.fetch_add(1, Ordering::SeqCst); + Ok(|_message: TransportMessage| async { Ok(TransportReply::text("unexpected")) }) + }) + .unwrap(); + + let error = pattern + .dispatch(TransportMessage::new("standalone.contextual", json!({}))) + .await + .unwrap_err(); + + assert!(matches!( + error, + BootError::Internal(message) + if message.contains("scoped") && message.contains("module context") + )); + assert_eq!(factory_calls.load(Ordering::SeqCst), 0); +} diff --git a/tests/websocket.rs b/tests/websocket.rs index 5a39dcc..5937bbe 100644 --- a/tests/websocket.rs +++ b/tests/websocket.rs @@ -1,13 +1,14 @@ use a3s_boot::{ - BootApplication, BootError, BootErrorKind, BootRequest, BoxFuture, ExecutionContext, - ExecutionProtocol, Guard, HttpMethod, Module, ModuleRef, ProviderDefinition, Result, Validate, - ValidationOptions, ValidationSchema, WebSocketContext, WebSocketExceptionResponse, - WebSocketGatewayConnection, WebSocketGatewayDefinition, WebSocketGatewayInitContext, - WebSocketGatewayServer, WebSocketInterceptor, WebSocketMessage, + BootApplication, BootError, BootErrorKind, BootRequest, BoxFuture, CallHandler, + ExecutionContext, ExecutionProtocol, Guard, HttpMethod, Module, ModuleRef, ProviderDefinition, + Result, Validate, ValidationOptions, ValidationSchema, WebSocketContext, + WebSocketExceptionResponse, WebSocketGatewayConnection, WebSocketGatewayDefinition, + WebSocketGatewayInitContext, WebSocketGatewayServer, WebSocketInterceptor, WebSocketMessage, WebSocketSubscriptionDefinition, }; use serde::{Deserialize, Serialize}; use serde_json::json; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; #[derive(Debug)] @@ -686,6 +687,117 @@ async fn global_websocket_guards_and_interceptors_wrap_gateway_and_subscription_ ); } +struct RecoverWsInterceptor; + +impl WebSocketInterceptor for RecoverWsInterceptor { + fn intercept<'a>( + &'a self, + _context: WebSocketContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + match next.handle().await { + Ok(reply) => Ok(reply), + Err(error) => Ok(Some(WebSocketMessage::text("recovered", error.to_string()))), + } + }) + } +} + +#[tokio::test] +async fn websocket_interceptors_can_recover_errors_before_filters() { + let filter_calls = Arc::new(AtomicUsize::new(0)); + let filter_log = Arc::clone(&filter_calls); + let gateway = WebSocketGatewayDefinition::new("/events") + .unwrap() + .with_interceptor(RecoverWsInterceptor) + .with_filter(move |_, _| { + let filter_log = Arc::clone(&filter_log); + async move { + filter_log.fetch_add(1, Ordering::SeqCst); + Ok(Some(WebSocketExceptionResponse::message( + WebSocketMessage::text("filtered", "unexpected"), + ))) + } + }) + .subscribe("ping", |_| async { + Err::(BootError::BadRequest("broken handler".to_string())) + }) + .unwrap(); + + let reply = gateway + .dispatch( + BootRequest::new(HttpMethod::Get, "/events"), + WebSocketMessage::new("ping", json!(null)), + ) + .await + .unwrap() + .unwrap(); + + assert_eq!(reply.event, "recovered"); + assert!(reply.data.as_str().unwrap().contains("broken handler")); + assert_eq!(filter_calls.load(Ordering::SeqCst), 0); +} + +struct RetryWsInterceptor; + +impl WebSocketInterceptor for RetryWsInterceptor { + fn intercept<'a>( + &'a self, + _context: WebSocketContext, + next: CallHandler<'a, Option>, + ) -> BoxFuture<'a, Result>> { + Box::pin(async move { + match next.handle().await { + Ok(reply) => Ok(reply), + Err(_) => next.handle().await, + } + }) + } +} + +#[tokio::test] +async fn websocket_call_handler_replays_pipes_and_handler_for_retries() { + let pipe_calls = Arc::new(AtomicUsize::new(0)); + let pipe_log = Arc::clone(&pipe_calls); + let handler_calls = Arc::new(AtomicUsize::new(0)); + let handler_log = Arc::clone(&handler_calls); + let gateway = WebSocketGatewayDefinition::new("/events") + .unwrap() + .with_interceptor(RetryWsInterceptor) + .with_pipe(move |message: WebSocketMessage| { + let pipe_log = Arc::clone(&pipe_log); + async move { + pipe_log.fetch_add(1, Ordering::SeqCst); + Ok(message) + } + }) + .subscribe("ping", move |message: WebSocketMessage| { + let handler_log = Arc::clone(&handler_log); + async move { + if handler_log.fetch_add(1, Ordering::SeqCst) == 0 { + Err(BootError::ServiceUnavailable("retry".to_string())) + } else { + Ok(WebSocketMessage::new("pong", message.data)) + } + } + }) + .unwrap(); + + let reply = gateway + .dispatch( + BootRequest::new(HttpMethod::Get, "/events"), + WebSocketMessage::new("ping", json!({ "id": 1 })), + ) + .await + .unwrap() + .unwrap(); + + assert_eq!(reply, WebSocketMessage::new("pong", json!({ "id": 1 }))); + assert_eq!(pipe_calls.load(Ordering::SeqCst), 2); + assert_eq!(handler_calls.load(Ordering::SeqCst), 2); +} + #[tokio::test] async fn global_validation_options_merge_into_websocket_payload_validators() { let app = BootApplication::builder()