use std::{
sync::Arc,
time::{Duration, Instant},
};
use num_bigint::BigUint;
use tokio::sync::{broadcast, Notify};
use tracing::{debug, error, info, warn};
use tycho_simulation::tycho_core::Bytes;
use crate::{
algorithm::Algorithm,
derived::{
computation::ComputationRequirements, events::DerivedDataEvent, tracker::ReadinessTracker,
SharedDerivedDataRef,
},
feed::{
events::{MarketEvent, MarketEventHandler},
exclusivity::{remove_exclusive_components, scope_event},
market_data::MarketData,
},
graph::{EdgeWeightUpdaterWithDerived, GraphManager},
types::internal::SolveTask,
worker_pool_router::LiquidityScope,
BlockInfo, Order, OrderQuote, QuoteStatus, SingleOrderQuote, SolveError, SolveParams,
};
fn record_task_pickup_metrics(pool_name: &str, queue_wait: Duration, queue_depth: usize) {
metrics::histogram!("worker_pool_queue_wait_seconds", "pool" => pool_name.to_string())
.record(queue_wait.as_secs_f64());
metrics::gauge!("worker_pool_queue_depth", "pool" => pool_name.to_string())
.set(queue_depth as f64);
}
fn record_solve_duration(pool_name: &str, solve_time: Duration) {
metrics::histogram!("worker_pool_solve_duration_seconds", "pool" => pool_name.to_string())
.record(solve_time.as_secs_f64());
}
pub(crate) struct SolverWorker<A>
where
A: Algorithm,
A::GraphManager: MarketEventHandler,
{
algorithm: A,
graph_manager: A::GraphManager,
market_data: MarketData,
derived_data: SharedDerivedDataRef,
requirements: ComputationRequirements,
readiness_tracker: ReadinessTracker,
ready_notify: Arc<Notify>,
initialized: bool,
worker_id: usize,
pool_name: String,
liquidity_scope: LiquidityScope,
}
impl<A> SolverWorker<A>
where
A: Algorithm,
A::GraphManager: MarketEventHandler,
{
pub fn new(
market_data: MarketData,
derived_data: SharedDerivedDataRef,
algorithm: A,
worker_id: usize,
pool_name: String,
) -> Self {
let requirements = algorithm.computation_requirements();
Self {
algorithm,
graph_manager: A::GraphManager::default(),
market_data,
derived_data,
requirements: requirements.clone(),
readiness_tracker: ReadinessTracker::new(requirements),
ready_notify: Arc::new(Notify::new()),
initialized: false,
worker_id,
pool_name,
liquidity_scope: LiquidityScope::default(),
}
}
pub(crate) fn with_liquidity_scope(mut self, scope: LiquidityScope) -> Self {
self.liquidity_scope = scope;
self
}
pub async fn initialize_graph(&mut self) {
let topology = {
let market = self.market_data.read().await;
let topology = market.component_topology().clone(); match self.liquidity_scope {
LiquidityScope::PublicOnly => {
remove_exclusive_components(market.base_market_state(), topology)
}
LiquidityScope::IncludeExclusive => topology,
}
};
self.graph_manager
.initialize_graph(&topology);
self.initialized = true;
}
pub async fn process_event(&mut self, event: MarketEvent) {
let event = {
let market = self.market_data.read().await;
match self.liquidity_scope {
LiquidityScope::PublicOnly => scope_event(market.base_market_state(), event),
LiquidityScope::IncludeExclusive => event,
}
};
match event {
MarketEvent::MarketUpdated { .. } => {
if let Err(e) = self
.graph_manager
.handle_event(&event)
.await
{
warn!("Error handling market event: {:?}", e);
}
}
}
}
pub async fn quote(
&mut self,
order: &Order,
params: SolveParams,
) -> Result<SingleOrderQuote, SolveError> {
let start_time = Instant::now();
debug!(
order_id = %order.id(),
token_in = ?order.token_in(),
token_out = ?order.token_out(),
amount = %order.amount(),
side = ?order.side(),
"processing order"
);
if self
.readiness_tracker
.has_requirements() &&
!self.readiness_tracker.is_ready()
{
return Err(SolveError::NotReady(format!(
"derived data not ready: missing {:?}",
self.readiness_tracker.missing()
)));
}
if !self.initialized {
self.initialize_graph().await;
}
let graph = self.graph_manager.graph();
let (block_info, solved_against) = {
let view = match params.state_label() {
Some(l) => self
.market_data
.read_labeled(l)
.await
.map_err(|e| SolveError::NotReady(e.to_string()))?,
None => self.market_data.read().await,
};
let last_block = view
.last_updated()
.ok_or(SolveError::NotReady("No block info".to_string()))?;
let block_info = BlockInfo::new(
last_block.number(),
last_block.hash().to_string(),
last_block.timestamp(),
);
let solved_against = view
.state_label()
.cloned()
.unwrap_or_else(|| last_block.number().to_string());
(block_info, solved_against)
};
let result = self
.algorithm
.find_best_route(
graph,
self.market_data.clone(),
params.state_label().cloned(),
Some(self.derived_data.clone()),
order,
)
.await;
let order_quote = match result {
Ok(result) => {
let amount_out_net_gas = result
.net_amount_out()
.to_biguint()
.unwrap_or(BigUint::ZERO);
let gas_price = result.gas_price().clone();
let algo_price_impact = result.price_impact();
let route = result.into_route();
if let Err(err) = route.validate() {
error!(
order_id = %order.id(),
algorithm = self.algorithm.name(),
error = %err,
"algorithm produced an invalid route"
);
return Err(SolveError::AlgorithmError(format!(
"{} produced an invalid route: {err}",
self.algorithm.name()
)));
}
let gas_estimate = route.total_gas();
let amount_in = if order.is_sell() {
order.amount().clone()
} else {
route
.swaps()
.first()
.map(|s| s.amount_in().clone())
.ok_or_else(|| {
error!(
order_id = %order.id(),
"route missing first swap for buy order"
);
SolveError::no_route_found(order.id())
})?
};
let amount_out = if order.is_sell() {
let output_token = route.output_token().ok_or_else(|| {
error!(
order_id = %order.id(),
"route missing swaps for sell order"
);
SolveError::no_route_found(order.id())
})?;
route
.swaps()
.iter()
.filter(|s| *s.token_out() == output_token)
.map(|s| s.amount_out().clone())
.fold(BigUint::ZERO, |acc, x| acc + x)
} else {
order.amount().clone()
};
let price_impact_bps = algo_price_impact
.or_else(|| {
super::price_impact::spot_price_impact(
&route,
&amount_in,
&amount_out,
&self.market_data,
)
})
.map(|f| (f * 10_000.0).round() as i32);
let mut quote = OrderQuote::new(
order.id().to_string(),
QuoteStatus::Success,
amount_in,
amount_out,
gas_estimate,
amount_out_net_gas,
block_info.clone(),
self.algorithm.name().to_string(),
Bytes::from(order.sender().as_ref()),
Bytes::from(order.effective_receiver().as_ref()),
solved_against,
)
.with_route(route)
.with_gas_price(gas_price);
if let Some(bps) = price_impact_bps {
quote = quote.with_price_impact_bps(bps);
}
quote
}
Err(err) => {
return Err(solve_error_from_algorithm_error(order.id(), order.amount(), err))
}
};
let solve_time = start_time.elapsed();
record_solve_duration(&self.pool_name, solve_time);
Ok(SingleOrderQuote::new(order_quote, solve_time.as_millis() as u64))
}
async fn wait_until_ready(&self, timeout: Duration) -> Result<(), SolveError> {
if !self
.readiness_tracker
.has_requirements() ||
self.readiness_tracker.is_ready()
{
return Ok(());
}
let deadline = Instant::now() + timeout;
loop {
let notified = self.ready_notify.notified();
if self.readiness_tracker.is_ready() {
return Ok(());
}
if self
.readiness_tracker
.is_blocked_for_current_block()
{
return Err(SolveError::ComputationFailed(format!(
"required computation failed for current block: {:?}",
self.readiness_tracker.missing()
)));
}
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(SolveError::NotReady(format!(
"timeout waiting for derived data: missing {:?}",
self.readiness_tracker.missing()
)));
}
tokio::select! {
_ = tokio::time::sleep(remaining) => {
return Err(SolveError::NotReady(format!(
"timeout waiting for derived data: missing {:?}",
self.readiness_tracker.missing()
)));
}
_ = notified => {
if self.readiness_tracker.is_blocked_for_current_block() {
return Err(SolveError::ComputationFailed(format!(
"required computation failed for current block: {:?}",
self.readiness_tracker.missing()
)));
}
continue;
}
}
}
}
pub async fn run(
&mut self,
mut event_rx: broadcast::Receiver<MarketEvent>,
mut derived_event_rx: broadcast::Receiver<DerivedDataEvent>,
task_rx: async_channel::Receiver<SolveTask>,
mut shutdown_rx: broadcast::Receiver<()>,
) where
A::GraphManager: EdgeWeightUpdaterWithDerived,
{
info!(self.worker_id, "worker started");
let mut derived_closed = false;
loop {
tokio::select! {
biased;
_ = shutdown_rx.recv() => {
info!(self.worker_id, "worker shutting down");
break;
}
event_result = event_rx.recv() => {
match event_result {
Ok(event) => {
self.process_event(event).await;
}
Err(broadcast::error::RecvError::Closed) => {
info!(self.worker_id, "event receiver closed, shutting down");
break;
}
Err(broadcast::error::RecvError::Lagged(skipped)) => {
warn!(
self.worker_id,
skipped = skipped,
"event receiver lagged, skipped {} events. Reinitializing graph from current market state",
skipped
);
self.initialize_graph().await;
}
}
}
derived_result = derived_event_rx.recv(), if !derived_closed => {
match derived_result {
Ok(event) => {
self.readiness_tracker.handle_event(&event);
self.ready_notify.notify_waiters();
if let DerivedDataEvent::ComputationComplete { computation_id, block, .. } = &event {
if self.requirements.is_required(computation_id) {
let market = self.market_data.read().await;
let derived = self.derived_data.read().await;
let updated = self.graph_manager.update_edge_weights_with_derived(market, &derived);
debug!(
self.worker_id,
computation_id,
block,
updated,
"updated edge weights with derived data"
);
}
}
}
Err(broadcast::error::RecvError::Closed) => {
warn!(self.worker_id, "derived event receiver closed; continuing with last derived data");
derived_closed = true;
}
Err(broadcast::error::RecvError::Lagged(skipped)) => {
warn!(
self.worker_id,
skipped,
"derived event receiver lagged, skipped {} events",
skipped
);
let market = self.market_data.read().await;
let derived = self.derived_data.read().await;
let updated = self.graph_manager.update_edge_weights_with_derived(market, &derived);
debug!(
self.worker_id,
updated,
"recovered edge weights after lag"
);
}
}
}
task = task_rx.recv() => {
match task.ok() {
Some(task) => {
let task_id = task.id();
record_task_pickup_metrics(
&self.pool_name,
task.wait_time(),
task_rx.len(),
);
if let Err(e) = self.wait_until_ready(self.algorithm.timeout()).await {
warn!(
self.worker_id,
task_id = %task_id,
error = %e,
"not ready to solve"
);
task.respond(Err(e));
continue;
}
let result = {
let params = task.params().clone();
let order = task.order();
self.quote(order, params).await
};
task.respond(result);
}
None => {
info!(self.worker_id, "task channel closed, exiting");
break;
}
}
}
}
}
}
}
fn solve_error_from_algorithm_error(
order_id: &str,
amount_in: &BigUint,
err: crate::AlgorithmError,
) -> SolveError {
match err {
crate::AlgorithmError::NoPath { reason, .. } => {
debug!(order_id = %order_id, error = %err, "no route found");
SolveError::no_route_found_with_reason(order_id, reason)
}
crate::AlgorithmError::Timeout { elapsed_ms } => {
warn!(order_id = %order_id, elapsed_ms, "solve timeout");
SolveError::Timeout { elapsed_ms }
}
crate::AlgorithmError::InsufficientLiquidity => {
debug!(order_id = %order_id, "insufficient liquidity on all paths");
SolveError::insufficient_liquidity(amount_in.clone(), BigUint::ZERO)
}
crate::AlgorithmError::DataNotFound { kind, id } => {
warn!(order_id = %order_id, kind, id = ?id, "required data not found");
SolveError::MissingData(match id {
Some(id) => format!("{kind}: {id}"),
None => kind.to_string(),
})
}
crate::AlgorithmError::SimulationFailed { component_id, error } => {
warn!(order_id = %order_id, %component_id, %error, "simulation failed");
SolveError::SimulationFailed(format!("{component_id}: {error}"))
}
crate::AlgorithmError::InvalidConfiguration { .. } |
crate::AlgorithmError::ExactOutNotSupported |
crate::AlgorithmError::Other(_) => {
error!(order_id = %order_id, error = %err, "algorithm error");
SolveError::AlgorithmError(err.to_string())
}
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use rustc_hash::FxHashMap;
use super::*;
use crate::{
algorithm::{
most_liquid::DepthAndPrice,
test_utils::{component, order, setup_market_weighted, token, MockProtocolSim},
},
derived::{
computation::DerivedComputation,
computations::{SpotPriceComputation, TokenGasPriceComputation},
DerivedData,
},
graph::petgraph::{PetgraphStableDiGraphManager, StableDiGraph},
types::{OrderSide, Route, RouteResult, Swap},
AlgorithmError,
};
struct MockAlgorithm {
requirements: ComputationRequirements,
timeout: Duration,
}
impl MockAlgorithm {
fn new() -> Self {
Self { requirements: ComputationRequirements::none(), timeout: Duration::from_secs(1) }
}
fn with_requirements(mut self, requirements: ComputationRequirements) -> Self {
self.requirements = requirements;
self
}
}
impl Algorithm for MockAlgorithm {
type GraphType = StableDiGraph<DepthAndPrice>;
type GraphManager = PetgraphStableDiGraphManager<DepthAndPrice>;
fn name(&self) -> &str {
"mock"
}
async fn find_best_route(
&self,
_graph: &Self::GraphType,
_market: MarketData,
_label: Option<crate::feed::market_data::StateLabel>,
_derived: Option<SharedDerivedDataRef>,
_order: &Order,
) -> Result<crate::types::RouteResult, crate::AlgorithmError> {
Err(crate::AlgorithmError::Other("not implemented".to_string()))
}
fn computation_requirements(&self) -> ComputationRequirements {
self.requirements.clone()
}
fn timeout(&self) -> Duration {
self.timeout
}
}
struct InvalidRouteAlgorithm;
impl Algorithm for InvalidRouteAlgorithm {
type GraphType = StableDiGraph<DepthAndPrice>;
type GraphManager = PetgraphStableDiGraphManager<DepthAndPrice>;
fn name(&self) -> &str {
"invalid_route_mock"
}
async fn find_best_route(
&self,
_graph: &Self::GraphType,
_market: MarketData,
_label: Option<crate::feed::market_data::StateLabel>,
_derived: Option<SharedDerivedDataRef>,
_order: &Order,
) -> Result<RouteResult, AlgorithmError> {
let token_a = token(0x01, "A");
let token_b = token(0x02, "B");
let token_c = token(0x03, "C");
let token_d = token(0x04, "D");
let swap_ab = Swap::new(
"p1".to_string(),
"mock".to_string(),
token_a.address.clone(),
token_b.address.clone(),
BigUint::from(100u64),
BigUint::from(90u64),
BigUint::from(1u64),
component("p1", &[token_a.clone(), token_b.clone()]),
Box::new(MockProtocolSim::new(2.0)),
);
let swap_cd = Swap::new(
"p2".to_string(),
"mock".to_string(),
token_c.address.clone(),
token_d.address.clone(),
BigUint::from(90u64),
BigUint::from(80u64),
BigUint::from(1u64),
component("p2", &[token_c.clone(), token_d.clone()]),
Box::new(MockProtocolSim::new(2.0)),
);
let route =
Route::new(vec![swap_ab, swap_cd], FxHashMap::default()).expect("non-empty route");
Ok(RouteResult::new(route, num_bigint::BigInt::from(0), BigUint::from(1u64)))
}
fn computation_requirements(&self) -> ComputationRequirements {
ComputationRequirements::none()
}
fn timeout(&self) -> Duration {
Duration::from_secs(1)
}
}
#[tokio::test]
async fn test_quote_rejects_invalid_route() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let mut worker =
SolverWorker::new(market, derived, InvalidRouteAlgorithm, 0, "test_pool".to_string());
let token_a = token(0x01, "A");
let token_b = token(0x02, "B");
let ord = order(&token_a, &token_b, 100, OrderSide::Sell);
let result = worker
.quote(&ord, SolveParams::default())
.await;
match result {
Err(SolveError::AlgorithmError(msg)) => {
assert!(msg.contains("invalid route"), "unexpected message: {msg}");
}
other => panic!("expected AlgorithmError for invalid route, got {other:?}"),
}
}
#[tokio::test]
async fn wait_until_ready_returns_immediately_when_no_requirements() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let algorithm = MockAlgorithm::new();
let worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let result = worker
.wait_until_ready(Duration::from_millis(10))
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn wait_until_ready_returns_immediately_when_already_ready() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.allow_stale(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::ComputationComplete {
computation_id: SpotPriceComputation::ID,
block: 1,
failed_items: vec![],
});
let result = worker
.wait_until_ready(Duration::from_millis(10))
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn wait_until_ready_times_out_when_not_ready() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let result = worker
.wait_until_ready(Duration::from_millis(50))
.await;
assert!(result.is_err());
match result {
Err(SolveError::NotReady(msg)) => {
assert!(msg.contains("timeout"));
assert!(msg.contains("spot_prices"));
}
other => panic!("Expected NotReady error, got {:?}", other),
}
}
#[tokio::test]
async fn wait_until_ready_wakes_up_on_notify() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let notify = worker.ready_notify.clone();
let handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
notify.notify_waiters();
});
let result = worker
.wait_until_ready(Duration::from_millis(100))
.await;
handle.await.unwrap();
assert!(result.is_err());
}
#[tokio::test]
async fn wait_until_ready_succeeds_when_notified_and_ready() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let notify = worker.ready_notify.clone();
let handle = tokio::spawn({
async move {
tokio::time::sleep(Duration::from_millis(20)).await;
notify.notify_waiters();
}
});
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::ComputationComplete {
computation_id: SpotPriceComputation::ID,
block: 1,
failed_items: vec![],
});
let result = worker
.wait_until_ready(Duration::from_millis(100))
.await;
handle.abort(); assert!(result.is_ok());
}
#[tokio::test]
async fn notify_pattern_handles_multiple_waiters() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.allow_stale(TokenGasPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let notify = worker.ready_notify.clone();
let notify1 = notify.clone();
let waiter1 = tokio::spawn(async move {
notify1.notified().await;
true
});
let notify2 = notify.clone();
let waiter2 = tokio::spawn(async move {
notify2.notified().await;
true
});
tokio::time::sleep(Duration::from_millis(10)).await;
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::ComputationComplete {
computation_id: TokenGasPriceComputation::ID,
block: 1,
failed_items: vec![],
});
notify.notify_waiters();
let (r1, r2) = tokio::join!(waiter1, waiter2);
assert!(r1.unwrap());
assert!(r2.unwrap());
}
#[tokio::test]
async fn wait_until_ready_returns_immediately_on_blocked_state() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::NewBlock { block: 1 });
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::ComputationFailed {
computation_id: SpotPriceComputation::ID,
block: 1,
});
let notify = worker.ready_notify.clone();
let notifier = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
notify.notify_waiters();
});
let result = worker
.wait_until_ready(Duration::from_secs(5))
.await;
notifier.await.unwrap();
match result {
Err(SolveError::ComputationFailed(msg)) => {
assert!(
msg.contains("required computation failed"),
"expected 'required computation failed' message, got: {msg}"
);
}
other => panic!("Expected ComputationFailed error, got {:?}", other),
}
}
#[tokio::test]
async fn wait_until_ready_returns_blocked_when_failure_already_processed() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::NewBlock { block: 1 });
worker
.readiness_tracker
.handle_event(&DerivedDataEvent::ComputationFailed {
computation_id: SpotPriceComputation::ID,
block: 1,
});
let result = worker
.wait_until_ready(Duration::from_secs(1))
.await;
match result {
Err(SolveError::ComputationFailed(msg)) => {
assert!(
msg.contains("required computation failed"),
"expected 'required computation failed' message, got: {msg}"
);
}
other => panic!("Expected ComputationFailed error, got {:?}", other),
}
}
#[tokio::test]
async fn worker_updates_tracker_and_notifies_on_derived_event() {
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let requirements = ComputationRequirements::none()
.require_fresh(SpotPriceComputation::ID)
.unwrap();
let algorithm = MockAlgorithm::new().with_requirements(requirements);
let mut worker = SolverWorker::new(market, derived, algorithm, 0, "test_pool".to_string());
let (_event_tx, event_rx) = broadcast::channel::<MarketEvent>(16);
let (derived_tx, derived_rx) = broadcast::channel::<DerivedDataEvent>(16);
let (_task_tx, task_rx) = async_channel::bounded::<crate::types::internal::SolveTask>(16);
let (shutdown_tx, shutdown_rx) = broadcast::channel::<()>(1);
let handle = tokio::spawn(async move {
worker
.run(event_rx, derived_rx, task_rx, shutdown_rx)
.await;
});
derived_tx
.send(DerivedDataEvent::ComputationComplete {
computation_id: SpotPriceComputation::ID,
block: 1,
failed_items: vec![],
})
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = shutdown_tx.send(());
tokio::time::timeout(Duration::from_secs(1), handle)
.await
.expect("worker should shutdown")
.expect("worker task should not panic");
}
#[derive(Clone, Default)]
struct SharedLogBuffer(std::sync::Arc<std::sync::Mutex<Vec<u8>>>);
impl std::io::Write for SharedLogBuffer {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0
.lock()
.unwrap()
.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for SharedLogBuffer {
type Writer = SharedLogBuffer;
fn make_writer(&'a self) -> Self::Writer {
self.clone()
}
}
#[tokio::test]
async fn worker_handles_derived_channel_close_once_without_spinning() {
let logs = SharedLogBuffer::default();
let subscriber = tracing_subscriber::fmt()
.with_writer(logs.clone())
.with_max_level(tracing::Level::WARN)
.finish();
let _guard = tracing::subscriber::set_default(subscriber);
let (market, _) = setup_market_weighted(vec![]);
let derived = DerivedData::new_shared();
let mut worker =
SolverWorker::new(market, derived, MockAlgorithm::new(), 0, "test_pool".to_string());
let (_event_tx, event_rx) = broadcast::channel::<MarketEvent>(16);
let (derived_tx, derived_rx) = broadcast::channel::<DerivedDataEvent>(16);
let (_task_tx, task_rx) = async_channel::bounded::<crate::types::internal::SolveTask>(16);
let (shutdown_tx, shutdown_rx) = broadcast::channel::<()>(1);
let handle = tokio::spawn(async move {
worker
.run(event_rx, derived_rx, task_rx, shutdown_rx)
.await;
});
drop(derived_tx);
tokio::time::sleep(Duration::from_millis(100)).await;
shutdown_tx.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(1), handle)
.await
.expect("worker should shutdown")
.expect("worker task should not panic");
let output = String::from_utf8(logs.0.lock().unwrap().clone()).unwrap();
let closed_warns = output
.matches("derived event receiver closed")
.count();
assert_eq!(
closed_warns, 1,
"closed channel must be handled once, not spun on ({closed_warns} warns)"
);
}
#[test]
fn no_route_found_with_reason_carries_reason() {
use crate::algorithm::NoPathReason;
let err = SolveError::no_route_found_with_reason(
"order-1",
NoPathReason::DestinationTokenNotInGraph,
);
match err {
SolveError::NoRouteFound { order_id, reason } => {
assert_eq!(order_id, "order-1");
assert_eq!(reason, Some(NoPathReason::DestinationTokenNotInGraph));
}
other => panic!("expected NoRouteFound, got {other:?}"),
}
}
#[test]
fn no_route_found_defaults_to_no_reason() {
match SolveError::no_route_found("order-1") {
SolveError::NoRouteFound { reason, .. } => assert_eq!(reason, None),
other => panic!("expected NoRouteFound, got {other:?}"),
}
}
#[test]
fn task_pickup_metrics_recorded() {
use metrics_util::debugging::{DebugValue, DebuggingRecorder};
let recorder = DebuggingRecorder::new();
let snapshotter = recorder.snapshotter();
metrics::with_local_recorder(&recorder, || {
record_task_pickup_metrics("test_pool", std::time::Duration::from_millis(25), 3);
});
let mut wait_seen = false;
let mut depth_seen = false;
for (key, _unit, _description, value) in snapshotter.snapshot().into_vec() {
let key = key.key();
let pool_label = key
.labels()
.find(|label| label.key() == "pool")
.map(|label| label.value().to_string());
match key.name() {
"worker_pool_queue_wait_seconds" => {
assert_eq!(pool_label.as_deref(), Some("test_pool"));
let DebugValue::Histogram(samples) = value else {
panic!("expected histogram, got {value:?}");
};
assert_eq!(samples.len(), 1);
assert!((samples[0].into_inner() - 0.025).abs() < 1e-9);
wait_seen = true;
}
"worker_pool_queue_depth" => {
assert_eq!(pool_label.as_deref(), Some("test_pool"));
let DebugValue::Gauge(depth) = value else {
panic!("expected gauge, got {value:?}");
};
assert!((depth.into_inner() - 3.0).abs() < f64::EPSILON);
depth_seen = true;
}
_ => {}
}
}
assert!(wait_seen, "queue wait histogram not recorded");
assert!(depth_seen, "queue depth gauge not recorded");
}
#[test]
fn solve_duration_metric_recorded() {
use metrics_util::debugging::{DebugValue, DebuggingRecorder};
let recorder = DebuggingRecorder::new();
let snapshotter = recorder.snapshotter();
metrics::with_local_recorder(&recorder, || {
record_solve_duration("test_pool", std::time::Duration::from_millis(120));
});
let mut solve_seen = false;
for (key, _unit, _description, value) in snapshotter.snapshot().into_vec() {
let key = key.key();
if key.name() != "worker_pool_solve_duration_seconds" {
continue;
}
let pool_label = key
.labels()
.find(|label| label.key() == "pool")
.map(|label| label.value().to_string());
assert_eq!(pool_label.as_deref(), Some("test_pool"));
let DebugValue::Histogram(samples) = value else {
panic!("expected histogram, got {value:?}");
};
assert_eq!(samples.len(), 1);
assert!((samples[0].into_inner() - 0.120).abs() < 1e-9);
solve_seen = true;
}
assert!(solve_seen, "solve duration histogram not recorded");
}
#[test]
fn test_algorithm_error_maps_data_not_found_to_missing_data() {
let err = crate::AlgorithmError::DataNotFound { kind: "gas price", id: None };
let mapped = solve_error_from_algorithm_error("o1", &num_bigint::BigUint::from(5u64), err);
assert!(matches!(mapped, SolveError::MissingData(_)), "got {mapped:?}");
}
#[test]
fn test_algorithm_error_maps_simulation_failed() {
let err = crate::AlgorithmError::SimulationFailed {
component_id: "pool-1".to_string(),
error: "revert".to_string(),
};
let mapped = solve_error_from_algorithm_error("o1", &num_bigint::BigUint::from(5u64), err);
assert!(matches!(mapped, SolveError::SimulationFailed(_)), "got {mapped:?}");
}
#[test]
fn test_algorithm_error_maps_insufficient_liquidity() {
let err = crate::AlgorithmError::InsufficientLiquidity;
let mapped = solve_error_from_algorithm_error("o1", &num_bigint::BigUint::from(5u64), err);
assert!(matches!(mapped, SolveError::InsufficientLiquidity { .. }), "got {mapped:?}");
}
#[test]
fn test_algorithm_error_other_stays_algorithm_error() {
let err = crate::AlgorithmError::Other("boom".to_string());
let mapped = solve_error_from_algorithm_error("o1", &num_bigint::BigUint::from(5u64), err);
assert!(matches!(mapped, SolveError::AlgorithmError(_)), "got {mapped:?}");
}
#[test]
fn test_algorithm_error_timeout_stays_timeout() {
let err = crate::AlgorithmError::Timeout { elapsed_ms: 7 };
let mapped = solve_error_from_algorithm_error("o1", &num_bigint::BigUint::from(5u64), err);
assert!(matches!(mapped, SolveError::Timeout { elapsed_ms: 7 }), "got {mapped:?}");
}
}