1use std::sync::Arc;
8
9use crate::cache::{ArticleAvailability, CachedArticle, UnifiedCache};
10use crate::metrics::MetricsCollector;
11use crate::pool::BufferPool;
12use crate::protocol::{RequestContext, StatusCode};
13use crate::router::{ArticleBackend, BackendSelector};
14use crate::session::backend;
15use crate::session::handlers::should_sample_backend_timing;
16use crate::types::{BackendId, MessageId};
17use futures::{StreamExt, stream::FuturesUnordered};
18
19#[derive(Debug)]
20#[cfg_attr(test, derive(PartialEq, Eq))]
21#[allow(clippy::large_enum_variant)]
22pub(crate) enum PrecheckHit {
23 Payload(crate::cache::CacheIngestResponse),
24 Availability(StatusCode),
25}
26
27impl PrecheckHit {
28 #[must_use]
29 const fn will_update_cache(&self, cache: &UnifiedCache) -> bool {
30 match self {
31 Self::Payload(_) => cache.stores_payload_responses(),
32 Self::Availability(_) => cache.records_backend_has_status(),
33 }
34 }
35}
36
37#[allow(clippy::large_enum_variant)]
39pub(crate) enum PrecheckResponse {
40 Cached(CachedArticle),
41 Direct(crate::cache::CacheIngestResponse),
42}
43
44#[derive(Debug)]
46#[cfg_attr(test, derive(PartialEq, Eq))]
47#[allow(clippy::large_enum_variant)]
48pub(crate) enum QueryResult {
49 Found(BackendId, PrecheckHit),
50 Missing(BackendId),
51 Error,
52}
53
54#[derive(Debug)]
55enum QueryAttemptResult {
56 Found(Box<PrecheckHit>),
57 Missing,
58 Error,
59}
60
61#[derive(Debug, PartialEq, Eq)]
62enum TierQuerySummary {
63 Exhausted(ArticleAvailability),
64 Inconclusive(ArticleAvailability),
65}
66
67impl TierQuerySummary {
68 const fn is_exhausted(&self) -> bool {
69 matches!(self, Self::Exhausted(_))
70 }
71}
72
73fn summarize_tier_results(results: &[QueryResult]) -> TierQuerySummary {
74 let mut availability = ArticleAvailability::new();
75 let mut exhausted = true;
76
77 for result in results {
78 match result {
79 QueryResult::Missing(id) => {
80 availability.record_missing(*id);
81 }
82 QueryResult::Found(_, _) => {
83 exhausted = false;
84 }
85 QueryResult::Error => {
86 exhausted = false;
87 }
88 }
89 }
90
91 if exhausted {
92 TierQuerySummary::Exhausted(availability)
93 } else {
94 TierQuerySummary::Inconclusive(availability)
95 }
96}
97
98#[derive(Debug)]
99#[cfg_attr(test, derive(PartialEq, Eq))]
100#[allow(clippy::large_enum_variant)]
101enum RacingQueryOutcome {
102 Hit(BackendId, PrecheckHit),
103 Missing,
104 Inconclusive,
105}
106
107#[derive(Clone, Copy)]
109pub struct PrecheckDeps<'a> {
110 pub router: &'a Arc<BackendSelector>,
111 pub cache: &'a Arc<UnifiedCache>,
112 pub buffer_pool: &'a BufferPool,
113 pub metrics: &'a MetricsCollector,
114 pub cache_articles: bool,
115}
116
117#[derive(Clone)]
118struct OwnedDeps {
119 router: Arc<BackendSelector>,
120 cache: Arc<UnifiedCache>,
121 buffer_pool: BufferPool,
122 metrics: MetricsCollector,
123 cache_articles: bool,
124}
125
126impl PrecheckDeps<'_> {
127 fn clone_deps(&self) -> OwnedDeps {
128 OwnedDeps {
129 router: Arc::clone(self.router),
130 cache: Arc::clone(self.cache),
131 buffer_pool: self.buffer_pool.clone(),
132 metrics: self.metrics.clone(),
133 cache_articles: self.cache_articles,
134 }
135 }
136}
137
138async fn cache_precheck_hit(
139 cache: &UnifiedCache,
140 msg_id: MessageId<'static>,
141 backend: BackendId,
142 hit: PrecheckHit,
143 tier: crate::cache::ttl::CacheTier,
144) {
145 if !hit.will_update_cache(cache) {
146 return;
147 }
148
149 match hit {
150 PrecheckHit::Payload(data) => {
151 cache.upsert_ingest(msg_id, data, backend, tier).await;
152 }
153 PrecheckHit::Availability(status_code) => {
154 cache
155 .record_backend_has_status(msg_id, status_code, backend, tier)
156 .await;
157 }
158 }
159}
160
161async fn record_precheck_result(
162 cache: &UnifiedCache,
163 msg_id: MessageId<'static>,
164 result: QueryResult,
165) -> Option<(BackendId, PrecheckHit)> {
166 match result {
167 QueryResult::Found(backend, hit) => Some((backend, hit)),
168 QueryResult::Missing(backend_id) => {
169 cache.record_backend_missing(msg_id, backend_id).await;
170 None
171 }
172 QueryResult::Error => None,
173 }
174}
175
176async fn query_backend(
177 deps: &OwnedDeps,
178 backend: ArticleBackend,
179 request: &RequestContext,
180) -> QueryResult {
181 let backend_id = backend.backend_id();
182 let Some(provider) = deps.router.backend_provider(backend_id) else {
183 return QueryResult::Error;
184 };
185
186 deps.router.mark_backend_pending(backend_id);
188
189 let query_result = crate::session::retry::retry_once!(
191 execute_backend_query(deps, provider, &backend, request).await,
192 backend = backend_id.as_index()
193 )
194 .unwrap_or(QueryAttemptResult::Error);
195
196 deps.router.complete_command(backend_id);
198
199 match query_result {
200 QueryAttemptResult::Found(hit) => QueryResult::Found(backend_id, *hit),
201 QueryAttemptResult::Missing => QueryResult::Missing(backend_id),
202 QueryAttemptResult::Error => QueryResult::Error,
203 }
204}
205
206async fn execute_backend_query(
211 deps: &OwnedDeps,
212 provider: &crate::pool::DeadpoolConnectionProvider,
213 backend: &ArticleBackend,
214 request: &RequestContext,
215) -> Result<QueryAttemptResult, ()> {
216 let backend_id = backend.backend_id();
217 let Ok(conn_raw) = provider.get_pooled_connection().await else {
218 return Ok(QueryAttemptResult::Error);
219 };
220 let mut conn = crate::pool::ConnectionGuard::new(conn_raw, provider.clone());
221
222 let mut buffer = deps.buffer_pool.acquire();
223
224 let response = if should_sample_backend_timing() {
225 backend::send_request_timed(&mut **conn, request, &mut buffer)
226 .await
227 .map(|(response, ttfb, send, recv)| (response, Some((ttfb, send, recv))))
228 } else {
229 backend::send_request(&mut **conn, request, &mut buffer)
230 .await
231 .map(|response| (response, None))
232 };
233
234 match response {
236 Ok((response, timings)) => {
237 let Some(status_code) = response.status_code() else {
238 response.log_warnings(&buffer, "adaptive-precheck", backend_id);
239 conn.retire_with_cooldown();
240 return Err(());
241 };
242 let single_line_payload = response
243 .single_line_bytes(&buffer)
244 .map(crate::cache::CacheIngestResponse::from);
245
246 let response = build_precheck_hit(
247 deps,
248 request,
249 status_code,
250 single_line_payload,
251 &mut conn,
252 &mut buffer,
253 )
254 .await?;
255
256 let result = classify_precheck_result(deps, backend, status_code, timings, response);
257
258 let _ = conn.release(); Ok(result)
260 }
261 Err(_) => {
262 conn.retire_with_cooldown();
263 Err(())
264 }
265 }
266}
267
268async fn build_precheck_hit(
269 deps: &OwnedDeps,
270 request: &RequestContext,
271 status_code: StatusCode,
272 single_line_payload: Option<crate::cache::CacheIngestResponse>,
273 conn: &mut crate::pool::ConnectionGuard,
274 buffer: &mut crate::pool::PooledBuffer,
275) -> Result<PrecheckHit, ()> {
276 if request.has_response_body(status_code) {
277 return read_complete_precheck_hit(deps, status_code, conn, buffer).await;
278 }
279
280 if let Some(payload) = single_line_payload {
281 Ok(PrecheckHit::Payload(payload))
282 } else {
283 Ok(PrecheckHit::Availability(status_code))
284 }
285}
286
287async fn read_complete_precheck_hit(
288 deps: &OwnedDeps,
289 status_code: StatusCode,
290 conn: &mut crate::pool::ConnectionGuard,
291 buffer: &mut crate::pool::PooledBuffer,
292) -> Result<PrecheckHit, ()> {
293 let mut response = deps
294 .cache_articles
295 .then(crate::pool::ChunkedResponse::default);
296
297 if let Some(response) = &mut response {
298 let retained =
299 crate::session::backend::capture_complete_multiline_response_chunked_optional(
300 conn.as_mut(),
301 buffer,
302 &deps.buffer_pool,
303 response,
304 )
305 .await
306 .map_err(|_| ())?;
307 if !retained {
308 response.clear();
309 return Ok(PrecheckHit::Availability(status_code));
310 }
311 } else {
312 crate::session::backend::observe_complete_multiline_response(conn.as_mut(), buffer)
313 .await
314 .map_err(|_| ())?;
315 }
316
317 if let Some(response) = response {
318 Ok(PrecheckHit::Payload(
319 crate::cache::CacheIngestResponse::Chunked(response),
320 ))
321 } else {
322 Ok(PrecheckHit::Availability(status_code))
323 }
324}
325
326fn classify_precheck_result(
327 deps: &OwnedDeps,
328 backend: &ArticleBackend,
329 status_code: StatusCode,
330 timings: Option<(u64, u64, u64)>,
331 response: PrecheckHit,
332) -> QueryAttemptResult {
333 let backend_id = backend.backend_id();
334 deps.metrics.record_command(backend_id);
335 if let Some((ttfb, send, recv)) = timings {
336 deps.metrics.record_ttfb_micros(backend_id, ttfb);
337 deps.metrics.record_send_recv_micros(backend_id, send, recv);
338 }
339
340 match status_code.as_u16() {
341 220..=223 => {
342 tracing::debug!(
343 backend = backend_id.as_index(),
344 "precheck recording command to metrics"
345 );
346 QueryAttemptResult::Found(Box::new(response))
347 }
348 430 => {
349 deps.metrics.record_error_4xx(backend_id);
350 QueryAttemptResult::Missing
351 }
352 400..=499 => {
353 deps.metrics.record_error_4xx(backend_id);
354 QueryAttemptResult::Error
355 }
356 500..=599 => {
357 deps.metrics.record_error_5xx(backend_id);
358 QueryAttemptResult::Error
359 }
360 _ => {
361 deps.metrics.record_error(backend_id);
362 QueryAttemptResult::Error
363 }
364 }
365}
366
367async fn query_all_backends(
369 deps: &OwnedDeps,
370 request: &RequestContext,
371 availability: &ArticleAvailability,
372) -> Vec<QueryResult> {
373 let mut results = Vec::with_capacity(deps.router.backend_count().get());
374
375 for tier in deps.router.tiers() {
376 let tier_results = query_tier_backends(deps, request, tier, availability).await;
377 let found = tier_results
378 .iter()
379 .any(|result| matches!(result, QueryResult::Found(_, _)));
380 let tier_exhausted = summarize_tier_results(&tier_results).is_exhausted();
381 results.extend(tier_results);
382
383 if found || !tier_exhausted {
384 break;
385 }
386 }
387
388 results
389}
390
391async fn query_all_backends_racing(
394 deps: &OwnedDeps,
395 request: &RequestContext,
396 msg_id: &MessageId<'_>,
397 initial_availability: &ArticleAvailability,
398) -> RacingQueryOutcome {
399 let mut results = Vec::with_capacity(deps.router.backend_count().get());
400
401 for tier in deps.router.tiers() {
402 let mut pending = spawn_backend_queries_for_tier(deps, request, tier, initial_availability);
403 let tier_results_start = results.len();
404
405 while let Some(result) = pending.next().await {
406 let Ok(result) = result else {
407 continue;
408 };
409 match result {
410 QueryResult::Found(id, response) => {
411 let cache = deps.cache.clone();
412 let msg_id_owned = msg_id.to_owned();
413 tokio::spawn(async move {
414 while let Some(result) = pending.next().await {
415 let Ok(result) = result else {
416 continue;
417 };
418 let Some((backend, hit)) =
419 record_precheck_result(&cache, msg_id_owned.clone(), result).await
420 else {
421 continue;
422 };
423 let tier = crate::cache::ttl::CacheTier::new(0);
424 cache_precheck_hit(&cache, msg_id_owned.clone(), backend, hit, tier)
425 .await;
426 }
427 });
428 return RacingQueryOutcome::Hit(id, response);
429 }
430 QueryResult::Missing(backend_id) => {
431 deps.cache
432 .record_backend_missing(msg_id.to_owned(), backend_id)
433 .await;
434 results.push(QueryResult::Missing(backend_id));
435 }
436 QueryResult::Error => {
437 results.push(QueryResult::Error);
438 }
439 }
440 }
441
442 let tier_summary = summarize_tier_results(&results[tier_results_start..]);
443 results.clear();
444 if !tier_summary.is_exhausted() {
445 return RacingQueryOutcome::Inconclusive;
446 }
447 }
448
449 RacingQueryOutcome::Missing
450}
451
452async fn query_tier_backends(
453 deps: &OwnedDeps,
454 request: &RequestContext,
455 tier: u8,
456 availability: &ArticleAvailability,
457) -> Vec<QueryResult> {
458 let mut pending = spawn_backend_queries_for_tier(deps, request, tier, availability);
459 let mut results = Vec::new();
460
461 while let Some(result) = pending.next().await {
462 if let Ok(result) = result {
463 results.push(result);
464 }
465 }
466
467 results
468}
469
470fn spawn_backend_queries_for_tier(
471 deps: &OwnedDeps,
472 request: &RequestContext,
473 tier: u8,
474 availability: &ArticleAvailability,
475) -> FuturesUnordered<tokio::task::JoinHandle<QueryResult>> {
476 deps.router
477 .backend_ids_in_tier(tier)
478 .filter_map(|id| ArticleBackend::from_availability(id, availability))
479 .map(|backend| {
480 let deps = deps.clone();
481 let request = request.clone();
482 tokio::spawn(async move { query_backend(&deps, backend, &request).await })
483 })
484 .collect()
485}
486
487#[cfg(test)]
493fn summarize(results: Vec<QueryResult>) -> (Option<(BackendId, PrecheckHit)>, ArticleAvailability) {
494 let mut availability = ArticleAvailability::new();
495 let mut found = None;
496
497 for r in results {
498 match r {
499 QueryResult::Found(id, response) => {
500 if found.is_none() {
501 found = Some((id, response));
502 }
503 }
504 QueryResult::Missing(id) => {
505 availability.record_missing(id);
506 }
507 QueryResult::Error => {}
508 }
509 }
510
511 (found, availability)
512}
513
514pub(crate) async fn precheck(
526 deps: &PrecheckDeps<'_>,
527 request: &RequestContext,
528 msg_id: &MessageId<'_>,
529 availability: &ArticleAvailability,
530) -> Option<PrecheckResponse> {
531 if let Some(cached) = deps.cache.get(msg_id).await
533 && cached.is_complete_article()
534 {
535 return Some(PrecheckResponse::Cached(cached));
536 }
537
538 let owned = deps.clone_deps();
539 let outcome = query_all_backends_racing(&owned, request, msg_id, availability).await;
540
541 match outcome {
543 RacingQueryOutcome::Hit(backend, data) => {
544 match data {
550 PrecheckHit::Payload(response) => {
551 if owned.cache.stores_payload_responses() {
552 let tier = crate::cache::ttl::CacheTier::new(0);
553 cache_precheck_hit(
554 &owned.cache,
555 msg_id.to_owned(),
556 backend,
557 PrecheckHit::Payload(response),
558 tier,
559 )
560 .await;
561 owned.cache.get(msg_id).await.map(PrecheckResponse::Cached)
562 } else {
563 Some(PrecheckResponse::Direct(response))
564 }
565 }
566 PrecheckHit::Availability(status_code) => {
567 if owned.cache.records_backend_has_status() {
568 let tier = crate::cache::ttl::CacheTier::new(0);
569 cache_precheck_hit(
570 &owned.cache,
571 msg_id.to_owned(),
572 backend,
573 PrecheckHit::Availability(status_code),
574 tier,
575 )
576 .await;
577 }
578 None
579 }
580 }
581 }
582 RacingQueryOutcome::Missing | RacingQueryOutcome::Inconclusive => None,
583 }
584}
585
586pub fn spawn_background_precheck(
590 deps: PrecheckDeps<'_>,
591 request: RequestContext,
592 msg_id: MessageId<'static>,
593) {
594 let owned = deps.clone_deps();
595 tokio::spawn(async move {
596 if let Some(cached) = owned.cache.get(&msg_id).await
598 && cached.is_complete_article()
599 {
600 return;
602 }
603
604 let availability = owned
605 .cache
606 .get(&msg_id)
607 .await
608 .map(|entry| entry.to_availability(owned.router.backend_count()))
609 .unwrap_or_default();
610 let results = query_all_backends(&owned, &request, &availability).await;
611 let mut found = None;
612
613 for result in results {
614 if let Some((backend_id, data)) =
615 record_precheck_result(&owned.cache, msg_id.to_owned(), result).await
616 {
617 if found.is_none() {
618 found = Some((backend_id, data));
619 } else {
620 let tier = crate::cache::ttl::CacheTier::new(0);
621 cache_precheck_hit(&owned.cache, msg_id.to_owned(), backend_id, data, tier)
622 .await;
623 }
624 }
625 }
626
627 if let Some((backend_id, data)) = found {
628 let tier = crate::cache::ttl::CacheTier::new(0);
632 if data.will_update_cache(&owned.cache) {
633 cache_precheck_hit(&owned.cache, msg_id.to_owned(), backend_id, data, tier).await;
634 }
635 }
636 });
637}
638
639#[cfg(test)]
640mod tests {
641 use super::*;
642 use std::sync::Arc;
643 use std::time::Duration;
644
645 use crate::cache::UnifiedCache;
646 use crate::metrics::MetricsCollector;
647 use crate::pool::BufferPool;
648 use crate::router::BackendSelector;
649 use crate::types::BufferSize;
650
651 fn eligible(backend_id: BackendId) -> ArticleBackend {
652 let availability = ArticleAvailability::new();
653 ArticleBackend::from_availability(backend_id, &availability)
654 .expect("backend should be eligible")
655 }
656
657 fn test_owned_deps(num_backends: usize) -> OwnedDeps {
658 OwnedDeps {
659 router: Arc::new(BackendSelector::new()),
660 cache: Arc::new(UnifiedCache::memory(100, Duration::from_secs(60))),
661 buffer_pool: BufferPool::new(BufferSize::try_new(4096).unwrap(), 1),
662 metrics: MetricsCollector::new(num_backends),
663 cache_articles: true,
664 }
665 }
666
667 #[test]
668 fn summarize_finds_first() {
669 let results = vec![
670 QueryResult::Missing(BackendId::from_index(0)),
671 QueryResult::Found(
672 BackendId::from_index(1),
673 PrecheckHit::Payload(b"first".to_vec().into()),
674 ),
675 QueryResult::Found(
676 BackendId::from_index(2),
677 PrecheckHit::Payload(b"second".to_vec().into()),
678 ),
679 ];
680 let (found, avail) = summarize(results);
681 assert_eq!(
682 found,
683 Some((
684 BackendId::from_index(1),
685 PrecheckHit::Payload(crate::cache::CacheIngestResponse::from(b"first".to_vec()))
686 ))
687 );
688 assert!(avail.is_missing(BackendId::from_index(0)));
689 assert!(!avail.is_missing(BackendId::from_index(1)));
690 assert!(!avail.is_missing(BackendId::from_index(2)));
691 }
692
693 #[test]
694 fn summarize_all_missing() {
695 let results = vec![
696 QueryResult::Missing(BackendId::from_index(0)),
697 QueryResult::Missing(BackendId::from_index(1)),
698 ];
699 let (found, avail) = summarize(results);
700 assert!(found.is_none());
701 assert!(avail.is_missing(BackendId::from_index(0)));
702 assert!(avail.is_missing(BackendId::from_index(1)));
703 }
704
705 #[test]
706 fn summarize_empty() {
707 let (found, avail) = summarize(vec![]);
708 assert!(found.is_none());
709 assert_eq!(avail.missing_bits(), 0);
710 }
711
712 #[tokio::test]
713 async fn record_precheck_result_records_missing_fact_immediately() {
714 let cache = UnifiedCache::memory(100, Duration::from_secs(60));
715 let msg_id = MessageId::from_borrowed("<precheck-missing@example.com>").unwrap();
716
717 let found = record_precheck_result(
718 &cache,
719 msg_id.to_owned(),
720 QueryResult::Missing(BackendId::from_index(0)),
721 )
722 .await;
723
724 assert!(found.is_none());
725 let cached = cache
726 .get(&msg_id)
727 .await
728 .expect("missing result should update authoritative availability");
729 assert!(!cached.should_try_backend(BackendId::from_index(0)));
730 }
731
732 #[tokio::test]
733 async fn record_precheck_result_does_not_record_transport_error_as_missing() {
734 let cache = UnifiedCache::memory(100, Duration::from_secs(60));
735 let msg_id = MessageId::from_borrowed("<precheck-error@example.com>").unwrap();
736
737 let found = record_precheck_result(&cache, msg_id.to_owned(), QueryResult::Error).await;
738
739 assert!(found.is_none());
740 assert!(cache.get(&msg_id).await.is_none());
741 }
742
743 #[test]
744 fn summarize_tier_results_treats_empty_tier_as_exhausted() {
745 let summary = summarize_tier_results(&[]);
746
747 assert!(matches!(
748 summary,
749 TierQuerySummary::Exhausted(availability) if availability.missing_bits() == 0
750 ));
751 }
752
753 #[test]
754 fn summarize_tier_results_requires_every_backend_to_miss() {
755 let summary = summarize_tier_results(&[
756 QueryResult::Missing(BackendId::from_index(0)),
757 QueryResult::Missing(BackendId::from_index(1)),
758 ]);
759
760 let TierQuerySummary::Exhausted(availability) = summary else {
761 panic!("all-missing tier should be exhausted");
762 };
763 assert!(availability.is_missing(BackendId::from_index(0)));
764 assert!(availability.is_missing(BackendId::from_index(1)));
765 }
766
767 #[test]
768 fn summarize_tier_results_treats_partial_missing_with_error_as_inconclusive() {
769 let summary = summarize_tier_results(&[
770 QueryResult::Missing(BackendId::from_index(0)),
771 QueryResult::Error,
772 ]);
773
774 let TierQuerySummary::Inconclusive(availability) = summary else {
775 panic!("tier with backend error should be inconclusive");
776 };
777 assert!(availability.is_missing(BackendId::from_index(0)));
778 assert!(!availability.is_missing(BackendId::from_index(1)));
779 }
780
781 #[test]
782 fn classify_precheck_result_records_non_430_4xx_as_error() {
783 let backend_id = BackendId::from_index(0);
784 let deps = test_owned_deps(1);
785
786 let result = classify_precheck_result(
787 &deps,
788 &eligible(backend_id),
789 StatusCode::new(412),
790 Some((11, 22, 33)),
791 PrecheckHit::Availability(StatusCode::new(412)),
792 );
793
794 assert!(matches!(result, QueryAttemptResult::Error));
795 let snapshot = deps.metrics.snapshot(None);
796 let backend = &snapshot.backend_stats[0];
797 assert_eq!(backend.total_commands.get(), 1);
798 assert_eq!(backend.ttfb_count.get(), 1);
799 assert_eq!(backend.errors_4xx.get(), 1);
800 assert_eq!(backend.errors_5xx.get(), 0);
801 assert_eq!(backend.errors.get(), 1);
802 }
803
804 #[test]
805 fn classify_precheck_result_records_5xx_as_error() {
806 let backend_id = BackendId::from_index(0);
807 let deps = test_owned_deps(1);
808
809 let result = classify_precheck_result(
810 &deps,
811 &eligible(backend_id),
812 StatusCode::new(502),
813 None,
814 PrecheckHit::Availability(StatusCode::new(502)),
815 );
816
817 assert!(matches!(result, QueryAttemptResult::Error));
818 let snapshot = deps.metrics.snapshot(None);
819 let backend = &snapshot.backend_stats[0];
820 assert_eq!(backend.total_commands.get(), 1);
821 assert_eq!(backend.errors_4xx.get(), 0);
822 assert_eq!(backend.errors_5xx.get(), 1);
823 assert_eq!(backend.errors.get(), 1);
824 }
825
826 async fn spawn_truncated_precheck_server() -> std::net::SocketAddr {
827 use tokio::io::{AsyncReadExt, AsyncWriteExt};
828 use tokio::net::TcpListener;
829
830 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
831 let addr = listener.local_addr().unwrap();
832
833 tokio::spawn(async move {
834 loop {
835 if let Ok((mut stream, _)) = listener.accept().await {
836 tokio::spawn(async move {
837 let _ = stream.write_all(b"200 mock\r\n").await;
838 let mut buf = [0u8; 1024];
839 let _ = stream.read(&mut buf).await;
840 let _ = stream
841 .write_all(b"220 0 <test@example.com>\r\npartial body")
842 .await;
843 let _ = stream.shutdown().await;
844 });
845 }
846 }
847 });
848
849 addr
850 }
851
852 async fn spawn_extra_response_precheck_server() -> std::net::SocketAddr {
853 use tokio::io::{AsyncReadExt, AsyncWriteExt};
854 use tokio::net::TcpListener;
855
856 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
857 let addr = listener.local_addr().unwrap();
858
859 tokio::spawn(async move {
860 loop {
861 if let Ok((mut stream, _)) = listener.accept().await {
862 tokio::spawn(async move {
863 let _ = stream.write_all(b"200 mock\r\n").await;
864 let mut buf = [0u8; 1024];
865 let _ = stream.read(&mut buf).await;
866 let _ = stream
867 .write_all(
868 b"220 0 <test@example.com>\r\nbody\r\n.\r\n430 No such article\r\n",
869 )
870 .await;
871 let _ = stream.shutdown().await;
872 });
873 }
874 }
875 });
876
877 addr
878 }
879
880 async fn spawn_large_article_precheck_server() -> std::net::SocketAddr {
881 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
882 use tokio::net::TcpListener;
883
884 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
885 let addr = listener.local_addr().unwrap();
886
887 tokio::spawn(async move {
888 loop {
889 if let Ok((mut stream, _)) = listener.accept().await {
890 tokio::spawn(async move {
891 let _ = stream.write_all(b"200 mock\r\n").await;
892 let mut reader = BufReader::new(stream);
893 let mut command = Vec::new();
894 loop {
895 command.clear();
896 if reader.read_until(b'\n', &mut command).await.is_err() {
897 return;
898 }
899 if command.starts_with(b"ARTICLE ") {
900 break;
901 }
902 let response = if command.starts_with(b"DATE") {
903 b"111 20260526120000\r\n".as_slice()
904 } else {
905 b"200 mock setup complete\r\n".as_slice()
906 };
907 if reader.get_mut().write_all(response).await.is_err() {
908 return;
909 }
910 }
911 let oversized_payload = vec![b'x'; 4 * 1024 * 1024 + 1];
912 let stream = reader.get_mut();
913 let _ = stream.write_all(b"220 0 <test@example.com>\r\n").await;
914 let _ = stream.write_all(&oversized_payload).await;
915 let _ = stream.write_all(b"\r\n.\r\n").await;
916 });
917 }
918 }
919 });
920
921 addr
922 }
923
924 fn selector_with_backend(
925 addr: std::net::SocketAddr,
926 max_connections: usize,
927 ) -> (BackendSelector, BackendId) {
928 use crate::pool::DeadpoolConnectionProvider;
929 use crate::types::ServerName;
930
931 let mut selector = BackendSelector::new();
932 let backend_id = BackendId::from_index(0);
933 let provider = DeadpoolConnectionProvider::new(
934 "127.0.0.1".to_string(),
935 addr.port(),
936 "test".to_string(),
937 max_connections,
938 None,
939 None,
940 );
941 selector.add_backend(
942 ServerName::try_new("test-server".to_string()).unwrap(),
943 provider,
944 0,
945 );
946 (selector, backend_id)
947 }
948
949 #[tokio::test]
950 async fn query_backend_returns_error_on_truncated_response_body() {
951 let addr = spawn_truncated_precheck_server().await;
952
953 let (selector, backend_id) = selector_with_backend(addr, 2);
954
955 let deps = OwnedDeps {
956 router: Arc::new(selector),
957 cache: Arc::new(UnifiedCache::memory(100, Duration::from_secs(60))),
958 buffer_pool: BufferPool::new(BufferSize::try_new(4096).unwrap(), 2),
959 metrics: MetricsCollector::new(1),
960 cache_articles: true,
961 };
962
963 let request =
964 RequestContext::parse(b"ARTICLE <test@example.com>\r\n").expect("valid request line");
965 let result = query_backend(&deps, eligible(backend_id), &request).await;
966 assert_eq!(result, QueryResult::Error);
967 }
968
969 #[tokio::test]
970 async fn query_backend_returns_error_on_extra_precheck_response() {
971 let addr = spawn_extra_response_precheck_server().await;
972
973 let (selector, backend_id) = selector_with_backend(addr, 2);
974
975 let deps = OwnedDeps {
976 router: Arc::new(selector),
977 cache: Arc::new(UnifiedCache::memory(100, Duration::from_secs(60))),
978 buffer_pool: BufferPool::new(BufferSize::try_new(4096).unwrap(), 2),
979 metrics: MetricsCollector::new(1),
980 cache_articles: true,
981 };
982
983 let request =
984 RequestContext::parse(b"ARTICLE <test@example.com>\r\n").expect("valid request line");
985 let result = query_backend(&deps, eligible(backend_id), &request).await;
986 assert_eq!(result, QueryResult::Error);
987 }
988
989 #[tokio::test]
990 async fn query_backend_treats_oversized_precheck_hit_as_availability() {
991 let addr = spawn_large_article_precheck_server().await;
992 let (selector, backend_id) = selector_with_backend(addr, 1);
993
994 let deps = OwnedDeps {
995 router: Arc::new(selector),
996 cache: Arc::new(UnifiedCache::memory(100, Duration::from_secs(60))),
997 buffer_pool: BufferPool::new(BufferSize::try_new(65536).unwrap(), 2),
998 metrics: MetricsCollector::new(1),
999 cache_articles: true,
1000 };
1001
1002 let request =
1003 RequestContext::parse(b"ARTICLE <test@example.com>\r\n").expect("valid request line");
1004 let result = query_backend(&deps, eligible(backend_id), &request).await;
1005
1006 assert_eq!(
1007 result,
1008 QueryResult::Found(backend_id, PrecheckHit::Availability(StatusCode::new(220)))
1009 );
1010 }
1011
1012 #[tokio::test]
1013 async fn cache_precheck_availability_hit_does_not_persist_positive_entry() {
1014 use crate::types::BackendId;
1015
1016 let backend_id = BackendId::from_index(0);
1017 let cache = UnifiedCache::availability(std::time::Duration::MAX);
1018 let msg_id = MessageId::new("<test@example.com>".to_string()).unwrap();
1019
1020 cache_precheck_hit(
1021 &cache,
1022 msg_id.to_owned(),
1023 backend_id,
1024 PrecheckHit::Availability(StatusCode::new(223)),
1025 crate::cache::ttl::CacheTier::new(0),
1026 )
1027 .await;
1028
1029 assert!(
1030 cache.get(&msg_id).await.is_none(),
1031 "availability-only index should persist negatives, not optimistic positive hits"
1032 );
1033 }
1034}