#![allow(dead_code)]
use std::{
future::{Future, poll_fn},
pin::Pin,
sync::Arc,
task::{Context, Poll},
time::Duration,
};
use saddle_admission::{
AdmissionError, AuditReport, DbCreditSnapshot, DbPermitDomain, DbRequestPermit, DbRequestPhase,
DbRouteResources, EntryIoAuditPlan, EntryReadPoll, HealthFailure, ManagedBytes,
ManagedResponse, OfficialTokioDomain, OfficialTokioEntryIoAttemptOutcome,
OfficialTokioRegistrationSnapshot, OfficialTokioTaskSlotOwner, ProcessAllocationProfile,
ProcessAllocationSnapshot, ProcessLedger, ProcessSnapshot, RequestMemory, RequestSignalHandle,
RequestSignalSet, RequestSignalStatus, ResponseWritePoll, RetryHint, WaitRegistration,
WaitRemoval,
};
use saddle_core::SaddleError;
use tokio::{
io::{AsyncRead, AsyncWrite, ReadBuf},
net::TcpStream,
task::{JoinError, JoinHandle},
time::{Instant as TokioInstant, sleep},
};
use crate::{
RequestLifecycle,
compiled_route::{
ClassifiedCompiledRouteAdapter, CompiledResponseOutcome, CompiledRouteAdapter,
},
request::{RequestClaim, RequestGuard},
response_framing::{FixedResponseWrite, ResponseFramingAdapter},
};
const TASK_STORAGE_BOUND: usize = 384;
type TaskOutput = Result<AuditReport, TaskExecutionError>;
#[derive(Debug)]
enum TaskExecutionError {
Admission(AdmissionError),
Deadline,
}
struct TaskEntry {
handle: JoinHandle<TaskOutput>,
lifecycle: RequestGuard,
signal: Option<RequestSignalHandle>,
official_task_slot: Option<OfficialTokioTaskSlotOwner>,
}
#[derive(Debug)]
enum SubmitError {
RuntimeUnavailable,
RegistryFull,
ProcessUnhealthy(HealthFailure),
Lifecycle(SaddleError),
Admission(AdmissionError),
RouteAdapter,
}
#[derive(Debug)]
enum TaskFailure {
Admission(AdmissionError),
Deadline,
Cancelled,
Panicked,
}
#[derive(Debug)]
struct ShutdownReport {
ledger: ProcessSnapshot,
first_task_failure: Option<TaskFailure>,
}
struct ManagedRequestTasks {
requests: RequestLifecycle,
ledger: Option<ProcessLedger>,
tasks: Vec<TaskEntry>,
capacity: usize,
official_tokio_domain: Option<OfficialTokioDomain>,
process_allocation_profile: Option<ProcessAllocationProfile>,
shutdown_complete: bool,
}
impl ManagedRequestTasks {
fn new(requests: RequestLifecycle, ledger: ProcessLedger, capacity: usize) -> Self {
Self {
requests,
ledger: Some(ledger),
tasks: Vec::with_capacity(capacity),
capacity,
official_tokio_domain: None,
process_allocation_profile: None,
shutdown_complete: false,
}
}
fn with_official_tokio(
requests: RequestLifecycle,
ledger: ProcessLedger,
capacity: usize,
domain: OfficialTokioDomain,
profile: ProcessAllocationProfile,
) -> Self {
Self {
requests,
ledger: Some(ledger),
tasks: Vec::with_capacity(capacity),
capacity,
official_tokio_domain: Some(domain),
process_allocation_profile: Some(profile),
shutdown_complete: false,
}
}
fn try_submit<F>(
&mut self,
managed_limit: usize,
factory: impl FnOnce(&RequestMemory) -> F,
) -> Result<(), SubmitError>
where
F: Future + Send + 'static,
{
if tokio::runtime::Handle::try_current().is_err() {
return Err(SubmitError::RuntimeUnavailable);
}
if self.tasks.len() == self.capacity {
return Err(SubmitError::RegistryFull);
}
let ledger = self
.ledger
.as_ref()
.expect("task manager owns ledger until shutdown");
let health = ledger.health();
if !health.healthy {
return Err(SubmitError::ProcessUnhealthy(health.failure));
}
let claim = self.requests.try_claim().map_err(SubmitError::Lifecycle)?;
let envelope = ledger
.try_envelope(managed_limit, TASK_STORAGE_BOUND, factory)
.map_err(SubmitError::Admission)?;
let lifecycle = RequestClaim::publish(claim);
let handle =
tokio::spawn(async move { envelope.await.map_err(TaskExecutionError::Admission) });
debug_assert!(self.tasks.len() < self.tasks.capacity());
self.tasks.push(TaskEntry {
handle,
lifecycle,
signal: None,
official_task_slot: None,
});
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn try_submit_official_tcp<M, B>(
&mut self,
managed_limit: usize,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
db_resources: DbRouteResources<'_>,
socket: TcpStream,
expected_body_bytes: usize,
deadline: Duration,
make_business: M,
) -> TcpSubmitOutcome
where
M: FnOnce(ManagedBytes, Option<DbRequestPermit>, &RequestMemory) -> B
+ Send
+ Unpin
+ 'static,
B: Future<Output = ManagedResponse> + Send + 'static,
{
self.try_submit_official_tcp_with_deadlines(
managed_limit,
task_storage_bound,
read_audit_plan,
write_audit_plan,
db_resources,
socket,
expected_body_bytes,
deadline,
deadline,
make_business,
)
}
#[allow(clippy::too_many_arguments)]
fn try_submit_official_tcp_with_deadlines<M, B>(
&mut self,
managed_limit: usize,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
db_resources: DbRouteResources<'_>,
socket: TcpStream,
expected_body_bytes: usize,
attempt_deadline: Duration,
termination_bound: Duration,
make_business: M,
) -> TcpSubmitOutcome
where
M: FnOnce(ManagedBytes, Option<DbRequestPermit>, &RequestMemory) -> B
+ Send
+ Unpin
+ 'static,
B: Future<Output = ManagedResponse> + Send + 'static,
{
self.try_submit_official_tcp_with_write(
managed_limit,
task_storage_bound,
read_audit_plan,
write_audit_plan,
db_resources,
socket,
TcpWriteState { written: 0 },
expected_body_bytes,
attempt_deadline,
termination_bound,
make_business,
)
}
#[allow(clippy::too_many_arguments)]
fn try_submit_official_tcp_with_write<M, B, O, W>(
&mut self,
managed_limit: usize,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
db_resources: DbRouteResources<'_>,
socket: TcpStream,
write_state: W,
expected_body_bytes: usize,
attempt_deadline: Duration,
termination_bound: Duration,
make_business: M,
) -> TcpSubmitOutcome
where
M: FnOnce(ManagedBytes, Option<DbRequestPermit>, &RequestMemory) -> B
+ Send
+ Unpin
+ 'static,
B: Future<Output = O> + Send + 'static,
O: saddle_admission::ManagedResponseOutput + Send + Unpin + 'static,
W: ResponseWritePoll<TcpConnection, O> + Send + 'static,
{
if tokio::runtime::Handle::try_current().is_err() {
return TcpSubmitOutcome::Stop(Some(socket), SubmitError::RuntimeUnavailable);
}
let ledger = self
.ledger
.as_ref()
.expect("task manager owns ledger until shutdown");
let health = ledger.health();
if !health.healthy {
return TcpSubmitOutcome::Stop(
Some(socket),
SubmitError::ProcessUnhealthy(health.failure),
);
}
let domain = self
.official_tokio_domain
.as_ref()
.expect("official Tokio submit requires its unique domain owner");
let claim = match self.requests.try_claim() {
Ok(claim) => claim,
Err(error) => {
return TcpSubmitOutcome::Stop(Some(socket), SubmitError::Lifecycle(error));
}
};
let mut socket = Some(socket);
let admission =
ledger.attempt_official_tokio_entry_io(
domain,
db_resources,
managed_limit,
task_storage_bound,
read_audit_plan,
write_audit_plan,
|memory| {
(
TcpConnection {
socket: socket.take().expect("accepted factory takes socket once"),
},
TcpReadState {
payload: Some(memory.try_bytes(&[]).expect(
"empty managed payload is covered by the admitted account",
)),
expected_body_bytes,
scratch: [0; 1_024],
},
write_state,
)
},
make_business,
);
match admission {
OfficialTokioEntryIoAttemptOutcome::Ready(ready) => {
let (envelope, task_slot) = ready.into_runtime_parts();
let signal = envelope.request_signal_handle();
let timer_ledger = ledger.clone();
let lifecycle = RequestClaim::publish(claim);
let handle = tokio::spawn(async move {
let timer = sleep(attempt_deadline);
tokio::pin!(timer);
tokio::pin!(envelope);
let mut terminating = false;
let mut attempt_expired = false;
loop {
let envelope_or_phase = poll_fn(|context| {
if !terminating && request_requires_termination(&timer_ledger, signal) {
return Poll::Ready(None);
}
match envelope.as_mut().poll(context) {
Poll::Ready(result) => Poll::Ready(Some(result)),
Poll::Pending
if !terminating
&& request_requires_termination(&timer_ledger, signal) =>
{
Poll::Ready(None)
}
Poll::Pending => Poll::Pending,
}
});
tokio::pin!(envelope_or_phase);
tokio::select! {
biased;
event = &mut envelope_or_phase => {
if let Some(result) = event {
return if attempt_expired {
match result {
Ok(_) | Err(AdmissionError::RequestCancelled) => {
Err(TaskExecutionError::Deadline)
}
Err(error) => Err(TaskExecutionError::Admission(error)),
}
} else {
result.map_err(TaskExecutionError::Admission)
};
}
terminating = true;
timer.as_mut().reset(
TokioInstant::now()
.checked_add(termination_bound)
.unwrap_or_else(|| std::process::abort()),
);
}
() = &mut timer => {
if terminating {
std::process::abort();
}
match timer_ledger.request_cancel(signal) {
RequestSignalSet::Set | RequestSignalSet::AlreadySet => {}
RequestSignalSet::Stale | RequestSignalSet::Unhealthy => {
std::process::abort();
}
}
attempt_expired = true;
terminating = true;
timer.as_mut().reset(
TokioInstant::now()
.checked_add(termination_bound)
.unwrap_or_else(|| std::process::abort()),
);
}
}
}
});
debug_assert!(self.tasks.len() < self.tasks.capacity());
self.tasks.push(TaskEntry {
handle,
lifecycle,
signal: Some(signal),
official_task_slot: Some(task_slot),
});
TcpSubmitOutcome::Accepted
}
OfficialTokioEntryIoAttemptOutcome::Registered(registration) => {
drop(claim);
TcpSubmitOutcome::Registered(RegisteredTcpInput {
socket: socket.expect("registered does not invoke the Entry I/O factory"),
registration,
})
}
OfficialTokioEntryIoAttemptOutcome::Reject(error) => {
drop(claim);
TcpSubmitOutcome::Reject(
socket.expect("reject does not invoke the Entry I/O factory"),
SubmitError::Admission(error),
)
}
OfficialTokioEntryIoAttemptOutcome::Stop(error) => {
drop(claim);
TcpSubmitOutcome::Stop(socket, SubmitError::Admission(error))
}
}
}
#[allow(clippy::too_many_arguments)]
fn try_submit_compiled_tcp<A>(
&mut self,
adapter: Arc<A>,
proof: A::Proof,
context: A::Context,
db_domain: Option<&DbPermitDomain>,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
socket: TcpStream,
expected_body_bytes: usize,
deadline: Duration,
) -> TcpSubmitOutcome
where
A: CompiledRouteAdapter,
{
self.try_submit_compiled_tcp_with_deadlines(
adapter,
proof,
context,
db_domain,
task_storage_bound,
read_audit_plan,
write_audit_plan,
socket,
expected_body_bytes,
deadline,
deadline,
)
}
#[allow(clippy::too_many_arguments)]
fn try_submit_compiled_tcp_with_deadlines<A>(
&mut self,
adapter: Arc<A>,
proof: A::Proof,
context: A::Context,
db_domain: Option<&DbPermitDomain>,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
socket: TcpStream,
expected_body_bytes: usize,
attempt_deadline: Duration,
termination_bound: Duration,
) -> TcpSubmitOutcome
where
A: CompiledRouteAdapter,
{
let managed_limit = match adapter.managed_commitment(proof) {
Ok(value) => value,
Err(_) => {
return TcpSubmitOutcome::Reject(socket, SubmitError::RouteAdapter);
}
};
let response_capacity = match adapter.response_capacity(proof) {
Ok(value) if value <= managed_limit => value,
Ok(_) | Err(_) => {
return TcpSubmitOutcome::Reject(socket, SubmitError::RouteAdapter);
}
};
let db_resources = match adapter.db_resources(proof, db_domain) {
Ok(value) => value,
Err(_) => {
return TcpSubmitOutcome::Stop(Some(socket), SubmitError::RouteAdapter);
}
};
self.try_submit_official_tcp_with_deadlines(
managed_limit,
task_storage_bound,
read_audit_plan,
write_audit_plan,
db_resources,
socket,
expected_body_bytes,
attempt_deadline,
termination_bound,
move |body, permit, memory| {
let error_response = memory
.try_response(&[])
.expect("legacy error response is inside the route commitment");
let future = adapter.execute(proof, context, body, permit, memory);
async move {
let _approved_capacity = response_capacity;
match future {
Ok(future) => future.await.unwrap_or(error_response),
Err(_) => error_response,
}
}
},
)
}
#[allow(clippy::too_many_arguments)]
fn try_submit_compiled_http<A, H>(
&mut self,
adapter: Arc<A>,
proof: A::Proof,
context: A::Context,
framing: Arc<H>,
framing_plan: H::Plan,
db_domain: Option<&DbPermitDomain>,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
socket: TcpStream,
expected_body_bytes: usize,
deadline: Duration,
) -> TcpSubmitOutcome
where
A: ClassifiedCompiledRouteAdapter,
H: ResponseFramingAdapter,
{
self.try_submit_compiled_http_with_deadlines(
adapter,
proof,
context,
framing,
framing_plan,
db_domain,
task_storage_bound,
read_audit_plan,
write_audit_plan,
socket,
expected_body_bytes,
deadline,
deadline,
)
}
#[allow(clippy::too_many_arguments)]
fn try_submit_compiled_http_with_deadlines<A, H>(
&mut self,
adapter: Arc<A>,
proof: A::Proof,
context: A::Context,
framing: Arc<H>,
framing_plan: H::Plan,
db_domain: Option<&DbPermitDomain>,
task_storage_bound: usize,
read_audit_plan: EntryIoAuditPlan,
write_audit_plan: EntryIoAuditPlan,
socket: TcpStream,
expected_body_bytes: usize,
attempt_deadline: Duration,
termination_bound: Duration,
) -> TcpSubmitOutcome
where
A: ClassifiedCompiledRouteAdapter,
H: ResponseFramingAdapter,
{
let managed_limit = match adapter.managed_commitment(proof) {
Ok(value) => value,
Err(_) => return TcpSubmitOutcome::Reject(socket, SubmitError::RouteAdapter),
};
let response_capacity = match adapter.response_capacity(proof) {
Ok(value) if value <= managed_limit => value,
Ok(_) | Err(_) => {
return TcpSubmitOutcome::Reject(socket, SubmitError::RouteAdapter);
}
};
let db_resources = match adapter.db_resources(proof, db_domain) {
Ok(value) => value,
Err(_) => return TcpSubmitOutcome::Stop(Some(socket), SubmitError::RouteAdapter),
};
let write_state = match framing.begin(framing_plan, response_capacity) {
Ok(state) => FramedTcpWriteState { inner: state },
Err(_) => return TcpSubmitOutcome::Reject(socket, SubmitError::RouteAdapter),
};
self.try_submit_official_tcp_with_write(
managed_limit,
task_storage_bound,
read_audit_plan,
write_audit_plan,
db_resources,
socket,
write_state,
expected_body_bytes,
attempt_deadline,
termination_bound,
move |body, permit, memory| {
let future = adapter.execute(proof, context, body, permit, memory);
async move {
let _approved_capacity = response_capacity;
future.await
}
},
)
}
fn claim_retry_hint(&self) -> Option<RetryHint> {
self.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.claim_retry_hint()
}
fn cancel_registered(&self, input: RegisteredTcpInput) -> (TcpStream, WaitRemoval) {
let removal = self
.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.cancel_wait(input.registration);
(input.socket, removal)
}
async fn reap_finished(&mut self) -> Option<TaskFailure> {
let mut first_failure = None;
let mut index = 0;
while index < self.tasks.len() {
if self.tasks[index].handle.is_finished() {
let TaskEntry {
handle,
lifecycle,
signal: _,
official_task_slot,
} = self.tasks.swap_remove(index);
record_failure(&mut first_failure, task_result(handle.await));
drop(lifecycle);
drop(official_task_slot);
} else {
index += 1;
}
}
first_failure
}
async fn shutdown(mut self, cancel: bool) -> Result<ShutdownReport, AdmissionError> {
assert!(
self.official_tokio_domain.is_none() && self.process_allocation_profile.is_none(),
"official Tokio manager requires two-stage driver-final shutdown"
);
self.requests.begin_draining();
if cancel {
for entry in &self.tasks {
entry.handle.abort();
}
}
let mut first_failure = None;
while let Some(entry) = self.tasks.pop() {
let TaskEntry {
handle,
lifecycle,
signal: _,
official_task_slot,
} = entry;
record_failure(&mut first_failure, task_result(handle.await));
drop(lifecycle);
drop(official_task_slot);
}
self.requests.wait_until_drained().await;
self.requests.mark_stopped();
let ledger = self
.ledger
.take()
.expect("shutdown consumes the final Runtime ledger owner");
let ledger = ledger.try_shutdown()?;
self.shutdown_complete = true;
Ok(ShutdownReport {
ledger,
first_task_failure: first_failure,
})
}
async fn shutdown_official_tokio(mut self, cancel: bool) -> DriverFinalizer {
self.requests.begin_draining();
self.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.stop_waiting();
drop(self.official_tokio_domain.take());
for entry in &self.tasks {
let Some(signal) = entry.signal else {
std::process::abort();
};
let result = if cancel {
self.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.request_cancel(signal)
} else {
self.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.request_shutdown(signal)
};
if !matches!(
result,
RequestSignalSet::Set | RequestSignalSet::AlreadySet | RequestSignalSet::Stale
) {
std::process::abort();
}
}
let mut first_task_failure = None;
while let Some(entry) = self.tasks.pop() {
let TaskEntry {
handle,
lifecycle,
signal: _,
official_task_slot,
} = entry;
record_failure(&mut first_task_failure, task_result(handle.await));
drop(lifecycle);
drop(official_task_slot);
}
self.requests.wait_until_drained().await;
self.requests.mark_stopped();
let ledger = self
.ledger
.take()
.expect("driver finalizer receives the unique ledger owner");
let profile = self
.process_allocation_profile
.take()
.expect("driver finalizer receives the process watermark owner");
self.shutdown_complete = true;
DriverFinalizer {
ledger: Some(ledger),
profile: Some(profile),
first_task_failure,
finished: false,
}
}
}
enum TcpSubmitOutcome {
Accepted,
Registered(RegisteredTcpInput),
Reject(TcpStream, SubmitError),
Stop(Option<TcpStream>, SubmitError),
}
struct RegisteredTcpInput {
socket: TcpStream,
registration: WaitRegistration,
}
impl RegisteredTcpInput {
fn matches(&self, hint: RetryHint) -> bool {
self.registration == hint.registration()
}
fn into_retry(self, hint: RetryHint) -> Result<TcpStream, Self> {
if self.matches(hint) {
Ok(self.socket)
} else {
Err(self)
}
}
}
pub struct OfficialCompiledRouteCoordinator<A>
where
A: CompiledRouteAdapter,
{
tasks: ManagedRequestTasks,
adapter: Arc<A>,
db_domain: Option<DbPermitDomain>,
accepted_generation: u64,
reaped_task_failed: bool,
}
#[derive(Clone, Copy)]
pub struct OfficialTcpAttemptProfile {
pub task_storage_bound: usize,
pub read_audit_plan: EntryIoAuditPlan,
pub write_audit_plan: EntryIoAuditPlan,
pub expected_body_bytes: usize,
pub deadline: Duration,
pub termination_bound: Duration,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct PublishedTcpIdentity {
generation: u64,
}
pub struct WaitingCompiledTcp<P, C> {
input: RegisteredTcpInput,
proof: P,
context: C,
}
pub struct CompiledRetryToken {
hint: RetryHint,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OfficialTcpRejectReason {
RuntimeCapacity,
Admission,
Route,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OfficialTcpStopReason {
Runtime,
Lifecycle,
Admission,
Route,
}
pub enum OfficialCompiledTcpOutcome<P, C> {
Accepted(PublishedTcpIdentity),
Registered(WaitingCompiledTcp<P, C>),
Reject(OfficialTcpRejectReason, TcpStream, P, C),
Stop(OfficialTcpStopReason, Option<TcpStream>, P, C),
}
pub struct OfficialCompiledDriverFinalizer {
inner: DriverFinalizer,
reaped_task_failed: bool,
}
pub struct OfficialCompiledRuntimeFinalizer {
inner: RuntimeDriverFinalizer,
reaped_task_failed: bool,
}
pub struct OfficialHttpRouteCoordinator<A, H>
where
A: ClassifiedCompiledRouteAdapter,
H: ResponseFramingAdapter,
{
tasks: ManagedRequestTasks,
adapter: Arc<A>,
framing: Arc<H>,
db_domain: Option<DbPermitDomain>,
accepted_generation: u64,
reaped_task_failed: bool,
}
pub struct WaitingHttpTcp<P, C, F> {
input: RegisteredTcpInput,
proof: P,
context: C,
framing_plan: F,
}
pub enum OfficialHttpTcpOutcome<P, C, F> {
Accepted(PublishedTcpIdentity),
Registered(WaitingHttpTcp<P, C, F>),
Reject(OfficialTcpRejectReason, TcpStream, P, C, F),
Stop(OfficialTcpStopReason, Option<TcpStream>, P, C, F),
}
#[derive(Debug)]
pub struct OfficialCompiledShutdownReport {
pub watermark: ProcessAllocationSnapshot,
pub ledger: ProcessSnapshot,
pub task_failed: bool,
}
#[doc(hidden)]
#[derive(Clone, Copy, Debug)]
pub struct OfficialCompiledLiveSnapshot {
pub watermark: ProcessAllocationSnapshot,
pub ledger: ProcessSnapshot,
pub registration: OfficialTokioRegistrationSnapshot,
pub db: Option<DbCreditSnapshot>,
pub task_entries: usize,
}
impl<A> OfficialCompiledRouteCoordinator<A>
where
A: CompiledRouteAdapter,
A::Context: Copy,
{
pub fn new(
adapter: Arc<A>,
ledger: ProcessLedger,
task_capacity: usize,
domain: OfficialTokioDomain,
profile: ProcessAllocationProfile,
db_domain: Option<DbPermitDomain>,
) -> Self {
let requests = RequestLifecycle::new();
requests.mark_ready();
Self {
tasks: ManagedRequestTasks::with_official_tokio(
requests,
ledger,
task_capacity,
domain,
profile,
),
adapter,
db_domain,
accepted_generation: 0,
reaped_task_failed: false,
}
}
pub fn attempt(
&mut self,
proof: A::Proof,
context: A::Context,
socket: TcpStream,
profile: OfficialTcpAttemptProfile,
) -> OfficialCompiledTcpOutcome<A::Proof, A::Context> {
if profile.deadline.is_zero() || profile.termination_bound.is_zero() {
return OfficialCompiledTcpOutcome::Stop(
OfficialTcpStopReason::Route,
Some(socket),
proof,
context,
);
}
let outcome = self.tasks.try_submit_compiled_tcp_with_deadlines(
Arc::clone(&self.adapter),
proof,
context,
self.db_domain.as_ref(),
profile.task_storage_bound,
profile.read_audit_plan,
profile.write_audit_plan,
socket,
profile.expected_body_bytes,
profile.deadline,
profile.termination_bound,
);
match outcome {
TcpSubmitOutcome::Accepted => {
self.accepted_generation = self
.accepted_generation
.checked_add(1)
.unwrap_or_else(|| std::process::abort());
OfficialCompiledTcpOutcome::Accepted(PublishedTcpIdentity {
generation: self.accepted_generation,
})
}
TcpSubmitOutcome::Registered(input) => {
OfficialCompiledTcpOutcome::Registered(WaitingCompiledTcp {
input,
proof,
context,
})
}
TcpSubmitOutcome::Reject(socket, error) => {
OfficialCompiledTcpOutcome::Reject(reject_reason(&error), socket, proof, context)
}
TcpSubmitOutcome::Stop(socket, error) => {
OfficialCompiledTcpOutcome::Stop(stop_reason(&error), socket, proof, context)
}
}
}
pub fn claim_retry(&self) -> Option<CompiledRetryToken> {
self.tasks
.claim_retry_hint()
.map(|hint| CompiledRetryToken { hint })
}
pub fn retry(
&mut self,
waiting: WaitingCompiledTcp<A::Proof, A::Context>,
token: CompiledRetryToken,
profile: OfficialTcpAttemptProfile,
) -> OfficialCompiledTcpOutcome<A::Proof, A::Context> {
let WaitingCompiledTcp {
input,
proof,
context,
} = waiting;
match input.into_retry(token.hint) {
Ok(socket) => self.attempt(proof, context, socket, profile),
Err(input) => OfficialCompiledTcpOutcome::Stop(
OfficialTcpStopReason::Admission,
Some(input.socket),
proof,
context,
),
}
}
pub fn cancel(
&self,
waiting: WaitingCompiledTcp<A::Proof, A::Context>,
) -> (TcpStream, A::Proof, A::Context, bool) {
let (socket, removal) = self.tasks.cancel_registered(waiting.input);
(socket, waiting.proof, waiting.context, removal.is_removed())
}
pub async fn reap_finished(&mut self) -> bool {
let before = self.tasks.tasks.len();
self.reaped_task_failed |= self.tasks.reap_finished().await.is_some();
self.tasks.tasks.len() != before
}
pub fn stop_accepting(&self) {
self.tasks.requests.begin_draining();
self.tasks
.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.stop_waiting();
}
#[doc(hidden)]
pub fn live_snapshot(&self) -> Result<OfficialCompiledLiveSnapshot, AdmissionError> {
let ledger = self
.tasks
.ledger
.as_ref()
.ok_or(AdmissionError::InvalidConfiguration)?;
let domain = self
.tasks
.official_tokio_domain
.as_ref()
.ok_or(AdmissionError::OfficialTokioDomainStopped)?;
let profile = self
.tasks
.process_allocation_profile
.as_ref()
.ok_or(AdmissionError::ProcessAllocationProfileAlreadyActive)?;
Ok(OfficialCompiledLiveSnapshot {
watermark: profile.snapshot()?,
ledger: ledger.snapshot(),
registration: domain.snapshot()?,
db: self
.db_domain
.as_ref()
.map(DbPermitDomain::snapshot)
.transpose()?,
task_entries: self.tasks.tasks.len(),
})
}
pub async fn shutdown(mut self, cancel: bool) -> OfficialCompiledDriverFinalizer {
drop(self.db_domain.take());
OfficialCompiledDriverFinalizer {
inner: self.tasks.shutdown_official_tokio(cancel).await,
reaped_task_failed: self.reaped_task_failed,
}
}
}
impl<A, H> OfficialHttpRouteCoordinator<A, H>
where
A: ClassifiedCompiledRouteAdapter,
A::Context: Copy,
H: ResponseFramingAdapter,
{
#[allow(clippy::too_many_arguments)]
pub fn new(
adapter: Arc<A>,
framing: Arc<H>,
ledger: ProcessLedger,
task_capacity: usize,
domain: OfficialTokioDomain,
profile: ProcessAllocationProfile,
db_domain: Option<DbPermitDomain>,
) -> Self {
let requests = RequestLifecycle::new();
requests.mark_ready();
Self {
tasks: ManagedRequestTasks::with_official_tokio(
requests,
ledger,
task_capacity,
domain,
profile,
),
adapter,
framing,
db_domain,
accepted_generation: 0,
reaped_task_failed: false,
}
}
pub fn attempt(
&mut self,
proof: A::Proof,
context: A::Context,
framing_plan: H::Plan,
socket: TcpStream,
profile: OfficialTcpAttemptProfile,
) -> OfficialHttpTcpOutcome<A::Proof, A::Context, H::Plan> {
if profile.deadline.is_zero() || profile.termination_bound.is_zero() {
return OfficialHttpTcpOutcome::Stop(
OfficialTcpStopReason::Route,
Some(socket),
proof,
context,
framing_plan,
);
}
let outcome = self.tasks.try_submit_compiled_http_with_deadlines(
Arc::clone(&self.adapter),
proof,
context,
Arc::clone(&self.framing),
framing_plan,
self.db_domain.as_ref(),
profile.task_storage_bound,
profile.read_audit_plan,
profile.write_audit_plan,
socket,
profile.expected_body_bytes,
profile.deadline,
profile.termination_bound,
);
match outcome {
TcpSubmitOutcome::Accepted => {
self.accepted_generation = self
.accepted_generation
.checked_add(1)
.unwrap_or_else(|| std::process::abort());
OfficialHttpTcpOutcome::Accepted(PublishedTcpIdentity {
generation: self.accepted_generation,
})
}
TcpSubmitOutcome::Registered(input) => {
OfficialHttpTcpOutcome::Registered(WaitingHttpTcp {
input,
proof,
context,
framing_plan,
})
}
TcpSubmitOutcome::Reject(socket, error) => OfficialHttpTcpOutcome::Reject(
reject_reason(&error),
socket,
proof,
context,
framing_plan,
),
TcpSubmitOutcome::Stop(socket, error) => OfficialHttpTcpOutcome::Stop(
stop_reason(&error),
socket,
proof,
context,
framing_plan,
),
}
}
pub fn claim_retry(&self) -> Option<CompiledRetryToken> {
self.tasks
.claim_retry_hint()
.map(|hint| CompiledRetryToken { hint })
}
pub fn retry(
&mut self,
waiting: WaitingHttpTcp<A::Proof, A::Context, H::Plan>,
token: CompiledRetryToken,
profile: OfficialTcpAttemptProfile,
) -> OfficialHttpTcpOutcome<A::Proof, A::Context, H::Plan> {
let WaitingHttpTcp {
input,
proof,
context,
framing_plan,
} = waiting;
match input.into_retry(token.hint) {
Ok(socket) => self.attempt(proof, context, framing_plan, socket, profile),
Err(input) => OfficialHttpTcpOutcome::Stop(
OfficialTcpStopReason::Admission,
Some(input.socket),
proof,
context,
framing_plan,
),
}
}
pub fn cancel(
&self,
waiting: WaitingHttpTcp<A::Proof, A::Context, H::Plan>,
) -> (TcpStream, A::Proof, A::Context, H::Plan, bool) {
let (socket, removal) = self.tasks.cancel_registered(waiting.input);
(
socket,
waiting.proof,
waiting.context,
waiting.framing_plan,
removal.is_removed(),
)
}
pub async fn reap_finished(&mut self) -> bool {
let before = self.tasks.tasks.len();
self.reaped_task_failed |= self.tasks.reap_finished().await.is_some();
self.tasks.tasks.len() != before
}
pub fn stop_accepting(&self) {
self.tasks.requests.begin_draining();
self.tasks
.ledger
.as_ref()
.expect("task manager owns ledger until shutdown")
.stop_waiting();
}
pub async fn shutdown(mut self, cancel: bool) -> OfficialCompiledDriverFinalizer {
drop(self.db_domain.take());
OfficialCompiledDriverFinalizer {
inner: self.tasks.shutdown_official_tokio(cancel).await,
reaped_task_failed: self.reaped_task_failed,
}
}
}
impl OfficialCompiledDriverFinalizer {
pub fn bind_runtime(
self,
runtime: tokio::runtime::Runtime,
) -> OfficialCompiledRuntimeFinalizer {
OfficialCompiledRuntimeFinalizer {
inner: RuntimeDriverFinalizer::new(runtime, self.inner),
reaped_task_failed: self.reaped_task_failed,
}
}
}
#[cfg(test)]
pub(crate) fn test_startup_driver_finalizer(
ledger: ProcessLedger,
profile: ProcessAllocationProfile,
) -> OfficialCompiledDriverFinalizer {
OfficialCompiledDriverFinalizer {
inner: DriverFinalizer {
ledger: Some(ledger),
profile: Some(profile),
first_task_failure: None,
finished: false,
},
reaped_task_failed: false,
}
}
impl OfficialCompiledRuntimeFinalizer {
pub fn finish(self) -> Result<OfficialCompiledShutdownReport, AdmissionError> {
let (watermark, report) = self.inner.finish()?;
Ok(OfficialCompiledShutdownReport {
watermark,
ledger: report.ledger,
task_failed: self.reaped_task_failed || report.first_task_failure.is_some(),
})
}
}
fn reject_reason(error: &SubmitError) -> OfficialTcpRejectReason {
match error {
SubmitError::RegistryFull => OfficialTcpRejectReason::RuntimeCapacity,
SubmitError::Admission(_) => OfficialTcpRejectReason::Admission,
SubmitError::RouteAdapter => OfficialTcpRejectReason::Route,
SubmitError::RuntimeUnavailable
| SubmitError::ProcessUnhealthy(_)
| SubmitError::Lifecycle(_) => OfficialTcpRejectReason::RuntimeCapacity,
}
}
fn stop_reason(error: &SubmitError) -> OfficialTcpStopReason {
match error {
SubmitError::RuntimeUnavailable | SubmitError::ProcessUnhealthy(_) => {
OfficialTcpStopReason::Runtime
}
SubmitError::Lifecycle(_) => OfficialTcpStopReason::Lifecycle,
SubmitError::Admission(_) => OfficialTcpStopReason::Admission,
SubmitError::RouteAdapter => OfficialTcpStopReason::Route,
SubmitError::RegistryFull => OfficialTcpStopReason::Runtime,
}
}
struct TcpConnection {
socket: TcpStream,
}
struct TcpReadState {
payload: Option<ManagedBytes>,
expected_body_bytes: usize,
scratch: [u8; 1_024],
}
impl EntryReadPoll<TcpConnection> for TcpReadState {
fn poll_read(
&mut self,
connection: &mut TcpConnection,
_memory: &RequestMemory,
context: &mut Context<'_>,
) -> Poll<Result<ManagedBytes, AdmissionError>> {
loop {
let length = self.payload.as_ref().expect("payload exists").len();
if length >= self.expected_body_bytes {
return Poll::Ready(Ok(self.payload.take().expect("payload completes once")));
}
let remaining = self.expected_body_bytes - length;
let capacity = remaining.min(self.scratch.len());
let mut read = ReadBuf::new(&mut self.scratch[..capacity]);
match Pin::new(&mut connection.socket).poll_read(context, &mut read) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(_)) => return Poll::Ready(Err(AdmissionError::EntryReadFailed)),
Poll::Ready(Ok(())) if read.filled().is_empty() => {
return Poll::Ready(Err(AdmissionError::EntryReadFailed));
}
Poll::Ready(Ok(())) => {
let bytes = read.filled();
self.payload
.as_mut()
.expect("payload exists")
.try_extend_from_slice(bytes)
.unwrap_or_else(|_| std::process::abort());
}
}
}
}
}
struct TcpWriteState {
written: usize,
}
struct FramedTcpWriteState<W> {
inner: W,
}
impl<W> ResponseWritePoll<TcpConnection, CompiledResponseOutcome> for FramedTcpWriteState<W>
where
W: FixedResponseWrite,
{
fn poll_write(
&mut self,
connection: &mut TcpConnection,
output: &CompiledResponseOutcome,
context: &mut Context<'_>,
) -> Poll<Result<(), AdmissionError>> {
self.inner
.poll_write(
&mut connection.socket,
output.class(),
output.payload(),
context,
)
.map_err(|_| AdmissionError::ResponseWriteFailed)
}
}
impl ResponseWritePoll<TcpConnection, ManagedResponse> for TcpWriteState {
fn poll_write(
&mut self,
connection: &mut TcpConnection,
response: &ManagedResponse,
context: &mut Context<'_>,
) -> Poll<Result<(), AdmissionError>> {
while self.written < response.as_slice().len() {
match Pin::new(&mut connection.socket)
.poll_write(context, &response.as_slice()[self.written..])
{
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(_)) | Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(AdmissionError::ResponseWriteFailed));
}
Poll::Ready(Ok(written)) => self.written += written,
}
}
Poll::Ready(Ok(()))
}
}
struct DriverFinalizer {
ledger: Option<ProcessLedger>,
profile: Option<ProcessAllocationProfile>,
first_task_failure: Option<TaskFailure>,
finished: bool,
}
struct RuntimeDriverFinalizer {
runtime: Option<tokio::runtime::Runtime>,
finalizer: Option<DriverFinalizer>,
finished: bool,
}
struct DriverStoppedProof {
_private: (),
}
impl RuntimeDriverFinalizer {
fn new(runtime: tokio::runtime::Runtime, finalizer: DriverFinalizer) -> Self {
Self {
runtime: Some(runtime),
finalizer: Some(finalizer),
finished: false,
}
}
fn finish(mut self) -> Result<(ProcessAllocationSnapshot, ShutdownReport), AdmissionError> {
drop(
self.runtime
.take()
.expect("Runtime/driver owner is consumed exactly once"),
);
let proof = DriverStoppedProof { _private: () };
let result = self
.finalizer
.take()
.expect("driver finalizer is consumed exactly once")
.finish_after_driver(proof)?;
self.finished = true;
Ok(result)
}
}
impl Drop for RuntimeDriverFinalizer {
fn drop(&mut self) {
if !self.finished {
std::process::abort();
}
}
}
impl DriverFinalizer {
fn finish_after_driver(
mut self,
_proof: DriverStoppedProof,
) -> Result<(ProcessAllocationSnapshot, ShutdownReport), AdmissionError> {
let watermark = self
.profile
.take()
.expect("watermark owner exists")
.finish()?;
let ledger = self
.ledger
.take()
.expect("ledger owner exists")
.try_shutdown()?;
self.finished = true;
Ok((
watermark,
ShutdownReport {
ledger,
first_task_failure: self.first_task_failure.take(),
},
))
}
}
impl Drop for DriverFinalizer {
fn drop(&mut self) {
if !self.finished {
std::process::abort();
}
}
}
impl Drop for ManagedRequestTasks {
fn drop(&mut self) {
if self.shutdown_complete {
return;
}
std::process::abort();
}
}
fn task_result(result: Result<TaskOutput, JoinError>) -> Option<TaskFailure> {
match result {
Ok(Ok(_)) => None,
Ok(Err(TaskExecutionError::Admission(
AdmissionError::RequestCancelled | AdmissionError::RequestShutdown,
))) => Some(TaskFailure::Cancelled),
Ok(Err(TaskExecutionError::Admission(error))) => Some(TaskFailure::Admission(error)),
Ok(Err(TaskExecutionError::Deadline)) => Some(TaskFailure::Deadline),
Err(error) if error.is_cancelled() => Some(TaskFailure::Cancelled),
Err(error) if error.is_panic() => Some(TaskFailure::Panicked),
Err(_) => Some(TaskFailure::Panicked),
}
}
fn record_failure(first: &mut Option<TaskFailure>, failure: Option<TaskFailure>) {
if first.is_none() {
*first = failure;
}
}
fn request_requires_termination(ledger: &ProcessLedger, signal: RequestSignalHandle) -> bool {
match ledger.request_signal_status(signal) {
RequestSignalStatus::Requested => return true,
RequestSignalStatus::Running => {}
RequestSignalStatus::Stale | RequestSignalStatus::Unhealthy => std::process::abort(),
}
match ledger.request_db_phase(signal) {
DbRequestPhase::Finalizing => true,
DbRequestPhase::Query => false,
DbRequestPhase::Stale | DbRequestPhase::Unhealthy => std::process::abort(),
}
}
#[cfg(test)]
pub(crate) mod tests {
use std::os::unix::process::ExitStatusExt;
use std::{
env,
future::{Future, Ready, ready},
pin::Pin,
process::{Command, Stdio},
sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
task::{Context, Poll, Waker},
thread::ThreadId,
time::{Duration, Instant},
};
use saddle_admission::{
DbCreditProfile, DbRouteCreditDemand, EntryIoAllocation, OfficialTokioRegistrationProfile,
RequestEnvelope, ResourceConfig,
};
use super::*;
use crate::ApplicationPhase;
pub(crate) static OFFICIAL_TOKIO_PROFILE_TEST: Mutex<()> = Mutex::new(());
fn ledger(max_tasks: usize) -> ProcessLedger {
let process_state_reserve =
ProcessLedger::minimum_process_state_reserve(max_tasks).unwrap();
ProcessLedger::new(ResourceConfig {
managed_capacity: max_tasks * 64,
entry_reserve: 64,
framework_reserve: max_tasks * 64,
task_reserve: max_tasks * TASK_STORAGE_BOUND,
process_state_reserve,
system_estimate: 4_096,
safety_margin: 2_048,
process_limit: max_tasks * (64 + 64 + TASK_STORAGE_BOUND)
+ process_state_reserve
+ 6_208,
max_active_requests: max_tasks,
})
.unwrap()
}
fn ready_manager(max_tasks: usize) -> ManagedRequestTasks {
let requests = RequestLifecycle::new();
requests.mark_ready();
ManagedRequestTasks::new(requests, ledger(max_tasks), max_tasks)
}
fn entry_io_plan(events: &[(usize, usize)]) -> EntryIoAuditPlan {
let allocations = events
.iter()
.map(|&(size, align)| EntryIoAllocation::new(size, align).unwrap())
.collect::<Vec<_>>();
EntryIoAuditPlan::locked_linux_x86_64_tokio_1_53_1(256, 256, &allocations).unwrap()
}
fn official_ledger(max_tasks: usize) -> (ProcessLedger, OfficialTokioRegistrationProfile) {
official_ledger_with_transport(max_tasks, 1)
}
fn official_ledger_with_transport(
max_tasks: usize,
transport_connections: usize,
) -> (ProcessLedger, OfficialTokioRegistrationProfile) {
let registration_profile = OfficialTokioRegistrationProfile {
listener: 1,
transport_connections,
runtime_fixed: 2,
};
let process_state_reserve = ProcessLedger::minimum_process_state_reserve(max_tasks)
.unwrap()
+ ProcessLedger::official_tokio_state_reserve(registration_profile).unwrap();
let managed_capacity = max_tasks * 4_096;
let framework_reserve = max_tasks * 8_192;
let task_reserve = max_tasks * 4_096;
let system_estimate = 64 * 1_024;
let safety_margin = 64 * 1_024;
let process_limit = managed_capacity
+ 64
+ framework_reserve
+ task_reserve
+ process_state_reserve
+ system_estimate
+ safety_margin;
let ledger = ProcessLedger::new(ResourceConfig {
managed_capacity,
entry_reserve: 64,
framework_reserve,
task_reserve,
process_state_reserve,
system_estimate,
safety_margin,
process_limit,
max_active_requests: max_tasks,
})
.unwrap();
(ledger, registration_profile)
}
fn official_manager(
max_tasks: usize,
threshold: usize,
) -> (ManagedRequestTasks, OfficialTokioRegistrationProfile) {
let (ledger, registration_profile) = official_ledger(max_tasks);
let domain = ledger
.prepare_official_tokio_domain(registration_profile)
.unwrap();
let profile = ledger
.prepare_process_allocation_profile(threshold)
.unwrap();
let requests = RequestLifecycle::new();
requests.mark_ready();
(
ManagedRequestTasks::with_official_tokio(requests, ledger, max_tasks, domain, profile),
registration_profile,
)
}
fn official_manager_with_transport(
max_tasks: usize,
transport_connections: usize,
threshold: usize,
) -> ManagedRequestTasks {
let (ledger, registration_profile) =
official_ledger_with_transport(max_tasks, transport_connections);
let domain = ledger
.prepare_official_tokio_domain(registration_profile)
.unwrap();
let profile = ledger
.prepare_process_allocation_profile(threshold)
.unwrap();
let requests = RequestLifecycle::new();
requests.mark_ready();
ManagedRequestTasks::with_official_tokio(requests, ledger, max_tasks, domain, profile)
}
fn official_wait_manager() -> ManagedRequestTasks {
official_wait_manager_with_accounts(1)
}
fn official_wait_manager_with_accounts(accounts: usize) -> ManagedRequestTasks {
let registration_profile = OfficialTokioRegistrationProfile {
listener: 1,
transport_connections: 2,
runtime_fixed: 2,
};
let process_state_reserve =
ProcessLedger::minimum_process_state_reserve_with_waiters(accounts, 1).unwrap()
+ ProcessLedger::official_tokio_state_reserve(registration_profile).unwrap();
let managed_capacity = 8_192;
let framework_reserve = 16_384;
let task_reserve = 8_192;
let system_estimate = 64 * 1_024;
let safety_margin = 64 * 1_024;
let process_limit = managed_capacity
+ 64
+ framework_reserve
+ task_reserve
+ process_state_reserve
+ system_estimate
+ safety_margin;
let ledger = ProcessLedger::new_with_waiters(
ResourceConfig {
managed_capacity,
entry_reserve: 64,
framework_reserve,
task_reserve,
process_state_reserve,
system_estimate,
safety_margin,
process_limit,
max_active_requests: accounts,
},
1,
)
.unwrap();
let domain = ledger
.prepare_official_tokio_domain(registration_profile)
.unwrap();
let profile = ledger
.prepare_process_allocation_profile(usize::MAX)
.unwrap();
let requests = RequestLifecycle::new();
requests.mark_ready();
ManagedRequestTasks::with_official_tokio(requests, ledger, accounts + 1, domain, profile)
}
#[derive(Clone, Copy)]
struct FixtureRouteProof {
db: bool,
response_capacity: usize,
}
#[derive(Clone, Copy)]
enum FixtureContext {
Success,
SuccessEmpty,
AsyncError,
}
#[derive(Clone, Copy, Debug)]
struct FixtureExecutionError;
struct FixtureCompiledRoute;
#[derive(Clone, Copy)]
struct FixtureFramingPlan {
_private: (),
}
struct FixtureFraming;
struct FixtureFramedWrite {
header: [u8; 128],
header_len: usize,
header_written: usize,
body_written: usize,
payload_capacity: usize,
encoded: bool,
closed: bool,
}
static FRAMED_WRITE_SEGMENTS: AtomicUsize = AtomicUsize::new(0);
impl ResponseFramingAdapter for FixtureFraming {
type Plan = FixtureFramingPlan;
type Error = ();
type WriteState = FixtureFramedWrite;
fn begin(
&self,
_: Self::Plan,
payload_capacity: usize,
) -> Result<Self::WriteState, Self::Error> {
Ok(FixtureFramedWrite {
header: [0; 128],
header_len: 0,
header_written: 0,
body_written: 0,
payload_capacity,
encoded: false,
closed: false,
})
}
}
impl FixedResponseWrite for FixtureFramedWrite {
type Error = ();
fn poll_write(
&mut self,
socket: &mut TcpStream,
class: crate::compiled_route::ResponseOutcomeClass,
payload: &ManagedResponse,
context: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
if !self.encoded {
if payload.as_slice().len() > self.payload_capacity {
return Poll::Ready(Err(()));
}
self.header_len =
encode_fixture_header(&mut self.header, class, payload.as_slice().len())?;
self.encoded = true;
}
let bytes = if self.header_written < self.header_len {
&self.header[self.header_written..self.header_len]
} else {
&payload.as_slice()[self.body_written..]
};
if !bytes.is_empty() {
let end = bytes.len().min(3);
match Pin::new(&mut *socket).poll_write(context, &bytes[..end]) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(_)) | Poll::Ready(Ok(0)) => return Poll::Ready(Err(())),
Poll::Ready(Ok(written)) => {
FRAMED_WRITE_SEGMENTS.fetch_add(1, Ordering::SeqCst);
if self.header_written < self.header_len {
self.header_written += written;
} else {
self.body_written += written;
}
context.waker().wake_by_ref();
return Poll::Pending;
}
}
}
if !self.closed {
match Pin::new(socket).poll_shutdown(context) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(_)) => return Poll::Ready(Err(())),
Poll::Ready(Ok(())) => self.closed = true,
}
}
Poll::Ready(Ok(()))
}
}
fn encode_fixture_header(
output: &mut [u8; 128],
class: crate::compiled_route::ResponseOutcomeClass,
length: usize,
) -> Result<usize, ()> {
const PREFIX: &[u8] = b"HTTP/1.1 fixture\r\nX-Saddle-Class: ";
const LENGTH: &[u8] = b"\r\nContent-Length: ";
const SUFFIX: &[u8] = b"\r\nConnection: close\r\n\r\n";
output[..PREFIX.len()].copy_from_slice(PREFIX);
let class = match class {
crate::compiled_route::ResponseOutcomeClass::Success => b"success".as_slice(),
crate::compiled_route::ResponseOutcomeClass::InvalidRequest => {
b"invalid-request".as_slice()
}
crate::compiled_route::ResponseOutcomeClass::BusinessRejected => {
b"business-rejected".as_slice()
}
crate::compiled_route::ResponseOutcomeClass::Unavailable => b"unavailable".as_slice(),
crate::compiled_route::ResponseOutcomeClass::Internal => b"internal".as_slice(),
};
let mut cursor = PREFIX.len();
output[cursor..cursor + class.len()].copy_from_slice(class);
cursor += class.len();
output[cursor..cursor + LENGTH.len()].copy_from_slice(LENGTH);
cursor += LENGTH.len();
let mut digits = [0_u8; 20];
let mut value = length;
let mut count = 0;
loop {
digits[count] = b'0' + u8::try_from(value % 10).map_err(|_| ())?;
count += 1;
value /= 10;
if value == 0 {
break;
}
}
for digit in digits[..count].iter().rev() {
output[cursor] = *digit;
cursor += 1;
}
output[cursor..cursor + SUFFIX.len()].copy_from_slice(SUFFIX);
Ok(cursor + SUFFIX.len())
}
impl CompiledRouteAdapter for FixtureCompiledRoute {
type Proof = FixtureRouteProof;
type Context = FixtureContext;
type Error = FixtureExecutionError;
type Future = Ready<Result<ManagedResponse, FixtureExecutionError>>;
fn managed_commitment(&self, _: Self::Proof) -> Result<usize, Self::Error> {
Ok(4_096)
}
fn response_capacity(&self, proof: Self::Proof) -> Result<usize, Self::Error> {
Ok(proof.response_capacity)
}
fn db_resources<'a>(
&self,
proof: Self::Proof,
domain: Option<&'a DbPermitDomain>,
) -> Result<DbRouteResources<'a>, Self::Error> {
if !proof.db {
return Ok(DbRouteResources::none());
}
let domain = domain.ok_or(FixtureExecutionError)?;
let demand = DbRouteCreditDemand::new(1, 1).map_err(|_| FixtureExecutionError)?;
Ok(DbRouteResources::required(domain, demand))
}
fn execute(
&self,
proof: Self::Proof,
context: Self::Context,
body: ManagedBytes,
permit: Option<DbRequestPermit>,
memory: &RequestMemory,
) -> Result<Self::Future, Self::Error> {
if proof.db != permit.is_some() || body.as_slice() != b"x" {
return Err(FixtureExecutionError);
}
if let Some(permit) = permit {
permit
.begin_finalizing()
.and_then(|claim| claim.complete_after_connection_return())
.map_err(|_| FixtureExecutionError)?;
}
let mut response = memory
.try_response_builder(proof.response_capacity)
.map_err(|_| FixtureExecutionError)?;
match context {
FixtureContext::Success | FixtureContext::SuccessEmpty => {
response
.try_extend_from_slice(b"ok")
.map_err(|_| FixtureExecutionError)?;
Ok(ready(response.finish().map_err(|_| FixtureExecutionError)))
}
FixtureContext::AsyncError => Ok(ready(Err(FixtureExecutionError))),
}
}
}
impl ClassifiedCompiledRouteAdapter for FixtureCompiledRoute {
type Proof = FixtureRouteProof;
type Context = FixtureContext;
type Error = FixtureExecutionError;
type Future = Ready<CompiledResponseOutcome>;
fn managed_commitment(&self, _: Self::Proof) -> Result<usize, Self::Error> {
Ok(4_096)
}
fn response_capacity(&self, proof: Self::Proof) -> Result<usize, Self::Error> {
Ok(proof.response_capacity)
}
fn db_resources<'a>(
&self,
proof: Self::Proof,
domain: Option<&'a DbPermitDomain>,
) -> Result<DbRouteResources<'a>, Self::Error> {
if !proof.db {
return Ok(DbRouteResources::none());
}
let domain = domain.ok_or(FixtureExecutionError)?;
let demand = DbRouteCreditDemand::new(1, 1).map_err(|_| FixtureExecutionError)?;
Ok(DbRouteResources::required(domain, demand))
}
fn execute(
&self,
proof: Self::Proof,
context: Self::Context,
body: ManagedBytes,
permit: Option<DbRequestPermit>,
memory: &RequestMemory,
) -> Self::Future {
let valid = proof.db == permit.is_some() && body.as_slice() == b"x";
let mut response = memory
.try_response_builder(proof.response_capacity)
.expect("fixture proof reserves its response");
if !valid {
return ready(CompiledResponseOutcome::internal(
response.finish().expect("empty response is reserved"),
));
}
match context {
FixtureContext::Success => {
response
.try_extend_from_slice(b"ok")
.expect("fixture payload is inside proof");
ready(CompiledResponseOutcome::success(
response.finish().expect("fixture response remains valid"),
))
}
FixtureContext::SuccessEmpty => ready(CompiledResponseOutcome::success(
response.finish().expect("empty response is valid"),
)),
FixtureContext::AsyncError => ready(CompiledResponseOutcome::business_rejected(
response.finish().expect("empty error response is valid"),
)),
}
}
}
async fn socket_pair() -> (TcpStream, TcpStream) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let client = TcpStream::connect(address);
let (client, accepted) = tokio::join!(client, listener.accept());
(client.unwrap(), accepted.unwrap().0)
}
async fn write_after_yield(client: TcpStream, bytes: &'static [u8]) {
tokio::time::sleep(Duration::from_millis(10)).await;
loop {
client.writable().await.unwrap();
match client.try_write(bytes) {
Ok(written) if written == bytes.len() => break,
Ok(_) => std::process::abort(),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
}
async fn exchange_after_yield(
client: TcpStream,
request: &'static [u8],
response: &'static [u8],
observed: Arc<AtomicBool>,
) {
tokio::time::sleep(Duration::from_millis(10)).await;
loop {
client.writable().await.unwrap();
match client.try_write(request) {
Ok(written) if written == request.len() => break,
Ok(_) => std::process::abort(),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
let mut received = [0_u8; 128];
let mut length = 0;
while length < response.len() {
client.readable().await.unwrap();
match client.try_read(&mut received[length..response.len()]) {
Ok(0) => panic!("response socket closed before the fixed payload completed"),
Ok(read) => length += read,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
assert_eq!(&received[..length], response);
observed.store(true, Ordering::SeqCst);
}
async fn exchange_until_close(
client: TcpStream,
request: &'static [u8],
response: &'static [u8],
observed: Arc<AtomicBool>,
) {
tokio::time::sleep(Duration::from_millis(10)).await;
client.writable().await.unwrap();
assert_eq!(client.try_write(request).unwrap(), request.len());
let mut received = [0_u8; 256];
let mut length = 0;
loop {
client.readable().await.unwrap();
match client.try_read(&mut received[length..]) {
Ok(0) => break,
Ok(read) => length += read,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
assert_eq!(&received[..length], response);
observed.store(true, Ordering::SeqCst);
}
async fn saturate_write_buffer(socket: &TcpStream) {
let bytes = [0_u8; 16 * 1_024];
let mut written = 0;
loop {
socket.writable().await.unwrap();
match socket.try_write(&bytes) {
Ok(0) => std::process::abort(),
Ok(count) => written += count,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock && written != 0 => {
break;
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
assert!(written > 0);
}
async fn wait_and_reap(tasks: &mut ManagedRequestTasks) -> Option<TaskFailure> {
while tasks.tasks.iter().any(|task| !task.handle.is_finished()) {
tokio::task::yield_now().await;
}
tasks.reap_finished().await
}
#[test]
fn db_role_completes_only_after_same_task_finalizer() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let (finalizer, returned) = runtime.block_on(async {
let (mut tasks, _) = official_manager(1, usize::MAX);
let db_domain = tasks
.ledger
.as_ref()
.unwrap()
.prepare_db_domain(DbCreditProfile {
connections: 1,
operations: 1,
})
.unwrap();
let returned = Arc::new(AtomicBool::new(false));
let (client, server) = socket_pair().await;
let response_observed = Arc::new(AtomicBool::new(false));
let exchange = tokio::spawn(exchange_after_yield(
client,
b"x",
b"ok",
Arc::clone(&response_observed),
));
let returned_by_future = Arc::clone(&returned);
let outcome = tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::required(&db_domain, DbRouteCreditDemand::new(1, 1).unwrap()),
server,
1,
Duration::from_secs(1),
move |_, permit, memory| {
let response = memory.try_response(b"ok").unwrap();
let output = crate::db_finalizer::DbFinalizingOutput::begin(
response,
permit.expect("DB route owns its finalizer role"),
async move {
tokio::task::yield_now().await;
returned_by_future.store(true, Ordering::SeqCst);
},
)
.unwrap();
async move {
crate::db_finalizer::drive_db_finalizer(async move { Ok(output) })
.await
.unwrap()
}
},
);
assert!(matches!(outcome, TcpSubmitOutcome::Accepted));
assert!(wait_and_reap(&mut tasks).await.is_none());
exchange.await.unwrap();
assert!(response_observed.load(Ordering::SeqCst));
drop(db_domain);
(tasks.shutdown_official_tokio(false).await, returned)
});
assert!(returned.load(Ordering::SeqCst));
let owner = RuntimeDriverFinalizer::new(runtime, finalizer);
let (_, report) = owner.finish().unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn finalizing_resets_the_single_timer_beyond_attempt_deadline() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let (finalizer, returned) = runtime.block_on(async {
let (mut tasks, _) = official_manager(1, usize::MAX);
let db_domain = tasks
.ledger
.as_ref()
.unwrap()
.prepare_db_domain(DbCreditProfile {
connections: 1,
operations: 1,
})
.unwrap();
let returned = Arc::new(AtomicBool::new(false));
let (client, server) = socket_pair().await;
client.writable().await.unwrap();
assert_eq!(client.try_write(b"x").unwrap(), 1);
let returned_by_future = Arc::clone(&returned);
let started = Instant::now();
assert!(matches!(
tasks.try_submit_official_tcp_with_deadlines(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::required(
&db_domain,
DbRouteCreditDemand::new(1, 1).unwrap(),
),
server,
1,
Duration::from_millis(5),
Duration::from_millis(200),
move |_, permit, memory| {
let response = memory.try_response(&[]).unwrap();
let output = crate::db_finalizer::DbFinalizingOutput::begin(
response,
permit.expect("DB route owns its finalizer role"),
async move {
sleep(Duration::from_millis(30)).await;
returned_by_future.store(true, Ordering::SeqCst);
},
)
.unwrap();
async move {
crate::db_finalizer::drive_db_finalizer(async move { Ok(output) })
.await
.unwrap()
}
},
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
assert!(started.elapsed() >= Duration::from_millis(20));
drop(client);
drop(db_domain);
(tasks.shutdown_official_tokio(false).await, returned)
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert!(returned.load(Ordering::SeqCst));
assert_reconciled(&report.ledger, true);
}
#[test]
fn official_shutdown_signals_business_and_joins_without_abort() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let (finalizer, signal_observed) = runtime.block_on(async {
let (mut tasks, _) = official_manager(1, usize::MAX);
let business_started = Arc::new(AtomicBool::new(false));
let signal_observed = Arc::new(AtomicBool::new(false));
let (client, server) = socket_pair().await;
client.writable().await.unwrap();
assert_eq!(client.try_write(b"x").unwrap(), 1);
let started_by_factory = Arc::clone(&business_started);
let observed_by_business = Arc::clone(&signal_observed);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_secs(1),
move |_, permit, memory| {
assert!(permit.is_none());
started_by_factory.store(true, Ordering::SeqCst);
let shutdown = memory.shutdown_requested().unwrap();
let response = memory.try_response(&[]).unwrap();
async move {
let _ = shutdown.await;
observed_by_business.store(true, Ordering::SeqCst);
response
}
},
),
TcpSubmitOutcome::Accepted
));
while !business_started.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
if tasks.tasks[0].handle.is_finished() {
panic!("business finished before manager delivered shutdown");
}
}
let finalizer = tasks.shutdown_official_tokio(false).await;
drop(client);
(finalizer, signal_observed)
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert!(signal_observed.load(Ordering::SeqCst));
assert!(report.first_task_failure.is_none());
assert_reconciled(&report.ledger, true);
}
#[test]
fn official_tokio_tcp_entry_io_uses_one_task_and_two_stage_finalization() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let (finalizer, observed) = runtime.block_on(async {
let (mut tasks, _) = official_manager(8, usize::MAX);
let observed = Arc::new(AtomicUsize::new(0));
let (client, server) = socket_pair().await;
let response_observed = Arc::new(AtomicBool::new(false));
tokio::spawn(exchange_after_yield(
client,
b"body",
b"response",
Arc::clone(&response_observed),
));
let business_observed = Arc::clone(&observed);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
4,
Duration::from_secs(1),
move |body, permit, memory| {
assert!(permit.is_none());
assert_eq!(body.as_slice(), b"body");
business_observed.fetch_add(1, Ordering::SeqCst);
ready(memory.try_response(b"response").unwrap())
},
),
TcpSubmitOutcome::Accepted
));
let failure = wait_and_reap(&mut tasks).await;
assert!(
failure.is_none(),
"unexpected response-owner failure: {failure:?}"
);
assert_eq!(observed.load(Ordering::SeqCst), 1);
while !response_observed.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
assert!(response_observed.load(Ordering::SeqCst));
let (client, server) = socket_pair().await;
drop(client);
let disconnect_observed = Arc::clone(&observed);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
4,
Duration::from_secs(1),
move |_, permit, memory| {
assert!(permit.is_none());
disconnect_observed.fetch_add(1, Ordering::SeqCst);
ready(memory.try_response(b"must-not-run").unwrap())
},
),
TcpSubmitOutcome::Accepted
));
assert!(matches!(
wait_and_reap(&mut tasks).await,
Some(TaskFailure::Admission(AdmissionError::EntryReadFailed))
));
let (_client, server) = socket_pair().await;
assert!(matches!(
tasks.try_submit_official_tcp_with_deadlines(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_millis(1),
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(b"response").unwrap())
},
),
TcpSubmitOutcome::Accepted
));
assert!(matches!(
wait_and_reap(&mut tasks).await,
Some(TaskFailure::Deadline)
));
let (client, server) = socket_pair().await;
loop {
client.writable().await.unwrap();
match client.try_write(b"x") {
Ok(1) => break,
Ok(_) => std::process::abort(),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(_) => std::process::abort(),
}
}
saturate_write_buffer(&server).await;
assert!(matches!(
tasks.try_submit_official_tcp_with_deadlines(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_millis(1),
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(b"response").unwrap())
},
),
TcpSubmitOutcome::Accepted
));
assert!(matches!(
wait_and_reap(&mut tasks).await,
Some(TaskFailure::Deadline)
));
drop(client);
let (_held_client, held_server) = socket_pair().await;
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
held_server,
1,
Duration::from_secs(60),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(b"response").unwrap())
},
),
TcpSubmitOutcome::Accepted
));
let (_rejected_client, rejected_server) = socket_pair().await;
match tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
rejected_server,
1,
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(b"response").unwrap())
},
) {
TcpSubmitOutcome::Reject(
socket,
SubmitError::Admission(AdmissionError::WaitRegistryFull),
) => drop(socket),
_ => panic!("credit miss without a wait slot must reject and return the socket"),
};
let finalizer = tasks.shutdown_official_tokio(true).await;
(finalizer, observed)
});
let (watermark, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert!(!watermark.breached);
assert!(matches!(
report.first_task_failure,
Some(TaskFailure::Cancelled)
));
assert_reconciled(&report.ledger, true);
assert_eq!(observed.load(Ordering::SeqCst), 1);
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let (mut tasks, _) = official_manager(1, usize::MAX);
let (client, server) = socket_pair().await;
tokio::spawn(write_after_yield(client, b"!"));
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_secs(1),
|_, _, _| async { panic!("injected business panic") },
),
TcpSubmitOutcome::Accepted
));
assert!(matches!(
wait_and_reap(&mut tasks).await,
Some(TaskFailure::Panicked)
));
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, false);
assert!(matches!(
report.ledger.health_failure,
HealthFailure::RequestPanic | HealthFailure::RequestAllocationEscape
));
let (baseline_ledger, _) = official_ledger(1);
let baseline = match baseline_ledger.prepare_process_allocation_profile(1) {
Err(AdmissionError::ProcessAllocationThresholdExceeded {
live_requested_bytes,
..
}) => live_requested_bytes,
_ => panic!("one byte cannot cover the running test process"),
};
baseline_ledger.try_shutdown().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let (finalizer, allocation) = runtime.block_on(async {
let (mut tasks, _) = official_manager(1, baseline + 256 * 1_024);
let allocation = vec![0_u8; 1024 * 1_024];
assert!(
tasks
.process_allocation_profile
.as_ref()
.unwrap()
.snapshot()
.unwrap()
.breached
);
let (_client, server) = socket_pair().await;
match tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(b"response").unwrap())
},
) {
TcpSubmitOutcome::Stop(
Some(socket),
SubmitError::ProcessUnhealthy(
HealthFailure::ProcessAllocationWatermarkExceeded,
),
) => drop(socket),
_ => panic!("watermark breach must stop before phase publication"),
};
(tasks.shutdown_official_tokio(false).await, allocation)
});
drop(allocation);
let (watermark, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert!(watermark.breached);
assert_reconciled(&report.ledger, false);
assert_eq!(
report.ledger.health_failure,
HealthFailure::ProcessAllocationWatermarkExceeded
);
}
#[test]
fn compiled_db_route_permit_reaches_business_and_miss_rolls_back_before_publish() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_wait_manager_with_accounts(2);
let ledger = tasks
.ledger
.as_ref()
.expect("manager owns the route ledger");
let db = ledger
.prepare_db_domain(DbCreditProfile {
connections: 1,
operations: 1,
})
.unwrap();
let demand = DbRouteCreditDemand::new(1, 1).unwrap();
let business_saw_permit = Arc::new(AtomicBool::new(false));
let release_query = Arc::new(AtomicBool::new(false));
let (client, server) = socket_pair().await;
tokio::spawn(write_after_yield(client, b"x"));
let observed = Arc::clone(&business_saw_permit);
let release = Arc::clone(&release_query);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::required(&db, demand),
server,
1,
Duration::from_secs(60),
move |_, permit, memory| {
let permit = permit.expect("DB route receives its opaque permit");
assert_eq!(permit.demand(), demand);
observed.store(true, Ordering::SeqCst);
let response = memory.try_response(b"pending").unwrap();
async move {
std::future::poll_fn(|context| {
if release.load(Ordering::SeqCst) {
Poll::Ready(())
} else {
context.waker().wake_by_ref();
Poll::Pending
}
})
.await;
permit
.begin_finalizing()
.unwrap()
.complete_after_connection_return()
.unwrap();
response
}
},
),
TcpSubmitOutcome::Accepted
));
while !business_saw_permit.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
assert_eq!(db.snapshot().unwrap().connections_in_use, 1);
assert_eq!(tasks.tasks.len(), 1);
let rejected_factory_ran = Arc::new(AtomicBool::new(false));
let rejected_observed = Arc::clone(&rejected_factory_ran);
let (_client, rejected_server) = socket_pair().await;
let waiting = match tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::required(&db, demand),
rejected_server,
1,
Duration::from_secs(1),
move |_, _, memory| {
rejected_observed.store(true, Ordering::SeqCst);
ready(memory.try_response(b"must-not-run").unwrap())
},
) {
TcpSubmitOutcome::Registered(waiting) => waiting,
TcpSubmitOutcome::Reject(_, error) => {
panic!("DB miss rejected instead of registering: {error:?}")
}
TcpSubmitOutcome::Stop(_, error) => {
panic!("DB miss stopped instead of registering: {error:?}")
}
TcpSubmitOutcome::Accepted => {
panic!("DB miss unexpectedly published a second task")
}
};
assert!(!rejected_factory_ran.load(Ordering::SeqCst));
assert_eq!(tasks.tasks.len(), 1, "DB miss must not submit a task");
assert_eq!(
tasks.ledger.as_ref().unwrap().snapshot().active_accounts,
1,
"DB miss must roll back its request account"
);
assert_eq!(
tasks
.official_tokio_domain
.as_ref()
.unwrap()
.snapshot()
.unwrap()
.transport_in_use,
1,
"DB miss must roll back its transport credit"
);
let (waiting_socket, removal) = tasks.cancel_registered(waiting);
assert!(removal.is_removed());
drop(waiting_socket);
release_query.store(true, Ordering::SeqCst);
assert!(wait_and_reap(&mut tasks).await.is_none());
drop(db);
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert!(report.first_task_failure.is_none());
assert_reconciled(&report.ledger, true);
}
#[test]
fn fixed_http_framing_writes_header_body_and_close_in_one_task() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
FRAMED_WRITE_SEGMENTS.store(0, Ordering::SeqCst);
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_manager_with_transport(1, 1, usize::MAX);
let (client, server) = socket_pair().await;
let observed = Arc::new(AtomicBool::new(false));
let exchange = tokio::spawn(exchange_until_close(
client,
b"x",
b"HTTP/1.1 fixture\r\nX-Saddle-Class: success\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok",
Arc::clone(&observed),
));
assert!(matches!(
tasks.try_submit_compiled_http(
Arc::new(FixtureCompiledRoute),
FixtureRouteProof {
db: false,
response_capacity: 2,
},
FixtureContext::Success,
Arc::new(FixtureFraming),
FixtureFramingPlan { _private: () },
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
exchange.await.unwrap();
assert!(observed.load(Ordering::SeqCst));
tasks.shutdown_official_tokio(false).await
});
assert!(
FRAMED_WRITE_SEGMENTS.load(Ordering::SeqCst) > 2,
"fixed header/body state must preserve partial write progress"
);
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn empty_success_and_empty_error_keep_distinct_closed_outcomes() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_manager_with_transport(2, 2, usize::MAX);
for (context, expected) in [
(
FixtureContext::SuccessEmpty,
b"HTTP/1.1 fixture\r\nX-Saddle-Class: success\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.as_slice(),
),
(
FixtureContext::AsyncError,
b"HTTP/1.1 fixture\r\nX-Saddle-Class: business-rejected\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.as_slice(),
),
] {
let (client, server) = socket_pair().await;
let observed = Arc::new(AtomicBool::new(false));
let exchange = tokio::spawn(exchange_until_close(
client,
b"x",
expected,
Arc::clone(&observed),
));
assert!(matches!(
tasks.try_submit_compiled_http(
Arc::new(FixtureCompiledRoute),
FixtureRouteProof {
db: false,
response_capacity: 0,
},
context,
Arc::new(FixtureFraming),
FixtureFramingPlan { _private: () },
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
exchange.await.unwrap();
assert!(observed.load(Ordering::SeqCst));
}
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn compiled_route_adapter_drives_real_tcp_and_bounds_service_error() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_manager_with_transport(2, 2, usize::MAX);
let adapter = Arc::new(FixtureCompiledRoute);
let proof = FixtureRouteProof {
db: false,
response_capacity: 2,
};
let (client, server) = socket_pair().await;
let observed = Arc::new(AtomicBool::new(false));
tokio::spawn(exchange_after_yield(
client,
b"x",
b"ok",
Arc::clone(&observed),
));
assert!(matches!(
tasks.try_submit_compiled_tcp(
Arc::clone(&adapter),
proof,
FixtureContext::Success,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
while !observed.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
assert!(observed.load(Ordering::SeqCst));
let (error_client, error_server) = socket_pair().await;
error_client.try_write(b"x").unwrap();
assert!(matches!(
tasks.try_submit_compiled_tcp(
adapter,
proof,
FixtureContext::AsyncError,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
error_server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(
wait_and_reap(&mut tasks).await.is_none(),
"Service execution error maps to the bounded empty response"
);
let mut empty = [0_u8; 1];
loop {
error_client.readable().await.unwrap();
match error_client.try_read(&mut empty) {
Ok(0) => break,
Ok(_) => panic!("bounded error mapping must not write unapproved bytes"),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(error) => panic!("unexpected error response read: {error}"),
}
}
let db = tasks
.ledger
.as_ref()
.unwrap()
.prepare_db_domain(DbCreditProfile {
connections: 1,
operations: 1,
})
.unwrap();
let db_proof = FixtureRouteProof {
db: true,
response_capacity: 2,
};
let (db_client, db_server) = socket_pair().await;
let db_observed = Arc::new(AtomicBool::new(false));
tokio::spawn(exchange_after_yield(
db_client,
b"x",
b"ok",
Arc::clone(&db_observed),
));
assert!(matches!(
tasks.try_submit_compiled_tcp(
Arc::new(FixtureCompiledRoute),
db_proof,
FixtureContext::Success,
Some(&db),
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
db_server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
while !db_observed.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
assert_eq!(db.snapshot().unwrap().connections_in_use, 0);
drop(db);
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn registered_tcp_rolls_back_claim_and_has_one_retry_generation() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_wait_manager();
let adapter = Arc::new(FixtureCompiledRoute);
let proof = FixtureRouteProof {
db: false,
response_capacity: 2,
};
let (held_client, held_server) = socket_pair().await;
tokio::spawn(write_after_yield(held_client, b"x"));
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
held_server,
1,
Duration::from_secs(60),
|_, permit, memory| {
assert!(permit.is_none());
let response = memory.try_response(b"held").unwrap();
async move {
std::future::pending::<()>().await;
response
}
},
),
TcpSubmitOutcome::Accepted
));
let (retry_client, retry_server) = socket_pair().await;
tokio::spawn(write_after_yield(retry_client, b"x"));
let waiting = match tasks.try_submit_compiled_tcp(
Arc::clone(&adapter),
proof,
FixtureContext::Success,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
retry_server,
1,
Duration::from_secs(1),
) {
TcpSubmitOutcome::Registered(waiting) => waiting,
_ => panic!("temporary account exhaustion must register"),
};
assert_eq!(tasks.tasks.len(), 1);
assert_eq!(
tasks.ledger.as_ref().unwrap().snapshot().active_accounts,
1,
"Registered must roll back its unpublished request account"
);
assert_eq!(tasks.ledger.as_ref().unwrap().wait_snapshot().registered, 1);
let (retry_server, removal) = tasks.cancel_registered(waiting);
assert!(removal.is_removed());
assert_eq!(tasks.ledger.as_ref().unwrap().wait_snapshot().registered, 0);
let waiting = match tasks.try_submit_compiled_tcp(
Arc::clone(&adapter),
proof,
FixtureContext::Success,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
retry_server,
1,
Duration::from_secs(1),
) {
TcpSubmitOutcome::Registered(waiting) => waiting,
_ => panic!("cancelled input can register one new generation"),
};
tasks.tasks[0].handle.abort();
assert!(matches!(
wait_and_reap(&mut tasks).await,
Some(TaskFailure::Cancelled)
));
let hint = tasks
.claim_retry_hint()
.expect("capacity release removes exactly one generation");
let retry_server = match waiting.into_retry(hint) {
Ok(socket) => socket,
Err(_) => panic!("hint matches the returned generation"),
};
assert!(tasks.claim_retry_hint().is_none());
assert!(matches!(
tasks.try_submit_compiled_tcp(
adapter,
proof,
FixtureContext::Success,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
retry_server,
1,
Duration::from_secs(1),
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
tasks.requests.begin_draining();
let (draining_client, draining_server) = socket_pair().await;
match tasks.try_submit_compiled_tcp(
Arc::new(FixtureCompiledRoute),
proof,
FixtureContext::Success,
None,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
draining_server,
1,
Duration::from_secs(1),
) {
TcpSubmitOutcome::Stop(Some(socket), SubmitError::Lifecycle(_)) => drop(socket),
_ => panic!("Draining must win before Admission attempt and keep socket owned"),
}
drop(draining_client);
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn task_slot_stays_owned_from_envelope_completion_through_join_reap() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let mut tasks = official_wait_manager();
tasks.capacity = 1;
let (client, server) = socket_pair().await;
client.writable().await.unwrap();
assert_eq!(client.try_write(b"x").unwrap(), 1);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(&[]).unwrap())
},
),
TcpSubmitOutcome::Accepted
));
while tasks.ledger.as_ref().unwrap().snapshot().active_accounts != 0 {
tokio::task::yield_now().await;
}
assert_eq!(tasks.tasks.len(), tasks.capacity);
let (second_client, second_server) = socket_pair().await;
second_client.writable().await.unwrap();
assert_eq!(second_client.try_write(b"y").unwrap(), 1);
let waiting = match tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
second_server,
1,
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(&[]).unwrap())
},
) {
TcpSubmitOutcome::Registered(waiting) => waiting,
_ => panic!("unreaped JoinHandle must keep the task slot unavailable"),
};
while !tasks.tasks[0].handle.is_finished() {
tokio::task::yield_now().await;
}
assert!(tasks.reap_finished().await.is_none());
let hint = tasks
.claim_retry_hint()
.expect("reaping the task releases exactly one waiter hint");
let retry_server = waiting
.into_retry(hint)
.unwrap_or_else(|_| panic!("hint must match the waiting generation"));
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
retry_server,
1,
Duration::from_secs(1),
|_, permit, memory| {
assert!(permit.is_none());
ready(memory.try_response(&[]).unwrap())
},
),
TcpSubmitOutcome::Accepted
));
assert!(wait_and_reap(&mut tasks).await.is_none());
drop(second_client);
drop(client);
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
}
#[test]
fn driver_finalization_destroys_runtime_before_profile_and_ledger_finish() {
let _profile_test = OFFICIAL_TOKIO_PROFILE_TEST.lock().unwrap();
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let stale_handle = runtime.handle().clone();
let finalizer = runtime.block_on(async {
let (tasks, _) = official_manager(1, usize::MAX);
tasks.shutdown_official_tokio(false).await
});
let (_, report) = RuntimeDriverFinalizer::new(runtime, finalizer)
.finish()
.unwrap();
assert_reconciled(&report.ledger, true);
let post_finalize_io = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
stale_handle.block_on(async {
tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
});
}));
assert!(
post_finalize_io.is_err(),
"final ledger verification must imply the unique I/O driver is gone"
);
}
#[test]
fn abnormal_runtime_driver_finalizer_drop_is_finite_fail_closed_in_subprocess() {
const CHILD_TEST: &str = "admission::tests::abnormal_runtime_driver_finalizer_drop_child";
const DEADLINE: Duration = Duration::from_secs(5);
for mode in ["direct", "during-unwind"] {
let mut child = Command::new(env::current_exe().unwrap())
.args(["--exact", CHILD_TEST, "--nocapture"])
.env("SADDLE_RUNTIME_DRIVER_FINALIZER_ABORT_CHILD", mode)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
let started = Instant::now();
let status = loop {
if let Some(status) = child.try_wait().unwrap() {
break status;
}
if started.elapsed() >= DEADLINE {
child.kill().unwrap();
child.wait().unwrap();
panic!("driver-finalizer Drop ({mode}) exceeded {DEADLINE:?}");
}
std::thread::sleep(Duration::from_millis(10));
};
assert_eq!(
status.signal(),
Some(6),
"driver-finalizer Drop ({mode}) must terminate with SIGABRT, got {status}"
);
}
}
#[test]
fn abnormal_runtime_driver_finalizer_drop_child() {
let Some(mode) = env::var_os("SADDLE_RUNTIME_DRIVER_FINALIZER_ABORT_CHILD") else {
return;
};
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
let finalizer = runtime.block_on(async {
let (tasks, _) = official_manager(1, usize::MAX);
tasks.shutdown_official_tokio(false).await
});
let owner = RuntimeDriverFinalizer::new(runtime, finalizer);
if mode == "during-unwind" {
let _owner = owner;
panic!("existing unwind must not make finalizer loss catchable");
}
drop(owner);
unreachable!("abnormal driver-finalizer Drop must abort");
}
#[test]
fn reviewed_task_bound_covers_locked_layout_and_registry() {
assert_eq!(std::mem::size_of::<RequestEnvelope<Ready<()>>>(), 80);
assert_eq!(std::mem::align_of::<RequestEnvelope<Ready<()>>>(), 8);
assert_eq!(std::mem::size_of::<TaskOutput>(), 56);
assert_eq!(std::mem::align_of::<TaskOutput>(), 8);
assert_eq!(std::mem::size_of::<TaskEntry>(), 80);
assert!(TASK_STORAGE_BOUND >= 256 + std::mem::size_of::<TaskEntry>());
}
#[test]
fn unavailable_runtime_rejects_before_claim_or_factory() {
let calls = AtomicUsize::new(0);
let mut tasks = ready_manager(1);
let error = tasks
.try_submit(0, |_| {
calls.fetch_add(1, Ordering::SeqCst);
ready(())
})
.unwrap_err();
assert!(matches!(error, SubmitError::RuntimeUnavailable));
assert_eq!(calls.load(Ordering::SeqCst), 0);
assert_eq!(tasks.requests.phase(), ApplicationPhase::Ready);
assert_eq!(tasks.ledger.as_ref().unwrap().snapshot().active_accounts, 0);
let runtime = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let report = runtime.block_on(tasks.shutdown(false)).unwrap();
assert_reconciled(&report.ledger, true);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn ready_abort_and_panic_each_reconcile_once() {
let mut ready_tasks = ready_manager(1);
ready_tasks
.try_submit(0, |_| ready(()))
.expect("ready envelope submits");
let report = ready_tasks.shutdown(false).await.unwrap();
assert!(report.first_task_failure.is_none());
assert_reconciled(&report.ledger, true);
let mut cancelled_tasks = ready_manager(1);
cancelled_tasks
.try_submit(0, |_| std::future::pending::<()>())
.expect("pending envelope submits");
let report = cancelled_tasks.shutdown(true).await.unwrap();
assert!(matches!(
report.first_task_failure,
Some(TaskFailure::Cancelled)
));
assert_reconciled(&report.ledger, true);
let mut panic_tasks = ready_manager(1);
panic_tasks
.try_submit(0, |_| async { panic!("injected request panic") })
.expect("panic envelope submits");
let report = panic_tasks.shutdown(false).await.unwrap();
assert!(matches!(
report.first_task_failure,
Some(TaskFailure::Panicked)
));
assert_reconciled(&report.ledger, false);
assert!(matches!(
report.ledger.health_failure,
HealthFailure::RequestPanic | HealthFailure::RequestAllocationEscape
));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn registry_reaps_before_reuse_and_rejects_without_factory_when_full() {
let calls = Arc::new(AtomicUsize::new(0));
let mut tasks = ready_manager(1);
tasks
.try_submit(0, |_| std::future::pending::<()>())
.unwrap();
let rejected_calls = Arc::clone(&calls);
assert!(matches!(
tasks.try_submit(0, move |_| {
rejected_calls.fetch_add(1, Ordering::SeqCst);
ready(())
}),
Err(SubmitError::RegistryFull)
));
assert_eq!(calls.load(Ordering::SeqCst), 0);
tasks.tasks[0].handle.abort();
while !tasks.tasks[0].handle.is_finished() {
tokio::task::yield_now().await;
}
assert!(matches!(
tasks.reap_finished().await,
Some(TaskFailure::Cancelled)
));
assert!(tasks.tasks.is_empty());
let accepted_calls = Arc::clone(&calls);
tasks
.try_submit(0, move |_| {
accepted_calls.fetch_add(1, Ordering::SeqCst);
ready(())
})
.unwrap();
let report = tasks.shutdown(false).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_reconciled(&report.ledger, true);
}
#[test]
fn abnormal_registry_drop_is_finite_fail_closed_in_subprocess() {
const CHILD_TEST: &str = "admission::tests::abnormal_registry_drop_fail_closed_child";
const DEADLINE: Duration = Duration::from_secs(5);
for mode in [
"live-direct",
"live-during-unwind",
"empty-after-reap-direct",
"empty-after-reap-during-unwind",
] {
let mut child = Command::new(env::current_exe().unwrap())
.args(["--exact", CHILD_TEST, "--nocapture"])
.env("SADDLE_RUNTIME_ABORT_REGISTRY_CHILD", mode)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
let started = Instant::now();
let status = loop {
if let Some(status) = child.try_wait().unwrap() {
break status;
}
if started.elapsed() >= DEADLINE {
child.kill().unwrap();
child.wait().unwrap();
panic!(
"live registry subprocess ({mode}) did not terminate within {DEADLINE:?}"
);
}
std::thread::sleep(Duration::from_millis(10));
};
assert_eq!(
status.signal(),
Some(6),
"live registry ({mode}) must terminate with SIGABRT, got {status}"
);
}
}
#[test]
fn uncooperative_official_shutdown_is_finite_fail_closed_in_subprocess() {
const CHILD_TEST: &str =
"admission::tests::uncooperative_official_shutdown_fail_closed_child";
const DEADLINE: Duration = Duration::from_secs(5);
let mut child = Command::new(env::current_exe().unwrap())
.args(["--exact", CHILD_TEST, "--nocapture"])
.env("SADDLE_RUNTIME_UNCOOPERATIVE_SHUTDOWN_CHILD", "1")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
let started = Instant::now();
let status = loop {
if let Some(status) = child.try_wait().unwrap() {
break status;
}
if started.elapsed() >= DEADLINE {
child.kill().unwrap();
child.wait().unwrap();
panic!("uncooperative shutdown did not terminate within {DEADLINE:?}");
}
std::thread::sleep(Duration::from_millis(10));
};
assert_eq!(
status.signal(),
Some(6),
"uncooperative shutdown must terminate with SIGABRT, got {status}"
);
}
#[test]
fn uncooperative_official_shutdown_fail_closed_child() {
if env::var_os("SADDLE_RUNTIME_UNCOOPERATIVE_SHUTDOWN_CHILD").is_none() {
return;
}
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_io()
.enable_time()
.build()
.unwrap();
runtime.block_on(async {
let (mut tasks, _) = official_manager(1, usize::MAX);
let business_started = Arc::new(AtomicBool::new(false));
let (client, server) = socket_pair().await;
client.writable().await.unwrap();
assert_eq!(client.try_write(b"x").unwrap(), 1);
let started_by_factory = Arc::clone(&business_started);
assert!(matches!(
tasks.try_submit_official_tcp(
4_096,
4_096,
entry_io_plan(&[]),
entry_io_plan(&[]),
DbRouteResources::none(),
server,
1,
Duration::from_millis(10),
move |_, permit, _| {
assert!(permit.is_none());
started_by_factory.store(true, Ordering::SeqCst);
std::future::pending::<ManagedResponse>()
},
),
TcpSubmitOutcome::Accepted
));
while !business_started.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
let _ = tasks.shutdown_official_tokio(false).await;
drop(client);
});
unreachable!("uncooperative shutdown must abort the process");
}
#[test]
fn abnormal_registry_drop_fail_closed_child() {
let Some(mode) = env::var_os("SADDLE_RUNTIME_ABORT_REGISTRY_CHILD") else {
return;
};
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.build()
.unwrap();
runtime.block_on(async {
let mut tasks = ready_manager(1);
if mode.to_string_lossy().starts_with("empty-after-reap") {
tasks.try_submit(0, |_| ready(())).unwrap();
while !tasks.tasks[0].handle.is_finished() {
tokio::task::yield_now().await;
}
assert!(tasks.reap_finished().await.is_none());
assert!(tasks.tasks.is_empty());
} else {
tasks
.try_submit(0, |_| std::future::pending::<()>())
.unwrap();
}
if mode.to_string_lossy().ends_with("during-unwind") {
panic!("force registry destruction during an existing unwind");
}
drop(tasks);
});
unreachable!("live registry drop must abort the process");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn phase_claim_rolls_back_or_publishes_across_shutdown() {
let requests = RequestLifecycle::new();
requests.mark_ready();
let claim = requests.try_claim().unwrap();
requests.begin_draining();
drop(claim);
requests.wait_until_drained().await;
assert_eq!(requests.phase(), ApplicationPhase::Draining);
let requests = RequestLifecycle::new();
requests.mark_ready();
let claim = requests.try_claim().unwrap();
requests.begin_draining();
let guard = claim.publish();
assert_eq!(requests.phase(), ApplicationPhase::Draining);
drop(guard);
requests.wait_until_drained().await;
}
struct WakeState {
waker: Mutex<Option<Waker>>,
polls: AtomicUsize,
ready: AtomicBool,
first_thread: Mutex<Option<ThreadId>>,
migrated: AtomicBool,
}
struct WokenFuture {
state: Arc<WakeState>,
}
impl Future for WokenFuture {
type Output = ();
fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<()> {
self.state.polls.fetch_add(1, Ordering::AcqRel);
let current = std::thread::current().id();
let mut first = self.state.first_thread.lock().unwrap();
if let Some(first) = *first {
if first != current {
self.state.migrated.store(true, Ordering::Release);
}
} else {
*first = Some(current);
}
drop(first);
if self.state.ready.load(Ordering::Acquire) {
Poll::Ready(())
} else {
*self.state.waker.lock().unwrap() = Some(context.waker().clone());
Poll::Pending
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn external_wake_reinstalls_identity_on_every_poll() {
let state = Arc::new(WakeState {
waker: Mutex::new(None),
polls: AtomicUsize::new(0),
ready: AtomicBool::new(false),
first_thread: Mutex::new(None),
migrated: AtomicBool::new(false),
});
let mut tasks = ready_manager(1);
let task_state = Arc::clone(&state);
tasks
.try_submit(0, move |_| WokenFuture { state: task_state })
.unwrap();
for expected_poll in 1..=10_000 {
while state.waker.lock().unwrap().is_none() {
tokio::task::yield_now().await;
}
state.waker.lock().unwrap().take().unwrap().wake();
while state.polls.load(Ordering::Acquire) <= expected_poll {
tokio::task::yield_now().await;
}
if state.migrated.load(Ordering::Acquire) {
break;
}
}
assert!(
state.migrated.load(Ordering::Acquire),
"Tokio task did not migrate during controlled external wakes"
);
while state.waker.lock().unwrap().is_none() {
tokio::task::yield_now().await;
}
state.ready.store(true, Ordering::Release);
state.waker.lock().unwrap().take().unwrap().wake();
let report = tasks.shutdown(false).await.unwrap();
assert_reconciled(&report.ledger, true);
assert!(state.polls.load(Ordering::Acquire) >= 3);
}
fn assert_reconciled(snapshot: &ProcessSnapshot, healthy: bool) {
assert_eq!(snapshot.active_accounts, 0);
assert_eq!(snapshot.committed, 0);
assert_eq!(snapshot.charged, 0);
assert_eq!(snapshot.framework_charged, 0);
assert_eq!(snapshot.task_charged, 0);
assert_eq!(snapshot.healthy, healthy);
}
}