1use std::{
4 collections::{BTreeMap, BTreeSet},
5 future::Future,
6 sync::{
7 Arc,
8 atomic::{AtomicUsize, Ordering},
9 },
10 time::{Duration, Instant},
11};
12
13use imagegen_bridge_core::{
14 BridgeError, ErrorCode, FallbackPolicy, ImageRequest, ImageResponse, Normalization,
15 ProviderAttempt, ProviderAttemptOutcome, ProviderContext, ProviderEvent, ProviderRoute,
16 RequestLimits, RevisedPromptPolicy, negotiate_request, validate_request,
17};
18use tokio::sync::{Mutex, Notify, OwnedSemaphorePermit, mpsc};
19use tokio_util::sync::CancellationToken;
20use tracing::Instrument as _;
21use uuid::Uuid;
22
23use crate::{
24 CircuitBreakerConfig, CircuitBreakerSnapshot,
25 admission::AdmissionGate,
26 circuit_breaker::CircuitBreaker,
27 idempotency::{IdempotencyAction, IdempotencyConfig, IdempotencyCoordinator},
28 materialize::{MaterializationConfig, OutputMaterializer},
29 registry::ProviderRegistry,
30 transparency::prepare_transparency,
31};
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct ConcurrencyLimit {
36 pub max_concurrent: usize,
38 pub max_queued: usize,
40}
41
42impl Default for ConcurrencyLimit {
43 fn default() -> Self {
44 Self {
45 max_concurrent: usize::MAX,
46 max_queued: usize::MAX,
47 }
48 }
49}
50
51#[derive(Debug, Clone)]
53pub struct RuntimeConfig {
54 pub request_limits: RequestLimits,
56 pub default_timeout: Duration,
58 pub cancellation_grace: Duration,
60 pub shutdown_grace: Duration,
62 pub global_limit: ConcurrencyLimit,
64 pub default_provider_limit: ConcurrencyLimit,
66 pub provider_limits: BTreeMap<String, ConcurrencyLimit>,
68 pub default_circuit_breaker: CircuitBreakerConfig,
70 pub circuit_breakers: BTreeMap<String, CircuitBreakerConfig>,
72 pub idempotency: IdempotencyConfig,
74 pub materialization: MaterializationConfig,
76}
77
78impl Default for RuntimeConfig {
79 fn default() -> Self {
80 Self {
81 request_limits: RequestLimits::default(),
82 default_timeout: Duration::from_secs(5 * 60),
83 cancellation_grace: Duration::from_secs(1),
84 shutdown_grace: Duration::from_secs(10),
85 global_limit: ConcurrencyLimit {
86 max_concurrent: usize::MAX,
87 max_queued: usize::MAX,
88 },
89 default_provider_limit: ConcurrencyLimit::default(),
90 provider_limits: BTreeMap::new(),
91 default_circuit_breaker: CircuitBreakerConfig::default(),
92 circuit_breakers: BTreeMap::new(),
93 idempotency: IdempotencyConfig::default(),
94 materialization: MaterializationConfig::default(),
95 }
96 }
97}
98
99#[derive(Clone)]
101pub struct ExecutionContext {
102 pub request_id: Option<String>,
104 pub idempotency_scope: String,
106 pub cancellation: CancellationToken,
108}
109
110impl std::fmt::Debug for ExecutionContext {
111 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
112 formatter
113 .debug_struct("ExecutionContext")
114 .field("request_id", &self.request_id)
115 .field("idempotency_scope", &"[REDACTED SCOPE]")
116 .field("cancelled", &self.cancellation.is_cancelled())
117 .finish()
118 }
119}
120
121impl Default for ExecutionContext {
122 fn default() -> Self {
123 Self {
124 request_id: None,
125 idempotency_scope: "default".to_owned(),
126 cancellation: CancellationToken::new(),
127 }
128 }
129}
130
131#[derive(Debug, Clone, PartialEq, Eq)]
133pub struct RuntimeQueueSnapshot {
134 pub global_queued: usize,
136 pub providers_queued: BTreeMap<String, usize>,
138}
139
140pub struct ImagegenRuntime {
142 registry: ProviderRegistry,
143 config: RuntimeConfig,
144 global_gate: Arc<AdmissionGate>,
145 provider_gates: BTreeMap<String, Arc<AdmissionGate>>,
146 circuit_breakers: BTreeMap<String, Arc<CircuitBreaker>>,
147 idempotency: IdempotencyCoordinator,
148 materializer: OutputMaterializer,
149 shutdown: CancellationToken,
150 activity: Arc<ActivityTracker>,
151 shutdown_result: Mutex<Option<Result<(), BridgeError>>>,
152}
153
154struct RouteExecution {
155 response: ImageResponse,
156 effective_request: ImageRequest,
157}
158
159impl std::fmt::Debug for ImagegenRuntime {
160 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
161 formatter
162 .debug_struct("ImagegenRuntime")
163 .field("registry", &self.registry)
164 .field("config", &self.config)
165 .field("shutdown", &self.shutdown.is_cancelled())
166 .finish_non_exhaustive()
167 }
168}
169
170impl ImagegenRuntime {
171 pub fn new(registry: ProviderRegistry, config: RuntimeConfig) -> Result<Self, BridgeError> {
173 if config.default_timeout.is_zero()
174 || config.default_timeout > Duration::from_millis(config.request_limits.max_timeout_ms)
175 {
176 return Err(configuration_error(
177 "default timeout must be within the configured request limit",
178 ));
179 }
180 if config.cancellation_grace.is_zero()
181 || config.cancellation_grace > Duration::from_secs(30)
182 {
183 return Err(configuration_error(
184 "cancellation grace must be greater than zero and at most 30 seconds",
185 ));
186 }
187 if config.shutdown_grace.is_zero() || config.shutdown_grace > Duration::from_secs(5 * 60) {
188 return Err(configuration_error(
189 "shutdown grace must be greater than zero and at most five minutes",
190 ));
191 }
192 for name in config.provider_limits.keys() {
193 registry.resolve(Some(name))?;
194 }
195 for name in config.circuit_breakers.keys() {
196 registry.resolve(Some(name))?;
197 }
198 let global_gate = Arc::new(AdmissionGate::new(
199 config.global_limit.max_concurrent,
200 config.global_limit.max_queued,
201 "global",
202 )?);
203 let mut provider_gates = BTreeMap::new();
204 let mut circuit_breakers = BTreeMap::new();
205 for (name, _) in registry.entries() {
206 let limit = config
207 .provider_limits
208 .get(name)
209 .copied()
210 .unwrap_or(config.default_provider_limit);
211 provider_gates.insert(
212 name.to_owned(),
213 Arc::new(AdmissionGate::new(
214 limit.max_concurrent,
215 limit.max_queued,
216 format!("provider:{name}"),
217 )?),
218 );
219 let breaker = config
220 .circuit_breakers
221 .get(name)
222 .copied()
223 .unwrap_or(config.default_circuit_breaker);
224 circuit_breakers.insert(name.to_owned(), Arc::new(CircuitBreaker::new(breaker)?));
225 }
226 let idempotency = IdempotencyCoordinator::new(config.idempotency)?;
227 let materializer = OutputMaterializer::new(config.materialization.clone())?;
228 Ok(Self {
229 registry,
230 config,
231 global_gate,
232 provider_gates,
233 circuit_breakers,
234 idempotency,
235 materializer,
236 shutdown: CancellationToken::new(),
237 activity: Arc::new(ActivityTracker::default()),
238 shutdown_result: Mutex::new(None),
239 })
240 }
241
242 #[must_use]
244 pub const fn registry(&self) -> &ProviderRegistry {
245 &self.registry
246 }
247
248 pub fn validate_request(&self, request: &ImageRequest) -> Result<(), BridgeError> {
250 validate_request(request, self.config.request_limits)
251 }
252
253 #[must_use]
255 pub const fn has_artifact_store(&self) -> bool {
256 self.materializer.has_artifact_store()
257 }
258
259 pub fn read_artifact(
261 &self,
262 artifact_id: &str,
263 ) -> Result<imagegen_bridge_artifacts::StoredArtifactContent, BridgeError> {
264 self.materializer.read_artifact(artifact_id)
265 }
266
267 pub fn read_artifact_thumbnail(
269 &self,
270 artifact_id: &str,
271 maximum_edge: u32,
272 ) -> Result<Vec<u8>, BridgeError> {
273 self.materializer.read_thumbnail(artifact_id, maximum_edge)
274 }
275
276 pub async fn execute(&self, request: ImageRequest) -> Result<ImageResponse, BridgeError> {
278 self.execute_with(request, ExecutionContext::default())
279 .await
280 }
281
282 pub async fn execute_with(
285 &self,
286 request: ImageRequest,
287 context: ExecutionContext,
288 ) -> Result<ImageResponse, BridgeError> {
289 self.execute_with_optional_events(request, context, None)
290 .await
291 }
292
293 pub async fn execute_with_events(
295 &self,
296 request: ImageRequest,
297 context: ExecutionContext,
298 events: mpsc::Sender<ProviderEvent>,
299 ) -> Result<ImageResponse, BridgeError> {
300 self.execute_with_optional_events(request, context, Some(events))
301 .await
302 }
303
304 async fn execute_with_optional_events(
305 &self,
306 request: ImageRequest,
307 context: ExecutionContext,
308 events: Option<mpsc::Sender<ProviderEvent>>,
309 ) -> Result<ImageResponse, BridgeError> {
310 if self.shutdown.is_cancelled() {
311 return Err(cancelled_error("runtime is shutting down"));
312 }
313 let _activity = self.activity.enter();
314 if self.shutdown.is_cancelled() {
315 return Err(cancelled_error("runtime is shutting down"));
316 }
317 validate_request(&request, self.config.request_limits)?;
318 let request_id = validated_request_id(context.request_id)?;
319 let timeout = request
320 .timeout_ms
321 .map_or(self.config.default_timeout, Duration::from_millis);
322 let deadline = Instant::now()
323 .checked_add(timeout)
324 .ok_or_else(|| configuration_error("request deadline overflowed"))?;
325 let operation = self.shutdown.child_token();
326
327 let mut token = None;
328 if let Some(key) = request.idempotency_key.as_deref() {
329 match self
330 .idempotency
331 .begin(&context.idempotency_scope, key, &request)
332 .await?
333 {
334 IdempotencyAction::Leader(leader) => token = Some(leader),
335 IdempotencyAction::Cached(response) => return Ok(*response),
336 IdempotencyAction::Wait(receiver) => {
337 return self
338 .await_controlled(
339 IdempotencyAction::wait(receiver, deadline, &operation),
340 deadline,
341 &context.cancellation,
342 &operation,
343 "request was cancelled while waiting for idempotent result",
344 )
345 .await;
346 }
347 }
348 }
349
350 let result = self
351 .execute_leader(
352 request,
353 request_id,
354 deadline,
355 &context.cancellation,
356 &operation,
357 events,
358 )
359 .await;
360 match (token, &result) {
361 (Some(token), Ok(response)) => {
362 self.idempotency.complete(token, response.clone()).await;
363 }
364 (Some(token), Err(error)) => {
365 self.idempotency.fail(token, error.clone()).await;
366 }
367 (None, _) => {}
368 }
369 result
370 }
371
372 #[must_use]
374 pub fn queue_snapshot(&self) -> RuntimeQueueSnapshot {
375 RuntimeQueueSnapshot {
376 global_queued: self.global_gate.queued(),
377 providers_queued: self
378 .provider_gates
379 .iter()
380 .map(|(name, gate)| (name.clone(), gate.queued()))
381 .collect(),
382 }
383 }
384
385 #[must_use]
387 pub fn circuit_breaker_snapshot(&self) -> BTreeMap<String, CircuitBreakerSnapshot> {
388 self.circuit_breakers
389 .iter()
390 .map(|(name, breaker)| (name.clone(), breaker.snapshot()))
391 .collect()
392 }
393
394 pub async fn shutdown(&self) -> Result<(), BridgeError> {
396 let mut stored_result = self.shutdown_result.lock().await;
397 if let Some(result) = stored_result.as_ref() {
398 return result.clone();
399 }
400 self.shutdown.cancel();
401 self.global_gate.close();
402 for gate in self.provider_gates.values() {
403 gate.close();
404 }
405 let mut first_error =
406 if tokio::time::timeout(self.config.shutdown_grace, self.activity.wait_until_idle())
407 .await
408 .is_err()
409 {
410 Some(
411 BridgeError::new(
412 ErrorCode::Timeout,
413 "runtime shutdown grace elapsed with active operations",
414 )
415 .retryable(true),
416 )
417 } else {
418 None
419 };
420 for (_, provider) in self.registry.entries() {
421 if let Err(error) = provider.shutdown().await
422 && first_error.is_none()
423 {
424 first_error = Some(error);
425 }
426 }
427 let result = first_error.map_or(Ok(()), Err);
428 *stored_result = Some(result.clone());
429 result
430 }
431
432 async fn execute_leader(
433 &self,
434 request: ImageRequest,
435 request_id: String,
436 deadline: Instant,
437 external_cancellation: &CancellationToken,
438 operation: &CancellationToken,
439 events: Option<mpsc::Sender<ProviderEvent>>,
440 ) -> Result<ImageResponse, BridgeError> {
441 let total_started = Instant::now();
442 let queue_started = Instant::now();
443 let global_permit = self
444 .acquire_controlled(
445 &self.global_gate,
446 deadline,
447 external_cancellation,
448 operation,
449 )
450 .await?;
451 let global_queue_ms = elapsed_ms(queue_started);
452 let primary_provider = request
453 .routing
454 .provider
455 .clone()
456 .unwrap_or_else(|| self.registry.default_name().to_owned());
457 let routes = std::iter::once(ProviderRoute {
458 provider: primary_provider.clone(),
459 model: request.routing.model.clone(),
460 })
461 .chain(request.routing.fallbacks.iter().cloned())
462 .collect::<Vec<_>>();
463 let mut unique_routes = BTreeSet::new();
464 for (index, route) in routes.iter().enumerate() {
465 if !unique_routes.insert((route.provider.as_str(), route.model.as_deref())) {
466 return Err(BridgeError::new(
467 ErrorCode::InvalidRequest,
468 "provider routing contains a duplicate resolved route",
469 )
470 .with_detail("field", format!("routing.fallbacks[{}]", index - 1))
471 .with_detail("provider", &route.provider)
472 .with_detail("model", route.model.as_deref()));
473 }
474 }
475 let mut attempts = Vec::with_capacity(routes.len());
476 let mut final_error = None;
477
478 for (index, route) in routes.iter().enumerate() {
479 let attempt_started = Instant::now();
480 let attempt_span = tracing::info_span!(
481 "imagegen_bridge.provider_attempt",
482 request_id = %request_id,
483 provider = %route.provider,
484 attempt_index = index
485 );
486 match self
487 .execute_route(
488 &request,
489 &request_id,
490 route,
491 deadline,
492 external_cancellation,
493 operation,
494 events.clone(),
495 )
496 .instrument(attempt_span)
497 .await
498 {
499 Ok(execution) => {
500 let mut response = execution.response;
501 attempts.push(ProviderAttempt {
502 provider: response.provider.clone(),
503 model: Some(response.model.clone()),
504 outcome: ProviderAttemptOutcome::Succeeded,
505 error_code: None,
506 duration_ms: elapsed_ms(attempt_started),
507 });
508 if index > 0 {
509 response.normalizations.insert(
510 0,
511 Normalization {
512 field: "routing.provider".to_owned(),
513 requested: Some(serde_json::json!(primary_provider)),
514 effective: Some(serde_json::json!(response.provider)),
515 reason: "provider_fallback".to_owned(),
516 },
517 );
518 response.warnings.push("provider_fallback_used".to_owned());
519 }
520 if !request.routing.fallbacks.is_empty() {
521 response.attempts = attempts;
522 }
523 response.timings.queue_ms =
524 response.timings.queue_ms.saturating_add(global_queue_ms);
525 response.timings.total_ms = elapsed_ms(total_started);
526 self.materializer.attach_metadata(
527 &request,
528 &execution.effective_request,
529 &mut response,
530 )?;
531 drop(global_permit);
532 return Ok(response);
533 }
534 Err(error) => {
535 attempts.push(ProviderAttempt {
536 provider: route.provider.clone(),
537 model: route.model.clone(),
538 outcome: ProviderAttemptOutcome::Failed,
539 error_code: Some(error.code),
540 duration_ms: elapsed_ms(attempt_started),
541 });
542 let has_next = index + 1 < routes.len();
543 if !has_next || !should_fallback(&error, request.routing.fallback_policy) {
544 final_error = Some(error);
545 break;
546 }
547 final_error = Some(error);
548 }
549 }
550 }
551 drop(global_permit);
552 Err(final_error
553 .unwrap_or_else(|| configuration_error("provider routing produced no attempts"))
554 .with_detail("attempts", attempts))
555 }
556
557 #[allow(clippy::too_many_arguments)]
558 async fn execute_route(
559 &self,
560 request: &ImageRequest,
561 request_id: &str,
562 route: &ProviderRoute,
563 deadline: Instant,
564 external_cancellation: &CancellationToken,
565 operation: &CancellationToken,
566 events: Option<mpsc::Sender<ProviderEvent>>,
567 ) -> Result<RouteExecution, BridgeError> {
568 let provider = self.registry.resolve(Some(&route.provider))?;
569 let descriptor = provider.descriptor();
570 let breaker = self
571 .circuit_breakers
572 .get(&descriptor.name)
573 .ok_or_else(|| configuration_error("provider has no configured circuit breaker"))?;
574 let breaker_permit = breaker.acquire(&descriptor.name)?;
575 let provider_gate = self
576 .provider_gates
577 .get(&descriptor.name)
578 .ok_or_else(|| configuration_error("provider has no configured admission gate"))?;
579 let queue_started = Instant::now();
580 let provider_permit = self
581 .acquire_controlled(provider_gate, deadline, external_cancellation, operation)
582 .await;
583 let _provider_permit = match provider_permit {
584 Ok(permit) => permit,
585 Err(error) => {
586 drop(breaker_permit);
587 return Err(error);
588 }
589 };
590 let queue_ms = elapsed_ms(queue_started);
591 let result = self
592 .execute_route_admitted(
593 request,
594 request_id,
595 route,
596 deadline,
597 external_cancellation,
598 operation,
599 events,
600 provider,
601 descriptor,
602 queue_ms,
603 )
604 .await;
605 breaker_permit.finish(&result);
606 result
607 }
608
609 #[allow(clippy::too_many_arguments)]
610 async fn execute_route_admitted(
611 &self,
612 request: &ImageRequest,
613 request_id: &str,
614 route: &ProviderRoute,
615 deadline: Instant,
616 external_cancellation: &CancellationToken,
617 operation: &CancellationToken,
618 events: Option<mpsc::Sender<ProviderEvent>>,
619 provider: Arc<dyn imagegen_bridge_core::ImageProvider>,
620 descriptor: imagegen_bridge_core::ProviderDescriptor,
621 queue_ms: u64,
622 ) -> Result<RouteExecution, BridgeError> {
623 let capabilities = self
624 .await_controlled(
625 provider.capabilities(route.model.as_deref()),
626 deadline,
627 external_cancellation,
628 operation,
629 "request was cancelled during capability discovery",
630 )
631 .await
632 .map_err(|error| attach_provider(error, &descriptor.name))?;
633 if capabilities.provider != descriptor.name {
634 return Err(protocol_error(
635 "provider capability identity does not match its registry identity",
636 )
637 .with_provider(descriptor.name));
638 }
639 let mut routed_request = request.clone();
640 routed_request.routing.provider = Some(route.provider.clone());
641 routed_request.routing.model.clone_from(&route.model);
642 routed_request.routing.fallbacks.clear();
643 let prepared = prepare_transparency(&routed_request, &capabilities)?;
644 validate_request(&prepared.provider_request, self.config.request_limits)?;
645 let negotiated = negotiate_request(&prepared.provider_request, &capabilities)?;
646 let effective_request = negotiated.effective_request;
647
648 let provider_started = Instant::now();
649 let provider_result = self
650 .await_provider_execution(
651 provider.execute(
652 effective_request.clone(),
653 ProviderContext {
654 request_id: request_id.to_owned(),
655 deadline,
656 cancellation: operation.clone(),
657 events,
658 },
659 ),
660 deadline,
661 external_cancellation,
662 operation,
663 )
664 .await;
665 let provider_ms = elapsed_ms(provider_started);
666 let mut response =
667 provider_result.map_err(|error| attach_provider(error, &descriptor.name))?;
668 if let Some(expected_model) = capabilities.model.as_deref()
669 && response.model != expected_model
670 {
671 return Err(protocol_error(
672 "provider response model does not match discovered capabilities",
673 )
674 .with_provider(descriptor.name)
675 .with_detail("expected_model", expected_model)
676 .with_detail("actual_model", response.model));
677 }
678
679 response.id = request_id.to_owned();
680 response.provider = descriptor.name;
681 response.requested = request.parameters.clone();
682 response.effective = effective_request.parameters.clone();
683 if prepared.plan.is_some() {
684 response.effective.background = imagegen_bridge_core::Background::Transparent;
685 response
686 .warnings
687 .push("transparent_background_postprocessed".to_owned());
688 }
689 response.normalizations.splice(
690 0..0,
691 prepared
692 .normalizations
693 .into_iter()
694 .chain(negotiated.normalizations),
695 );
696 if effective_request.policies.revised_prompt == RevisedPromptPolicy::Omit {
697 response.revised_prompt = None;
698 }
699 response.timings.queue_ms = queue_ms;
700 response.timings.provider_ms = provider_ms;
701 let artifact_started = Instant::now();
702 response = self
703 .await_controlled(
704 self.materializer
705 .materialize(response, request, &effective_request, prepared.plan),
706 deadline,
707 external_cancellation,
708 operation,
709 "request was cancelled during output materialization",
710 )
711 .await?;
712 response.timings.artifact_ms = elapsed_ms(artifact_started);
713 Ok(RouteExecution {
714 response,
715 effective_request,
716 })
717 }
718
719 async fn acquire_controlled(
720 &self,
721 gate: &AdmissionGate,
722 deadline: Instant,
723 external: &CancellationToken,
724 operation: &CancellationToken,
725 ) -> Result<OwnedSemaphorePermit, BridgeError> {
726 self.await_controlled(
727 gate.acquire(deadline, operation),
728 deadline,
729 external,
730 operation,
731 "request was cancelled while waiting for capacity",
732 )
733 .await
734 }
735
736 async fn await_provider_execution<T, F>(
737 &self,
738 future: F,
739 deadline: Instant,
740 external: &CancellationToken,
741 operation: &CancellationToken,
742 ) -> Result<T, BridgeError>
743 where
744 F: Future<Output = Result<T, BridgeError>>,
745 {
746 if deadline <= Instant::now() {
747 operation.cancel();
748 return Err(timeout_error());
749 }
750 tokio::pin!(future);
751 let interrupted = tokio::select! {
752 result = &mut future => return classify_provider_execution_result(result),
753 () = external.cancelled() => {
754 operation.cancel();
755 (ErrorCode::Cancelled, "request was cancelled during provider execution")
756 }
757 () = self.shutdown.cancelled() => {
758 operation.cancel();
759 (ErrorCode::Cancelled, "runtime shut down during provider execution")
760 }
761 () = tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)) => {
762 operation.cancel();
763 (ErrorCode::Timeout, "request deadline elapsed during provider execution")
764 }
765 };
766 match tokio::time::timeout(self.config.cancellation_grace, &mut future).await {
767 Ok(Ok(value)) => Ok(value),
768 Ok(Err(error)) if !matches!(error.code, ErrorCode::Timeout | ErrorCode::Cancelled) => {
769 Err(error)
770 }
771 Ok(Err(_)) | Err(_) => Err(unknown_outcome_error(interrupted.0, interrupted.1)),
772 }
773 }
774
775 async fn await_controlled<T, F>(
776 &self,
777 future: F,
778 deadline: Instant,
779 external: &CancellationToken,
780 operation: &CancellationToken,
781 cancellation_message: &'static str,
782 ) -> Result<T, BridgeError>
783 where
784 F: Future<Output = Result<T, BridgeError>>,
785 {
786 if deadline <= Instant::now() {
787 operation.cancel();
788 return Err(timeout_error());
789 }
790 tokio::pin!(future);
791 tokio::select! {
792 result = &mut future => result,
793 () = external.cancelled() => {
794 operation.cancel();
795 let _ = tokio::time::timeout(self.config.cancellation_grace, &mut future).await;
796 Err(cancelled_error(cancellation_message))
797 }
798 () = self.shutdown.cancelled() => {
799 operation.cancel();
800 let _ = tokio::time::timeout(self.config.cancellation_grace, &mut future).await;
801 Err(cancelled_error("runtime is shutting down"))
802 }
803 () = tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)) => {
804 operation.cancel();
805 let _ = tokio::time::timeout(self.config.cancellation_grace, &mut future).await;
806 Err(timeout_error())
807 }
808 }
809 }
810}
811
812#[derive(Default)]
813struct ActivityTracker {
814 active: AtomicUsize,
815 idle: Notify,
816}
817
818impl ActivityTracker {
819 fn enter(self: &Arc<Self>) -> ActivityGuard {
820 self.active.fetch_add(1, Ordering::AcqRel);
821 ActivityGuard {
822 tracker: Arc::clone(self),
823 }
824 }
825
826 async fn wait_until_idle(&self) {
827 loop {
828 let notified = self.idle.notified();
829 if self.active.load(Ordering::Acquire) == 0 {
830 return;
831 }
832 notified.await;
833 }
834 }
835}
836
837struct ActivityGuard {
838 tracker: Arc<ActivityTracker>,
839}
840
841impl Drop for ActivityGuard {
842 fn drop(&mut self) {
843 if self.tracker.active.fetch_sub(1, Ordering::AcqRel) == 1 {
844 self.tracker.idle.notify_waiters();
845 }
846 }
847}
848
849fn validated_request_id(value: Option<String>) -> Result<String, BridgeError> {
850 let Some(value) = value else {
851 return Ok(Uuid::now_v7().to_string());
852 };
853 let valid = !value.is_empty()
854 && value.len() <= 128
855 && value
856 .bytes()
857 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b':'));
858 if valid {
859 Ok(value)
860 } else {
861 Err(BridgeError::new(
862 ErrorCode::InvalidRequest,
863 "request ID is invalid",
864 ))
865 }
866}
867
868fn attach_provider(mut error: BridgeError, provider: &str) -> BridgeError {
869 if error.provider.is_none() {
870 error.provider = Some(provider.to_owned());
871 }
872 error
873}
874
875fn should_fallback(error: &BridgeError, policy: FallbackPolicy) -> bool {
876 if error
877 .details
878 .get("outcome")
879 .and_then(serde_json::Value::as_str)
880 == Some("unknown")
881 {
882 return false;
883 }
884 if matches!(
885 error.code,
886 ErrorCode::InvalidRequest
887 | ErrorCode::PermissionDenied
888 | ErrorCode::SafetyRejected
889 | ErrorCode::Cancelled
890 | ErrorCode::Session
891 | ErrorCode::IdempotencyConflict
892 ) {
893 return false;
894 }
895 if matches!(
896 error.code,
897 ErrorCode::Configuration
898 | ErrorCode::Authentication
899 | ErrorCode::UnsupportedCapability
900 | ErrorCode::RateLimited
901 | ErrorCode::Overloaded
902 ) {
903 return true;
904 }
905 policy == FallbackPolicy::OnError
906 && matches!(
907 error.code,
908 ErrorCode::Upstream
909 | ErrorCode::Protocol
910 | ErrorCode::Artifact
911 | ErrorCode::Input
912 | ErrorCode::Internal
913 )
914}
915
916fn elapsed_ms(started: Instant) -> u64 {
917 u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX)
918}
919
920fn timeout_error() -> BridgeError {
921 BridgeError::new(ErrorCode::Timeout, "request deadline elapsed").retryable(true)
922}
923
924fn unknown_outcome_error(code: ErrorCode, message: impl Into<String>) -> BridgeError {
925 BridgeError::new(code, message)
926 .retryable(false)
927 .with_detail("outcome", "unknown")
928}
929
930fn classify_provider_execution_result<T>(result: Result<T, BridgeError>) -> Result<T, BridgeError> {
931 match result {
932 Err(error)
933 if matches!(error.code, ErrorCode::Timeout | ErrorCode::Cancelled)
934 && error
935 .details
936 .get("outcome")
937 .and_then(serde_json::Value::as_str)
938 != Some("unknown") =>
939 {
940 Err(error.retryable(false).with_detail("outcome", "unknown"))
941 }
942 result => result,
943 }
944}
945
946fn cancelled_error(message: impl Into<String>) -> BridgeError {
947 BridgeError::new(ErrorCode::Cancelled, message)
948}
949
950fn configuration_error(message: impl Into<String>) -> BridgeError {
951 BridgeError::new(ErrorCode::Configuration, message)
952}
953
954fn protocol_error(message: impl Into<String>) -> BridgeError {
955 BridgeError::new(ErrorCode::Protocol, message)
956}
957
958#[cfg(test)]
959mod tests {
960 #![allow(clippy::panic, clippy::unwrap_used)]
961
962 use std::{
963 collections::BTreeSet,
964 sync::atomic::{AtomicBool, AtomicUsize, Ordering},
965 };
966
967 use async_trait::async_trait;
968 use base64::{Engine as _, engine::general_purpose::STANDARD};
969 use imagegen_bridge_artifacts::{ArtifactStore, ImageLimits, inspect_image};
970 use imagegen_bridge_core::{
971 Background, BatchCapabilities, BatchMode, CompatibilityMode, GeneratedImage,
972 GenerationParameters, ImageAction, ImagePayload, ImageProvider, ImageSize,
973 InputCapabilities, InputFidelity, Moderation, OutputFormat, ProviderCapabilities,
974 ProviderDescriptor, Quality, ResponseFormat, SizeCapabilities, SupportLevel, Timings,
975 U8Range, Usage,
976 };
977
978 use super::*;
979 use crate::CircuitState;
980
981 const ONE_PIXEL_PNG: &str = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=";
982
983 struct FakeProvider {
984 calls: AtomicUsize,
985 delay: Duration,
986 saw_cancellation: AtomicBool,
987 corrupt_metadata: bool,
988 }
989
990 impl FakeProvider {
991 fn new(delay: Duration) -> Self {
992 Self {
993 calls: AtomicUsize::new(0),
994 delay,
995 saw_cancellation: AtomicBool::new(false),
996 corrupt_metadata: false,
997 }
998 }
999
1000 fn corrupting() -> Self {
1001 Self {
1002 corrupt_metadata: true,
1003 ..Self::new(Duration::ZERO)
1004 }
1005 }
1006 }
1007
1008 #[async_trait]
1009 impl ImageProvider for FakeProvider {
1010 fn descriptor(&self) -> ProviderDescriptor {
1011 ProviderDescriptor {
1012 name: "fake".to_owned(),
1013 display_name: "Fake".to_owned(),
1014 version: "test".to_owned(),
1015 experimental: false,
1016 models: vec!["test-image".to_owned()],
1017 }
1018 }
1019
1020 async fn capabilities(
1021 &self,
1022 model: Option<&str>,
1023 ) -> Result<ProviderCapabilities, BridgeError> {
1024 Ok(fake_capabilities(model))
1025 }
1026
1027 async fn execute(
1028 &self,
1029 request: ImageRequest,
1030 context: ProviderContext,
1031 ) -> Result<ImageResponse, BridgeError> {
1032 self.calls.fetch_add(1, Ordering::AcqRel);
1033 tokio::select! {
1034 () = tokio::time::sleep(self.delay) => {}
1035 () = context.cancellation.cancelled() => {
1036 self.saw_cancellation.store(true, Ordering::Release);
1037 return Err(BridgeError::new(ErrorCode::Cancelled, "fake cancelled"));
1038 }
1039 }
1040 let bytes = STANDARD.decode(ONE_PIXEL_PNG).unwrap();
1041 let metadata = inspect_image(&bytes, ImageLimits::default()).unwrap();
1042 let image = GeneratedImage {
1043 index: 0,
1044 payload: ImagePayload::B64Json {
1045 b64_json: ONE_PIXEL_PNG.to_owned(),
1046 },
1047 format: metadata.format,
1048 width: metadata.width,
1049 height: metadata.height,
1050 bytes: metadata.bytes,
1051 sha256: if self.corrupt_metadata {
1052 "0".repeat(64)
1053 } else {
1054 metadata.sha256
1055 },
1056 generation_ms: None,
1057 metadata_name: None,
1058 };
1059 Ok(ImageResponse {
1060 id: context.request_id,
1061 created: 0,
1062 provider: "fake".to_owned(),
1063 model: request
1064 .routing
1065 .model
1066 .clone()
1067 .unwrap_or_else(|| "fake-image".to_owned()),
1068 requested: request.parameters.clone(),
1069 effective: request.parameters.clone(),
1070 normalizations: Vec::new(),
1071 attempts: Vec::new(),
1072 data: (0..request.parameters.n)
1073 .map(|index| GeneratedImage {
1074 index,
1075 ..image.clone()
1076 })
1077 .collect(),
1078 failures: Vec::new(),
1079 revised_prompt: Some("safe revised prompt".to_owned()),
1080 usage: Some(Usage::default()),
1081 session: None,
1082 timings: Timings::default(),
1083 warnings: Vec::new(),
1084 })
1085 }
1086
1087 async fn check_ready(&self) -> Result<(), BridgeError> {
1088 Ok(())
1089 }
1090 }
1091
1092 #[derive(Clone, Copy)]
1093 enum RouteFailure {
1094 Authentication,
1095 KnownUpstream,
1096 Safety,
1097 UnknownOutcome,
1098 }
1099
1100 struct RouteProvider {
1101 name: &'static str,
1102 failure: Option<RouteFailure>,
1103 execute_calls: AtomicUsize,
1104 }
1105
1106 impl RouteProvider {
1107 fn new(name: &'static str, failure: Option<RouteFailure>) -> Self {
1108 Self {
1109 name,
1110 failure,
1111 execute_calls: AtomicUsize::new(0),
1112 }
1113 }
1114 }
1115
1116 #[async_trait]
1117 impl ImageProvider for RouteProvider {
1118 fn descriptor(&self) -> ProviderDescriptor {
1119 ProviderDescriptor {
1120 name: self.name.to_owned(),
1121 display_name: self.name.to_owned(),
1122 version: "test".to_owned(),
1123 experimental: false,
1124 models: vec!["route-image".to_owned()],
1125 }
1126 }
1127
1128 async fn capabilities(
1129 &self,
1130 model: Option<&str>,
1131 ) -> Result<ProviderCapabilities, BridgeError> {
1132 if matches!(self.failure, Some(RouteFailure::Authentication)) {
1133 return Err(BridgeError::new(
1134 ErrorCode::Authentication,
1135 "test provider is unavailable",
1136 ));
1137 }
1138 let mut capabilities = fake_capabilities(model);
1139 capabilities.provider = self.name.to_owned();
1140 Ok(capabilities)
1141 }
1142
1143 async fn execute(
1144 &self,
1145 request: ImageRequest,
1146 context: ProviderContext,
1147 ) -> Result<ImageResponse, BridgeError> {
1148 self.execute_calls.fetch_add(1, Ordering::AcqRel);
1149 match self.failure {
1150 Some(RouteFailure::Safety) => {
1151 return Err(BridgeError::safety_rejected("test safety rejection"));
1152 }
1153 Some(RouteFailure::UnknownOutcome) => {
1154 return Err(
1155 BridgeError::new(ErrorCode::Upstream, "test ambiguous failure")
1156 .with_detail("outcome", "unknown"),
1157 );
1158 }
1159 Some(RouteFailure::KnownUpstream) => {
1160 return Err(BridgeError::new(ErrorCode::Upstream, "test known failure")
1161 .with_detail("outcome", "failed"));
1162 }
1163 Some(RouteFailure::Authentication) | None => {}
1164 }
1165 let bytes = STANDARD.decode(ONE_PIXEL_PNG).unwrap();
1166 let metadata = inspect_image(&bytes, ImageLimits::default()).unwrap();
1167 Ok(ImageResponse {
1168 id: context.request_id,
1169 created: 0,
1170 provider: self.name.to_owned(),
1171 model: request
1172 .routing
1173 .model
1174 .clone()
1175 .unwrap_or_else(|| "route-image".to_owned()),
1176 requested: request.parameters.clone(),
1177 effective: request.parameters.clone(),
1178 normalizations: Vec::new(),
1179 attempts: Vec::new(),
1180 data: vec![GeneratedImage {
1181 index: 0,
1182 payload: ImagePayload::B64Json {
1183 b64_json: ONE_PIXEL_PNG.to_owned(),
1184 },
1185 format: metadata.format,
1186 width: metadata.width,
1187 height: metadata.height,
1188 bytes: metadata.bytes,
1189 sha256: metadata.sha256,
1190 generation_ms: None,
1191 metadata_name: None,
1192 }],
1193 failures: Vec::new(),
1194 revised_prompt: None,
1195 usage: None,
1196 session: None,
1197 timings: Timings::default(),
1198 warnings: Vec::new(),
1199 })
1200 }
1201
1202 async fn check_ready(&self) -> Result<(), BridgeError> {
1203 Ok(())
1204 }
1205 }
1206
1207 fn fake_capabilities(model: Option<&str>) -> ProviderCapabilities {
1208 let unsupported_inputs = InputCapabilities {
1209 support: SupportLevel::Unsupported,
1210 max_count: 0,
1211 max_bytes_each: 0,
1212 max_bytes_total: 0,
1213 };
1214 ProviderCapabilities {
1215 provider: "fake".to_owned(),
1216 implementation_version: "test".to_owned(),
1217 model: model.map(str::to_owned),
1218 experimental: false,
1219 generation: true,
1220 edits: false,
1221 count: U8Range { min: 1, max: 1 },
1222 batching: BatchCapabilities {
1223 mode: BatchMode::Native,
1224 native_count: U8Range { min: 1, max: 1 },
1225 max_parallel_outputs: 1,
1226 },
1227 sizes: SizeCapabilities {
1228 auto: true,
1229 allowed: BTreeSet::from([ImageSize::exact(2, 2).unwrap()]),
1230 arbitrary: false,
1231 min_edge: None,
1232 max_edge: None,
1233 edge_multiple: None,
1234 min_pixels: None,
1235 max_pixels: None,
1236 max_aspect_ratio: None,
1237 },
1238 aspect_ratio: SupportLevel::Unsupported,
1239 resolution: SupportLevel::Unsupported,
1240 qualities: BTreeSet::from([Quality::Auto]),
1241 output_formats: BTreeSet::from([OutputFormat::Png]),
1242 backgrounds: BTreeSet::from([Background::Auto]),
1243 transparent_background: SupportLevel::Emulated,
1244 moderation: BTreeSet::from([Moderation::Auto]),
1245 negative_prompt: SupportLevel::Emulated,
1246 revised_prompt: SupportLevel::Native,
1247 user_attribution: SupportLevel::Native,
1248 input_fidelities: BTreeSet::from([InputFidelity::Low, InputFidelity::High]),
1249 actions: BTreeSet::from([ImageAction::Auto, ImageAction::Generate, ImageAction::Edit]),
1250 reference_images: unsupported_inputs.clone(),
1251 edit_images: unsupported_inputs.clone(),
1252 masks: unsupported_inputs,
1253 partial_images: U8Range { min: 0, max: 0 },
1254 persistent_sessions: false,
1255 explicit_threads: false,
1256 }
1257 }
1258
1259 fn runtime(
1260 provider: &Arc<FakeProvider>,
1261 mutate: impl FnOnce(&mut RuntimeConfig),
1262 ) -> ImagegenRuntime {
1263 let registry =
1264 ProviderRegistry::new([Arc::clone(provider) as Arc<dyn ImageProvider>], "fake")
1265 .unwrap();
1266 let mut config = RuntimeConfig::default();
1267 mutate(&mut config);
1268 ImagegenRuntime::new(registry, config).unwrap()
1269 }
1270
1271 fn routing_runtime(
1272 primary: &Arc<RouteProvider>,
1273 fallback: &Arc<RouteProvider>,
1274 ) -> ImagegenRuntime {
1275 routing_runtime_with(primary, fallback, |_| {})
1276 }
1277
1278 fn routing_runtime_with(
1279 primary: &Arc<RouteProvider>,
1280 fallback: &Arc<RouteProvider>,
1281 mutate: impl FnOnce(&mut RuntimeConfig),
1282 ) -> ImagegenRuntime {
1283 let registry = ProviderRegistry::new(
1284 [
1285 Arc::clone(primary) as Arc<dyn ImageProvider>,
1286 Arc::clone(fallback) as Arc<dyn ImageProvider>,
1287 ],
1288 primary.name,
1289 )
1290 .unwrap();
1291 let mut config = RuntimeConfig::default();
1292 mutate(&mut config);
1293 ImagegenRuntime::new(registry, config).unwrap()
1294 }
1295
1296 #[tokio::test]
1297 async fn negotiates_then_independently_verifies_and_projects_metadata() {
1298 let provider = Arc::new(FakeProvider::new(Duration::ZERO));
1299 let runtime = runtime(&provider, |_| {});
1300 let mut request = ImageRequest::generate("test");
1301 request.parameters = GenerationParameters {
1302 n: 2,
1303 ..GenerationParameters::default()
1304 };
1305 request.policies.compatibility = CompatibilityMode::Normalize;
1306 request.policies.revised_prompt = RevisedPromptPolicy::Omit;
1307 request.output.response_format = ResponseFormat::Metadata;
1308 let response = runtime
1309 .execute_with(
1310 request,
1311 ExecutionContext {
1312 request_id: Some("request-1".to_owned()),
1313 ..ExecutionContext::default()
1314 },
1315 )
1316 .await
1317 .unwrap();
1318 assert_eq!(response.id, "request-1");
1319 assert_eq!(response.requested.n, 2);
1320 assert_eq!(response.effective.n, 1);
1321 assert_eq!(response.data.len(), 1);
1322 assert!(matches!(response.data[0].payload, ImagePayload::Metadata));
1323 assert!(response.revised_prompt.is_none());
1324 assert_eq!(provider.calls.load(Ordering::Acquire), 1);
1325 }
1326
1327 #[tokio::test]
1328 async fn strict_mode_reports_verified_dimension_mismatch() {
1329 let provider = Arc::new(FakeProvider::new(Duration::ZERO));
1330 let runtime = runtime(&provider, |_| {});
1331 let mut request = ImageRequest::generate("test");
1332 request.parameters.size = ImageSize::exact(2, 2).unwrap();
1333 let response = runtime.execute(request).await.unwrap();
1334 assert_eq!(response.effective.size, ImageSize::exact(1, 1).unwrap());
1335 assert!(response.normalizations.iter().any(|entry| {
1336 entry.field == "parameters.size"
1337 && entry.requested == Some(serde_json::json!("2x2"))
1338 && entry.effective == Some(serde_json::json!("1x1"))
1339 && entry.reason == "provider_output_dimensions_differed"
1340 }));
1341 assert!(
1342 response
1343 .warnings
1344 .iter()
1345 .any(|warning| warning == "provider_output_dimensions_differed")
1346 );
1347 }
1348
1349 #[tokio::test]
1350 async fn normalize_mode_reports_verified_dimension_mismatch() {
1351 let provider = Arc::new(FakeProvider::new(Duration::ZERO));
1352 let runtime = runtime(&provider, |_| {});
1353 let mut request = ImageRequest::generate("test");
1354 request.parameters.size = ImageSize::exact(2, 2).unwrap();
1355 request.policies.compatibility = CompatibilityMode::Normalize;
1356 let response = runtime.execute(request).await.unwrap();
1357 assert_eq!(response.effective.size, ImageSize::exact(1, 1).unwrap());
1358 assert!(response.normalizations.iter().any(|entry| {
1359 entry.field == "parameters.size"
1360 && entry.reason == "provider_output_dimensions_differed"
1361 }));
1362 assert!(
1363 response
1364 .warnings
1365 .iter()
1366 .any(|warning| warning == "provider_output_dimensions_differed")
1367 );
1368 }
1369
1370 #[tokio::test]
1371 async fn replays_idempotent_responses_without_a_second_provider_call() {
1372 let provider = Arc::new(FakeProvider::new(Duration::ZERO));
1373 let runtime = runtime(&provider, |_| {});
1374 let mut request = ImageRequest::generate("test");
1375 request.idempotency_key = Some("stable-key".to_owned());
1376 let first = runtime.execute(request.clone()).await.unwrap();
1377 let second = runtime.execute(request).await.unwrap();
1378 assert_eq!(first, second);
1379 assert_eq!(provider.calls.load(Ordering::Acquire), 1);
1380 }
1381
1382 #[tokio::test]
1383 async fn publishes_bridge_owned_artifacts_only_after_verification() {
1384 let directory = tempfile::tempdir().unwrap();
1385 let store = Arc::new(ArtifactStore::new(directory.path(), ImageLimits::default()).unwrap());
1386 let provider = Arc::new(FakeProvider::new(Duration::ZERO));
1387 let runtime = runtime(&provider, |config| {
1388 config.materialization.artifact_store = Some(store);
1389 });
1390 let mut request = ImageRequest::generate("test");
1391 request.output.response_format = ResponseFormat::Artifact;
1392 request.output.filename_prefix = Some("Result".to_owned());
1393 let response = runtime.execute(request).await.unwrap();
1394 assert!(matches!(
1395 &response.data[0].payload,
1396 ImagePayload::Artifact { name: Some(name), .. } if name.starts_with("result-")
1397 ));
1398 }
1399
1400 #[tokio::test]
1401 async fn deadline_keeps_an_unknown_idempotency_tombstone() {
1402 let provider = Arc::new(FakeProvider::new(Duration::from_secs(5)));
1403 let runtime = runtime(&provider, |_| {});
1404 let mut request = ImageRequest::generate("test");
1405 request.timeout_ms = Some(100);
1406 request.idempotency_key = Some("retryable-key".to_owned());
1407 let first = runtime.execute(request.clone()).await.unwrap_err();
1408 assert_eq!(first.code, ErrorCode::Timeout);
1409 assert!(!first.retryable);
1410 assert_eq!(first.details["outcome"], "unknown");
1411 tokio::task::yield_now().await;
1412 assert!(provider.saw_cancellation.load(Ordering::Acquire));
1413 let second = runtime.execute(request).await.unwrap_err();
1414 assert_eq!(second.code, ErrorCode::Timeout);
1415 assert!(!second.retryable);
1416 assert_eq!(second.details["outcome"], "unknown");
1417 assert_eq!(provider.calls.load(Ordering::Acquire), 1);
1418 }
1419
1420 #[tokio::test]
1421 async fn rejects_provider_metadata_that_does_not_match_decoded_output() {
1422 let provider = Arc::new(FakeProvider::corrupting());
1423 let runtime = runtime(&provider, |_| {});
1424 let error = runtime
1425 .execute(ImageRequest::generate("test"))
1426 .await
1427 .unwrap_err();
1428 assert_eq!(error.code, ErrorCode::Protocol);
1429 }
1430
1431 #[tokio::test]
1432 async fn runtime_queue_rejects_work_beyond_its_explicit_bound() {
1433 let provider = Arc::new(FakeProvider::new(Duration::from_millis(40)));
1434 let runtime = Arc::new(runtime(&provider, |config| {
1435 config.global_limit = ConcurrencyLimit {
1436 max_concurrent: 1,
1437 max_queued: 1,
1438 };
1439 config.default_provider_limit = ConcurrencyLimit {
1440 max_concurrent: 1,
1441 max_queued: 1,
1442 };
1443 }));
1444 let first_runtime = Arc::clone(&runtime);
1445 let first =
1446 tokio::spawn(
1447 async move { first_runtime.execute(ImageRequest::generate("first")).await },
1448 );
1449 while provider.calls.load(Ordering::Acquire) == 0 {
1450 tokio::task::yield_now().await;
1451 }
1452 let second_runtime = Arc::clone(&runtime);
1453 let second = tokio::spawn(async move {
1454 second_runtime
1455 .execute(ImageRequest::generate("second"))
1456 .await
1457 });
1458 while runtime.queue_snapshot().global_queued == 0 {
1459 tokio::task::yield_now().await;
1460 }
1461 let error = runtime
1462 .execute(ImageRequest::generate("rejected"))
1463 .await
1464 .unwrap_err();
1465 assert_eq!(error.code, ErrorCode::Overloaded);
1466 assert!(first.await.unwrap().is_ok());
1467 assert!(second.await.unwrap().is_ok());
1468 }
1469
1470 #[tokio::test]
1471 async fn concurrent_idempotent_callers_share_one_provider_operation() {
1472 let provider = Arc::new(FakeProvider::new(Duration::from_millis(20)));
1473 let runtime = Arc::new(runtime(&provider, |_| {}));
1474 let mut request = ImageRequest::generate("same");
1475 request.idempotency_key = Some("shared-key".to_owned());
1476 let first_runtime = Arc::clone(&runtime);
1477 let first_request = request.clone();
1478 let first = tokio::spawn(async move { first_runtime.execute(first_request).await });
1479 while provider.calls.load(Ordering::Acquire) == 0 {
1480 tokio::task::yield_now().await;
1481 }
1482 let second = runtime.execute(request).await.unwrap();
1483 let first = first.await.unwrap().unwrap();
1484 assert_eq!(first, second);
1485 assert_eq!(provider.calls.load(Ordering::Acquire), 1);
1486 }
1487
1488 #[tokio::test]
1489 async fn unavailable_primary_falls_back_in_order_and_records_attempts() {
1490 let primary = Arc::new(RouteProvider::new(
1491 "primary",
1492 Some(RouteFailure::Authentication),
1493 ));
1494 let fallback = Arc::new(RouteProvider::new("fallback", None));
1495 let runtime = routing_runtime(&primary, &fallback);
1496 let mut request = ImageRequest::generate("test routing");
1497 request.routing.fallbacks.push(ProviderRoute {
1498 provider: "fallback".to_owned(),
1499 model: None,
1500 });
1501 let response = runtime.execute(request).await.unwrap();
1502 assert_eq!(response.provider, "fallback");
1503 assert_eq!(response.attempts.len(), 2);
1504 assert_eq!(
1505 response.attempts[0].error_code,
1506 Some(ErrorCode::Authentication)
1507 );
1508 assert_eq!(
1509 response.attempts[1].outcome,
1510 ProviderAttemptOutcome::Succeeded
1511 );
1512 assert!(
1513 response
1514 .warnings
1515 .contains(&"provider_fallback_used".to_owned())
1516 );
1517 assert_eq!(primary.execute_calls.load(Ordering::Acquire), 0);
1518 assert_eq!(fallback.execute_calls.load(Ordering::Acquire), 1);
1519 }
1520
1521 #[tokio::test]
1522 async fn open_primary_circuit_fails_fast_and_preserves_explicit_fallback() {
1523 let primary = Arc::new(RouteProvider::new(
1524 "primary",
1525 Some(RouteFailure::KnownUpstream),
1526 ));
1527 let fallback = Arc::new(RouteProvider::new("fallback", None));
1528 let runtime = routing_runtime_with(&primary, &fallback, |config| {
1529 config.default_circuit_breaker.failure_threshold = 1;
1530 });
1531 let mut request = ImageRequest::generate("test breaker fallback");
1532 request.routing.fallback_policy = FallbackPolicy::OnError;
1533 request.routing.fallbacks.push(ProviderRoute {
1534 provider: "fallback".to_owned(),
1535 model: None,
1536 });
1537 assert_eq!(
1538 runtime.execute(request.clone()).await.unwrap().provider,
1539 "fallback"
1540 );
1541 assert_eq!(
1542 runtime.circuit_breaker_snapshot()["primary"].state,
1543 CircuitState::Open
1544 );
1545 let second = runtime.execute(request).await.unwrap();
1546 assert_eq!(second.provider, "fallback");
1547 assert_eq!(primary.execute_calls.load(Ordering::Acquire), 1);
1548 assert_eq!(fallback.execute_calls.load(Ordering::Acquire), 2);
1549 }
1550
1551 #[tokio::test]
1552 async fn unknown_outcome_opens_only_for_later_calls_without_retrying_current_call() {
1553 let primary = Arc::new(RouteProvider::new(
1554 "primary",
1555 Some(RouteFailure::UnknownOutcome),
1556 ));
1557 let fallback = Arc::new(RouteProvider::new("fallback", None));
1558 let runtime = routing_runtime_with(&primary, &fallback, |config| {
1559 config.default_circuit_breaker.failure_threshold = 1;
1560 });
1561 let mut request = ImageRequest::generate("test unknown outcome");
1562 request.routing.fallback_policy = FallbackPolicy::OnError;
1563 request.routing.fallbacks.push(ProviderRoute {
1564 provider: "fallback".to_owned(),
1565 model: None,
1566 });
1567 let first = runtime.execute(request.clone()).await.unwrap_err();
1568 assert_eq!(first.details["outcome"], "unknown");
1569 assert_eq!(fallback.execute_calls.load(Ordering::Acquire), 0);
1570 assert_eq!(
1571 runtime.circuit_breaker_snapshot()["primary"].state,
1572 CircuitState::Open
1573 );
1574 assert_eq!(runtime.execute(request).await.unwrap().provider, "fallback");
1575 assert_eq!(primary.execute_calls.load(Ordering::Acquire), 1);
1576 }
1577
1578 #[tokio::test]
1579 async fn fallback_never_bypasses_safety_or_unknown_outcomes() {
1580 for failure in [RouteFailure::Safety, RouteFailure::UnknownOutcome] {
1581 let primary = Arc::new(RouteProvider::new("primary", Some(failure)));
1582 let fallback = Arc::new(RouteProvider::new("fallback", None));
1583 let runtime = routing_runtime(&primary, &fallback);
1584 let mut request = ImageRequest::generate("test routing guard");
1585 request.routing.fallback_policy = FallbackPolicy::OnError;
1586 request.routing.fallbacks.push(ProviderRoute {
1587 provider: "fallback".to_owned(),
1588 model: None,
1589 });
1590 let error = runtime.execute(request).await.unwrap_err();
1591 assert_eq!(fallback.execute_calls.load(Ordering::Acquire), 0);
1592 assert_eq!(error.details["attempts"].as_array().unwrap().len(), 1);
1593 }
1594 }
1595
1596 #[tokio::test]
1597 async fn on_error_expands_but_does_not_replace_unavailable_policy() {
1598 for (policy, expected_calls) in [
1599 (FallbackPolicy::OnUnavailable, 0),
1600 (FallbackPolicy::OnError, 1),
1601 ] {
1602 let primary = Arc::new(RouteProvider::new(
1603 "primary",
1604 Some(RouteFailure::KnownUpstream),
1605 ));
1606 let fallback = Arc::new(RouteProvider::new("fallback", None));
1607 let runtime = routing_runtime(&primary, &fallback);
1608 let mut request = ImageRequest::generate("test known outcome policy");
1609 request.routing.fallback_policy = policy;
1610 request.routing.fallbacks.push(ProviderRoute {
1611 provider: "fallback".to_owned(),
1612 model: None,
1613 });
1614 let result = runtime.execute(request).await;
1615 assert_eq!(
1616 fallback.execute_calls.load(Ordering::Acquire),
1617 expected_calls
1618 );
1619 assert_eq!(result.is_ok(), policy == FallbackPolicy::OnError);
1620 }
1621 }
1622
1623 #[tokio::test]
1624 async fn implicit_primary_cannot_be_repeated_as_a_fallback() {
1625 let primary = Arc::new(RouteProvider::new("primary", None));
1626 let fallback = Arc::new(RouteProvider::new("fallback", None));
1627 let runtime = routing_runtime(&primary, &fallback);
1628 let mut request = ImageRequest::generate("test duplicate route");
1629 request.routing.fallbacks.push(ProviderRoute {
1630 provider: "primary".to_owned(),
1631 model: None,
1632 });
1633 let error = runtime.execute(request).await.unwrap_err();
1634 assert_eq!(error.code, ErrorCode::InvalidRequest);
1635 assert_eq!(primary.execute_calls.load(Ordering::Acquire), 0);
1636 }
1637
1638 #[tokio::test]
1639 async fn shutdown_cancels_and_drains_active_work_before_provider_release() {
1640 let provider = Arc::new(FakeProvider::new(Duration::from_secs(5)));
1641 let runtime = Arc::new(runtime(&provider, |_| {}));
1642 let executing_runtime = Arc::clone(&runtime);
1643 let executing = tokio::spawn(async move {
1644 executing_runtime
1645 .execute(ImageRequest::generate("active"))
1646 .await
1647 });
1648 while provider.calls.load(Ordering::Acquire) == 0 {
1649 tokio::task::yield_now().await;
1650 }
1651 runtime.shutdown().await.unwrap();
1652 runtime.shutdown().await.unwrap();
1653 let error = executing.await.unwrap().unwrap_err();
1654 assert_eq!(error.code, ErrorCode::Cancelled);
1655 assert!(provider.saw_cancellation.load(Ordering::Acquire));
1656 let error = runtime
1657 .execute(ImageRequest::generate("after shutdown"))
1658 .await
1659 .unwrap_err();
1660 assert_eq!(error.code, ErrorCode::Cancelled);
1661 }
1662}