use crate::{
idl::arcium::{
accounts::{
ArxNode,
ClockAccount,
Cluster,
LargeExecPool,
LargeMempool,
MXEAccount,
MediumExecPool,
MediumMempool,
SmallExecPool,
SmallMempool,
TinyExecPool,
TinyMempool,
},
types::{ClusterMembership, ComputationReference, LeaderSelector, MempoolSize},
},
pda::{arx_acc, clock_acc, cluster_acc, mempool_acc},
};
use anchor_client::anchor_lang::{AccountDeserialize, Discriminator};
use anchor_lang::prelude::Pubkey;
use bytemuck::Zeroable;
use solana_account_decoder_client_types::{UiAccountData, UiAccountEncoding, UiDataSliceConfig};
use solana_rpc_client::nonblocking::rpc_client::RpcClient as AsyncRpcClient;
use solana_rpc_client_api::{
client_error::Error as SolanaClientError,
config::{RpcAccountInfoConfig, RpcProgramAccountsConfig},
filter::{Memcmp, MemcmpEncodedBytes, RpcFilterType},
};
use std::{collections::HashSet, hash::Hash};
use thiserror::Error;
pub const MIN_CLUSTER_CONTEXT_SLOT: u64 = 0;
#[derive(Error, Debug, Clone, PartialEq)]
pub enum ClusterOffsetError {
#[error("Failed to fetch node accounts from RPC: {0}")]
AccountFetchFailed(String),
#[error("Failed to deserialize node account data: {0}")]
DeserializationFailed(String),
#[error("Found inactive node in cluster at offset {0}")]
InactiveNode(u32),
#[error("Cluster has no nodes")]
EmptyCluster,
#[error("No cluster membership found for target cluster")]
NoClusterMembership,
}
#[derive(Debug, Clone, PartialEq)]
pub enum ClusterOffsetState {
Available(u32),
NotLookedUp,
Unavailable(ClusterOffsetError),
}
impl ClusterOffsetState {
pub fn is_available(&self) -> bool {
matches!(self, ClusterOffsetState::Available(_))
}
pub fn get(&self) -> Option<u32> {
match self {
ClusterOffsetState::Available(offset) => Some(*offset),
_ => None,
}
}
pub fn error(&self) -> Option<&ClusterOffsetError> {
match self {
ClusterOffsetState::Unavailable(err) => Some(err),
_ => None,
}
}
}
impl std::fmt::Display for ClusterOffsetState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ClusterOffsetState::Available(offset) => write!(f, "Available (offset: {})", offset),
ClusterOffsetState::NotLookedUp => write!(f, "Not looked up"),
ClusterOffsetState::Unavailable(err) => write!(f, "Unavailable: {}", err),
}
}
}
pub async fn arx_acc_active(
rpc_client: &AsyncRpcClient,
node_offset: u32,
) -> Result<bool, Box<dyn std::error::Error>> {
let arx_acc = arx_acc(node_offset);
let bytes = rpc_client
.get_account(&arx_acc)
.await
.map_err(|e| format!("Failed to get account data: {}", e))?
.data;
let arx_data = ArxNode::try_deserialize(&mut bytes.as_slice())
.map_err(|e| format!("Failed to deserialize account data: {}", e))?;
Ok(arx_data.is_active)
}
pub async fn active_proposals(
rpc_client: &AsyncRpcClient,
cluster_offset: u32,
) -> Result<Vec<u64>, Box<dyn std::error::Error>> {
let cluster_acc = cluster_acc(cluster_offset);
let bytes = rpc_client
.get_account(&cluster_acc)
.await
.map_err(|e| format!("Failed to get account data: {}", e))?
.data;
let cluster_data = Cluster::try_deserialize(&mut bytes.as_slice())
.map_err(|e| format!("Failed to deserialize account data: {}", e))?;
Ok(dedupe(cluster_data.cu_price_proposals.to_vec()))
}
fn dedupe<T: PartialEq + Eq + Hash + Copy>(arr: Vec<T>) -> Vec<T> {
let mut seen = HashSet::new();
let mut result = Vec::new();
for &item in arr.iter() {
if seen.insert(item) {
result.push(item);
}
}
result
}
pub async fn get_all_cluster_accounts(
rpc_client: &AsyncRpcClient,
min_context_slot: Option<u64>,
) -> Result<Vec<(Pubkey, Cluster)>, Box<dyn std::error::Error>> {
let program_id = crate::idl::arcium::ID;
let discriminator = Cluster::DISCRIMINATOR;
let memcmp_filter = RpcFilterType::Memcmp(Memcmp::new(
0,
MemcmpEncodedBytes::Bytes(discriminator.to_vec()),
));
let config = RpcProgramAccountsConfig {
filters: Some(vec![memcmp_filter]),
account_config: RpcAccountInfoConfig {
encoding: Some(UiAccountEncoding::Base64),
commitment: None,
data_slice: None,
min_context_slot,
},
with_context: None,
sort_results: None,
};
let ui_accounts = rpc_client
.get_program_ui_accounts_with_config(&program_id, config)
.await?;
let mut clusters = Vec::new();
for (pubkey, ui_account) in ui_accounts {
let data = ui_account.data.decode().ok_or_else(|| {
let variant = match &ui_account.data {
UiAccountData::Json(_) => "JsonParsed".to_string(),
UiAccountData::LegacyBinary(_) => "LegacyBinary".to_string(),
UiAccountData::Binary(_, enc) => format!("Binary({enc:?})"),
};
format!("Failed to decode account data for {pubkey}: got {variant}, expected Base64")
})?;
match Cluster::try_deserialize(&mut data.as_slice()) {
Ok(cluster) => clusters.push((pubkey, cluster)),
Err(_) => continue, }
}
Ok(clusters)
}
async fn get_mxe_count(
rpc_client: &AsyncRpcClient,
min_context_slot: Option<u64>,
cluster_offset: u32,
) -> Result<usize, Box<dyn std::error::Error>> {
let program_id = crate::idl::arcium::ID;
let discriminator = MXEAccount::DISCRIMINATOR;
let mut comparative_bytes = Vec::with_capacity(13);
comparative_bytes.extend_from_slice(discriminator);
comparative_bytes.push(1);
comparative_bytes.extend_from_slice(&cluster_offset.to_le_bytes());
let memcmp_filter =
RpcFilterType::Memcmp(Memcmp::new(0, MemcmpEncodedBytes::Bytes(comparative_bytes)));
let data_slice_config = UiDataSliceConfig {
offset: 0,
length: 0,
};
let config = RpcProgramAccountsConfig {
filters: Some(vec![memcmp_filter]),
account_config: RpcAccountInfoConfig {
encoding: Some(UiAccountEncoding::Base64),
commitment: None,
data_slice: Some(data_slice_config),
min_context_slot,
},
with_context: None,
sort_results: None,
};
let ui_accounts = rpc_client
.get_program_ui_accounts_with_config(&program_id, config)
.await?;
Ok(ui_accounts.len())
}
#[derive(Debug, Clone)]
pub struct ClusterInfo {
pub pubkey: Pubkey,
pub offset: ClusterOffsetState,
pub cluster: Cluster,
pub mxe_count: usize,
}
impl ClusterInfo {
pub fn node_count(&self) -> usize {
self.cluster.nodes.len()
}
pub fn max_nodes(&self) -> u16 {
self.cluster.cluster_size
}
pub fn pending_node_count(&self) -> usize {
self.cluster.pending_nodes.len()
}
pub fn node_utilization_percent(&self) -> f32 {
let max = self.max_nodes();
if max > 0 {
(self.node_count() as f32 / max as f32) * 100.0
} else {
0.0
}
}
}
pub async fn get_cluster_discovery_info(
rpc_client: &AsyncRpcClient,
_current_epoch: u64,
min_context_slot: Option<u64>,
) -> Result<Vec<ClusterInfo>, Box<dyn std::error::Error>> {
let clusters = get_all_cluster_accounts(rpc_client, min_context_slot).await?;
let mut infos = Vec::with_capacity(clusters.len());
for (pubkey, cluster) in clusters {
let offset = find_cluster_offset_via_node(rpc_client, &cluster).await;
let mxe_count = if let ClusterOffsetState::Available(offset) = offset {
let mxe_count = get_mxe_count(rpc_client, min_context_slot, offset).await?;
mxe_count
} else {
0
};
infos.push(ClusterInfo {
pubkey,
offset,
cluster,
mxe_count,
});
}
Ok(infos)
}
async fn find_cluster_offset_via_node(
rpc_client: &AsyncRpcClient,
cluster: &Cluster,
) -> ClusterOffsetState {
if cluster.nodes.is_empty() {
return ClusterOffsetState::Unavailable(ClusterOffsetError::EmptyCluster);
}
let node_pubkeys: Vec<Pubkey> = cluster
.nodes
.iter()
.map(|node_ref| arx_acc(node_ref.offset))
.collect();
let accounts = match rpc_client.get_multiple_accounts(&node_pubkeys).await {
Ok(accounts) => accounts,
Err(e) => {
return ClusterOffsetState::Unavailable(ClusterOffsetError::AccountFetchFailed(
e.to_string(),
))
}
};
debug_assert_eq!(
accounts.len(),
cluster.nodes.len(),
"RPC returned {} accounts but requested {} nodes",
accounts.len(),
cluster.nodes.len()
);
let mut found_offset = None;
for (i, maybe_account) in accounts.iter().enumerate() {
let node_offset = cluster.nodes[i].offset;
let account = match maybe_account.as_ref() {
Some(acc) => acc,
None => {
return ClusterOffsetState::Unavailable(ClusterOffsetError::AccountFetchFailed(
format!("Node account not found: {}", node_offset),
))
}
};
let node = match ArxNode::try_deserialize(&mut account.data.as_slice()) {
Ok(node) => node,
Err(e) => {
return ClusterOffsetState::Unavailable(ClusterOffsetError::DeserializationFailed(
e.to_string(),
))
}
};
if !node.is_active {
return ClusterOffsetState::Unavailable(ClusterOffsetError::InactiveNode(node_offset));
}
if found_offset.is_none() {
if let ClusterMembership::Active(cluster_offset) = &node.cluster_membership {
found_offset = Some(*cluster_offset);
}
}
}
match found_offset {
Some(offset) => ClusterOffsetState::Available(offset),
None => ClusterOffsetState::Unavailable(ClusterOffsetError::NoClusterMembership),
}
}
pub async fn get_current_epoch(
rpc_client: &AsyncRpcClient,
) -> Result<u64, Box<dyn std::error::Error>> {
let clock_pubkey = clock_acc();
let account = rpc_client.get_account(&clock_pubkey).await?;
let clock_data = ClockAccount::try_deserialize(&mut account.data.as_slice())?;
Ok(clock_data.current_epoch.0)
}
pub async fn get_mempool_acc_data(
rpc: &AsyncRpcClient,
mempool_acc: &Pubkey,
) -> Result<MempoolWrapper, ComputationPoolError> {
let mempool_data = rpc
.get_account_data(mempool_acc)
.await
.map_err(ComputationPoolError::new_solana_error)?;
MempoolWrapper::from_raw(&mempool_data)
}
pub async fn get_mempool_acc_data_raw(
rpc: &AsyncRpcClient,
mempool_acc: &Pubkey,
) -> Result<Vec<u8>, SolanaClientError> {
let mempool_data = rpc.get_account_data(mempool_acc).await?;
Ok(mempool_data)
}
pub fn mempool_tier_from_discriminator(disc: &[u8]) -> Option<MempoolSize> {
match disc {
TinyMempool::DISCRIMINATOR => Some(MempoolSize::Tiny),
SmallMempool::DISCRIMINATOR => Some(MempoolSize::Small),
MediumMempool::DISCRIMINATOR => Some(MempoolSize::Medium),
LargeMempool::DISCRIMINATOR => Some(MempoolSize::Large),
_ => None,
}
}
pub async fn get_mempool_tiers(
rpc: &AsyncRpcClient,
cluster_offsets: &[u32],
min_context_slot: Option<u64>,
) -> Result<Vec<Option<MempoolSize>>, SolanaClientError> {
const MAX_MULTIPLE_ACCOUNTS: usize = 100;
let pdas: Vec<Pubkey> = cluster_offsets
.iter()
.map(|offset| mempool_acc(*offset))
.collect();
let config = RpcAccountInfoConfig {
encoding: Some(UiAccountEncoding::Base64),
data_slice: Some(UiDataSliceConfig {
offset: 0,
length: 8,
}),
commitment: None,
min_context_slot,
};
let mut tiers = Vec::with_capacity(pdas.len());
for chunk in pdas.chunks(MAX_MULTIPLE_ACCOUNTS) {
let accounts = rpc
.get_multiple_ui_accounts_with_config(chunk, config.clone())
.await?
.value;
for account in accounts {
tiers.push(
account
.and_then(|acc| acc.data.decode())
.and_then(|data| mempool_tier_from_discriminator(&data)),
);
}
}
Ok(tiers)
}
pub async fn get_execpool_acc_data(
rpc: &AsyncRpcClient,
execpool_acc: &Pubkey,
) -> Result<ExecpoolWrapper, ComputationPoolError> {
let execpool_data = rpc
.get_account_data(execpool_acc)
.await
.map_err(ComputationPoolError::new_solana_error)?;
ExecpoolWrapper::from_raw(&execpool_data)
}
pub async fn get_execpool_acc_data_raw(
rpc: &AsyncRpcClient,
execpool_acc: &Pubkey,
) -> Result<Vec<u8>, SolanaClientError> {
let execpool_data = rpc.get_account_data(execpool_acc).await?;
Ok(execpool_data)
}
#[derive(Debug, Hash, PartialEq, Eq, Copy, Clone)]
pub struct MempoolInfo {
pub cluster: Pubkey,
pub mxe: Pubkey,
pub mempool: Pubkey,
}
pub enum MempoolWrapper {
Tiny(Box<TinyMempool>),
Small(Box<SmallMempool>),
Medium(Box<MediumMempool>),
Large(Box<LargeMempool>),
}
#[derive(Debug)]
pub enum ComputationPoolError {
InvalidDiscriminator,
InvalidSize,
InvalidStartIndex { start_index: usize, capacity: usize },
ClientError(Box<SolanaClientError>),
}
impl ComputationPoolError {
pub fn new_solana_error(err: SolanaClientError) -> Self {
ComputationPoolError::ClientError(Box::new(err))
}
}
macro_rules! extract_computations {
($inner:expr) => {{
let start_index = $inner.computations.start_index as usize;
let buffer_size = $inner.computations.elems.len();
let length = $inner.computations.length as usize;
$inner
.computations
.elems
.iter()
.enumerate()
.filter(|(i, _)| {
let normalized_i = if *i >= start_index {
*i - start_index
} else {
buffer_size - start_index + *i
};
Self::is_valid(&$inner.computations.valid_bits, *i) && normalized_i < length
})
.flat_map(|(_, h)| h.entries.iter().copied())
.filter(|computation| !is_empty_computation_ref(computation))
.collect()
}};
}
macro_rules! extract_computations_highest_prio {
($inner:expr) => {{
let start_index = $inner.computations.start_index as usize;
let buffer_size = $inner.computations.elems.len();
let length = $inner.computations.length as usize;
$inner
.computations
.elems
.iter()
.enumerate()
.filter_map(|(i, h)| {
let normalized_i = if i >= start_index {
i - start_index
} else {
buffer_size - start_index + i
};
if Self::is_valid(&$inner.computations.valid_bits, i) && normalized_i < length {
let first = h.entries.first().copied()?;
if !is_empty_computation_ref(&first) {
return Some(first);
}
}
None
})
.collect()
}};
}
macro_rules! extract_computations_with_offset {
($inner:expr) => {{
let start_index = $inner.computations.start_index as usize;
let buffer_size = $inner.computations.elems.len();
let length = $inner.computations.length as usize;
(0..length)
.filter_map(|offset| {
let physical_index = (start_index + offset) % buffer_size;
if Self::is_valid(&$inner.computations.valid_bits, physical_index) {
Some((offset, physical_index))
} else {
None
}
})
.flat_map(|(offset, physical_index)| {
$inner.computations.elems[physical_index]
.entries
.iter()
.copied()
.filter(|c| !is_empty_computation_ref(c))
.map(move |c| (offset, c))
})
.collect()
}};
}
macro_rules! deserialize_mempool {
($raw:expr, $mempool:ty, $variant:ident) => {{
let offset = <$mempool as Discriminator>::DISCRIMINATOR.len();
if offset + std::mem::size_of::<$mempool>() > $raw.len() {
return Err(ComputationPoolError::InvalidSize);
}
let data = bytemuck::pod_read_unaligned::<$mempool>(
&$raw[offset..offset + std::mem::size_of::<$mempool>()],
);
let capacity = data.inner.computations.elems.len();
let start_index = data.inner.computations.start_index as usize;
if start_index >= capacity {
return Err(ComputationPoolError::InvalidStartIndex {
start_index,
capacity,
});
}
Ok(MempoolWrapper::$variant(Box::new(data)))
}};
}
impl MempoolWrapper {
pub fn last_updated_slot(&self) -> u64 {
match self {
MempoolWrapper::Tiny(m) => m.inner.last_updated_slot,
MempoolWrapper::Small(m) => m.inner.last_updated_slot,
MempoolWrapper::Medium(m) => m.inner.last_updated_slot,
MempoolWrapper::Large(m) => m.inner.last_updated_slot,
}
}
pub fn tier(&self) -> MempoolSize {
match self {
MempoolWrapper::Tiny(_) => MempoolSize::Tiny,
MempoolWrapper::Small(_) => MempoolSize::Small,
MempoolWrapper::Medium(_) => MempoolSize::Medium,
MempoolWrapper::Large(_) => MempoolSize::Large,
}
}
pub fn slot_capacity(&self) -> usize {
match self {
MempoolWrapper::Tiny(m) => m.inner.computations.elems[0].entries.len(),
MempoolWrapper::Small(m) => m.inner.computations.elems[0].entries.len(),
MempoolWrapper::Medium(m) => m.inner.computations.elems[0].entries.len(),
MempoolWrapper::Large(m) => m.inner.computations.elems[0].entries.len(),
}
}
pub fn computations_raw(self) -> Vec<(bool, Vec<ComputationReference>, usize, usize)> {
match self {
MempoolWrapper::Tiny(tm) => {
let len = tm.inner.computations.elems.len();
let start_index = tm.inner.computations.start_index as usize;
let mut res = vec![Default::default(); len];
for (i, h) in tm.inner.computations.elems.iter().enumerate() {
let normalized_i = if i >= start_index {
i - start_index
} else {
len - start_index + i
};
res[normalized_i] = (
Self::is_valid(&tm.inner.computations.valid_bits, i),
h.entries
.iter()
.copied()
.filter(|c| !is_empty_computation_ref(c))
.collect(),
normalized_i,
i,
);
}
res
}
MempoolWrapper::Small(sm) => {
let len = sm.inner.computations.elems.len();
let start_index = sm.inner.computations.start_index as usize;
let mut res = vec![Default::default(); len];
for (i, h) in sm.inner.computations.elems.iter().enumerate() {
let normalized_i = if i >= start_index {
i - start_index
} else {
len - start_index + i
};
res[normalized_i] = (
Self::is_valid(&sm.inner.computations.valid_bits, i),
h.entries
.iter()
.copied()
.filter(|c| !is_empty_computation_ref(c))
.collect(),
normalized_i,
i,
);
}
res
}
MempoolWrapper::Medium(mm) => {
let len = mm.inner.computations.elems.len();
let start_index = mm.inner.computations.start_index as usize;
let mut res = vec![Default::default(); len];
for (i, h) in mm.inner.computations.elems.iter().enumerate() {
let normalized_i = if i >= start_index {
i - start_index
} else {
len - start_index + i
};
res[normalized_i] = (
Self::is_valid(&mm.inner.computations.valid_bits, i),
h.entries
.iter()
.copied()
.filter(|c| !is_empty_computation_ref(c))
.collect(),
normalized_i,
i,
);
}
res
}
MempoolWrapper::Large(lm) => {
let len = lm.inner.computations.elems.len();
let start_index = lm.inner.computations.start_index as usize;
let mut res = vec![Default::default(); len];
for (i, h) in lm.inner.computations.elems.iter().enumerate() {
let normalized_i = if i >= start_index {
i - start_index
} else {
len - start_index + i
};
res[normalized_i] = (
Self::is_valid(&lm.inner.computations.valid_bits, i),
h.entries
.iter()
.copied()
.filter(|c| !is_empty_computation_ref(c))
.collect(),
normalized_i,
i,
);
}
res
}
}
}
pub fn computations(self) -> Vec<ComputationReference> {
match self {
MempoolWrapper::Tiny(tm) => extract_computations!(tm.inner),
MempoolWrapper::Small(sm) => extract_computations!(sm.inner),
MempoolWrapper::Medium(mm) => extract_computations!(mm.inner),
MempoolWrapper::Large(lm) => extract_computations!(lm.inner),
}
}
pub fn computations_with_offset(self) -> Vec<(usize, ComputationReference)> {
match self {
MempoolWrapper::Tiny(tm) => extract_computations_with_offset!(tm.inner),
MempoolWrapper::Small(sm) => extract_computations_with_offset!(sm.inner),
MempoolWrapper::Medium(mm) => extract_computations_with_offset!(mm.inner),
MempoolWrapper::Large(lm) => extract_computations_with_offset!(lm.inner),
}
}
pub fn computations_highest_prio(self) -> Vec<ComputationReference> {
match self {
MempoolWrapper::Tiny(tm) => extract_computations_highest_prio!(tm.inner),
MempoolWrapper::Small(sm) => extract_computations_highest_prio!(sm.inner),
MempoolWrapper::Medium(mm) => extract_computations_highest_prio!(mm.inner),
MempoolWrapper::Large(lm) => extract_computations_highest_prio!(lm.inner),
}
}
pub fn from_raw(raw_mempool: &[u8]) -> Result<Self, ComputationPoolError> {
if raw_mempool.len() < 8 {
return Err(ComputationPoolError::InvalidSize);
}
match &raw_mempool[0..8] {
TinyMempool::DISCRIMINATOR => deserialize_mempool!(raw_mempool, TinyMempool, Tiny),
SmallMempool::DISCRIMINATOR => deserialize_mempool!(raw_mempool, SmallMempool, Small),
MediumMempool::DISCRIMINATOR => {
deserialize_mempool!(raw_mempool, MediumMempool, Medium)
}
LargeMempool::DISCRIMINATOR => deserialize_mempool!(raw_mempool, LargeMempool, Large),
_ => Err(ComputationPoolError::InvalidDiscriminator),
}
}
fn is_valid(valid_bits: &[u8], idx: usize) -> bool {
let byte = idx / 8;
let bit = idx - (byte * 8);
if byte >= valid_bits.len() {
return false;
}
(valid_bits[byte] & (1 << bit)) != 0
}
}
pub fn is_empty_computation_ref(c: &ComputationReference) -> bool {
*c == ComputationReference::zeroed()
}
impl std::fmt::Display for ComputationReference {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Computation offset: {}, priority fee: {}",
self.computation_offset, self.priority_fee
)
}
}
#[derive(Debug, Hash, PartialEq, Eq, Copy, Clone)]
pub struct ExecpoolInfo {
pub cluster: Pubkey,
pub mxe: Pubkey,
pub execpool: Pubkey,
}
pub enum ExecpoolWrapper {
Tiny(Box<TinyExecPool>),
Small(Box<SmallExecPool>),
Medium(Box<MediumExecPool>),
Large(Box<LargeExecPool>),
}
impl ExecpoolWrapper {
pub fn tier(&self) -> MempoolSize {
match self {
ExecpoolWrapper::Tiny(_) => MempoolSize::Tiny,
ExecpoolWrapper::Small(_) => MempoolSize::Small,
ExecpoolWrapper::Medium(_) => MempoolSize::Medium,
ExecpoolWrapper::Large(_) => MempoolSize::Large,
}
}
pub fn max_parallel(&self) -> usize {
match self {
ExecpoolWrapper::Tiny(p) => p.inner.currently_executing.len(),
ExecpoolWrapper::Small(p) => p.inner.currently_executing.len(),
ExecpoolWrapper::Medium(p) => p.inner.currently_executing.len(),
ExecpoolWrapper::Large(p) => p.inner.currently_executing.len(),
}
}
pub fn computations_unfiltered(self) -> Vec<ComputationReferenceWIndex> {
match self {
ExecpoolWrapper::Tiny(tm) => tm
.inner
.currently_executing
.into_iter()
.enumerate()
.map(|(i, reference)| ComputationReferenceWIndex {
reference,
index: tm.inner.meta[i].index,
})
.collect(),
ExecpoolWrapper::Small(sm) => sm
.inner
.currently_executing
.into_iter()
.enumerate()
.map(|(i, reference)| ComputationReferenceWIndex {
reference,
index: sm.inner.meta[i].index,
})
.collect(),
ExecpoolWrapper::Medium(mm) => mm
.inner
.currently_executing
.into_iter()
.enumerate()
.map(|(i, reference)| ComputationReferenceWIndex {
reference,
index: mm.inner.meta[i].index,
})
.collect(),
ExecpoolWrapper::Large(lm) => lm
.inner
.currently_executing
.into_iter()
.enumerate()
.map(|(i, reference)| ComputationReferenceWIndex {
reference,
index: lm.inner.meta[i].index,
})
.collect(),
}
}
pub fn computations(self) -> Vec<ComputationReferenceWIndex> {
self.computations_unfiltered()
.into_iter()
.filter(|computation| !is_empty_computation_ref(&computation.reference))
.collect()
}
pub fn from_raw(raw_mempool: &[u8]) -> Result<Self, ComputationPoolError> {
if raw_mempool.len() < 8 {
return Err(ComputationPoolError::InvalidSize);
}
match &raw_mempool[0..8] {
TinyExecPool::DISCRIMINATOR => {
let offset = TinyExecPool::DISCRIMINATOR.len();
if offset + std::mem::size_of::<TinyExecPool>() > raw_mempool.len() {
return Err(ComputationPoolError::InvalidSize);
}
let te = bytemuck::pod_read_unaligned::<TinyExecPool>(
&raw_mempool[offset..offset + std::mem::size_of::<TinyExecPool>()],
);
Ok(ExecpoolWrapper::Tiny(Box::new(te)))
}
SmallExecPool::DISCRIMINATOR => {
let offset = SmallExecPool::DISCRIMINATOR.len();
if offset + std::mem::size_of::<SmallExecPool>() > raw_mempool.len() {
return Err(ComputationPoolError::InvalidSize);
}
let se = bytemuck::pod_read_unaligned::<SmallExecPool>(
&raw_mempool[offset..offset + std::mem::size_of::<SmallExecPool>()],
);
Ok(ExecpoolWrapper::Small(Box::new(se)))
}
MediumExecPool::DISCRIMINATOR => {
let offset = MediumExecPool::DISCRIMINATOR.len();
if offset + std::mem::size_of::<MediumExecPool>() > raw_mempool.len() {
return Err(ComputationPoolError::InvalidSize);
}
let me = bytemuck::pod_read_unaligned::<MediumExecPool>(
&raw_mempool[offset..offset + std::mem::size_of::<MediumExecPool>()],
);
Ok(ExecpoolWrapper::Medium(Box::new(me)))
}
LargeExecPool::DISCRIMINATOR => {
let offset = LargeExecPool::DISCRIMINATOR.len();
if offset + std::mem::size_of::<LargeExecPool>() > raw_mempool.len() {
return Err(ComputationPoolError::InvalidSize);
}
let le = bytemuck::pod_read_unaligned::<LargeExecPool>(
&raw_mempool[offset..offset + std::mem::size_of::<LargeExecPool>()],
);
Ok(ExecpoolWrapper::Large(Box::new(le)))
}
_ => Err(ComputationPoolError::InvalidDiscriminator),
}
}
}
#[derive(PartialEq, Eq, Debug, Clone)]
pub struct ComputationReferenceWIndex {
pub reference: ComputationReference,
pub index: u64,
}
impl LeaderSelector {
pub fn new_with_size(size: usize) -> Self {
let mut selector = Self::default();
selector.info.resize(size, Default::default());
selector
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mempool_from_raw_empty_returns_invalid_size() {
let result = MempoolWrapper::from_raw(&[]);
assert!(matches!(result, Err(ComputationPoolError::InvalidSize)));
}
#[test]
fn mempool_from_raw_short_returns_invalid_size() {
let result = MempoolWrapper::from_raw(&[0u8; 7]);
assert!(matches!(result, Err(ComputationPoolError::InvalidSize)));
}
#[test]
fn mempool_from_raw_unknown_discriminator() {
let result = MempoolWrapper::from_raw(&[0xFF; 8]);
assert!(matches!(
result,
Err(ComputationPoolError::InvalidDiscriminator)
));
}
#[test]
fn execpool_from_raw_empty_returns_invalid_size() {
let result = ExecpoolWrapper::from_raw(&[]);
assert!(matches!(result, Err(ComputationPoolError::InvalidSize)));
}
#[test]
fn execpool_from_raw_short_returns_invalid_size() {
let result = ExecpoolWrapper::from_raw(&[0u8; 7]);
assert!(matches!(result, Err(ComputationPoolError::InvalidSize)));
}
#[test]
fn execpool_from_raw_unknown_discriminator() {
let result = ExecpoolWrapper::from_raw(&[0xFF; 8]);
assert!(matches!(
result,
Err(ComputationPoolError::InvalidDiscriminator)
));
}
#[test]
fn mempool_from_raw_valid_discriminator_short_body_returns_invalid_size() {
use anchor_lang::Discriminator;
let mut raw = Vec::with_capacity(32);
raw.extend_from_slice(TinyMempool::DISCRIMINATOR);
raw.resize(32, 0u8); let result = MempoolWrapper::from_raw(&raw);
assert!(matches!(result, Err(ComputationPoolError::InvalidSize)));
}
#[test]
fn is_valid_returns_false_for_index_past_valid_bits() {
let bits = [0xFF_u8];
assert!(!MempoolWrapper::is_valid(&bits, 8));
assert!(!MempoolWrapper::is_valid(&bits, usize::MAX));
}
fn expected_normalized(i: usize, start_index: usize, len: usize) -> usize {
if i >= start_index {
i - start_index
} else {
len - start_index + i
}
}
fn encoded_mempool_bytes<M>(mp: &M) -> Vec<u8>
where
M: bytemuck::Pod + anchor_lang::Discriminator,
{
let mut raw = Vec::with_capacity(M::DISCRIMINATOR.len() + std::mem::size_of::<M>());
raw.extend_from_slice(M::DISCRIMINATOR);
raw.extend_from_slice(bytemuck::bytes_of(mp));
raw
}
#[test]
fn from_raw_rejects_out_of_range_start_index_tiny() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let capacity = mp.inner.computations.elems.len();
assert!(capacity > 0 && capacity < u8::MAX as usize);
mp.inner.computations.start_index = u8::MAX;
let raw = encoded_mempool_bytes(&*mp);
assert!(matches!(
MempoolWrapper::from_raw(&raw),
Err(ComputationPoolError::InvalidStartIndex { start_index, capacity: cap })
if start_index == u8::MAX as usize && cap == capacity,
));
}
#[test]
fn from_raw_rejects_start_index_equal_to_capacity_tiny() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let capacity = mp.inner.computations.elems.len();
assert!(capacity > 0 && capacity <= u8::MAX as usize);
mp.inner.computations.start_index = capacity as u8;
let raw = encoded_mempool_bytes(&*mp);
assert!(matches!(
MempoolWrapper::from_raw(&raw),
Err(ComputationPoolError::InvalidStartIndex { .. }),
));
}
#[test]
fn from_raw_accepts_in_range_start_index_and_rotates_tiny() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let len = mp.inner.computations.elems.len();
assert!(len > 1);
let start = len / 2;
mp.inner.computations.start_index = start as u8;
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw).expect("in-range start_index must parse");
let out = wrapper.computations_raw();
assert_eq!(out.len(), len);
for i in 0..len {
let expected = expected_normalized(i, start, len);
assert_eq!(out[expected].3, i);
}
}
#[test]
fn from_raw_accepts_start_index_at_upper_bound_tiny() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let capacity = mp.inner.computations.elems.len();
assert!(capacity > 0 && capacity <= u8::MAX as usize);
mp.inner.computations.start_index = (capacity - 1) as u8;
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw)
.expect("start_index == capacity - 1 is in-range and must parse");
let out = wrapper.computations_raw();
assert_eq!(out.len(), capacity);
for i in 0..capacity {
let expected = expected_normalized(i, capacity - 1, capacity);
assert_eq!(out[expected].3, i);
}
}
#[test]
fn computations_raw_tiny_rotates_mid_buffer_correctly() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let len = mp.inner.computations.elems.len();
assert!(len > 1);
let start = len / 2;
mp.inner.computations.start_index = start as u8;
let out = MempoolWrapper::Tiny(mp).computations_raw();
assert_eq!(out.len(), len);
for i in 0..len {
let expected = expected_normalized(i, start, len);
assert_eq!(out[expected].3, i);
}
}
#[test]
fn last_updated_slot_roundtrips_through_deserialization() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
mp.inner.last_updated_slot = 4_242;
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw).expect("must parse");
assert_eq!(wrapper.last_updated_slot(), 4_242);
}
fn fabricate_comp(offset_id: u64) -> ComputationReference {
let mut comp: ComputationReference = bytemuck::Zeroable::zeroed();
comp.computation_offset = offset_id;
comp
}
#[test]
fn computations_with_offset_yields_logical_offsets_unrotated() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
mp.inner.last_updated_slot = 100;
mp.inner.computations.start_index = 0;
mp.inner.computations.length = 3;
mp.inner.computations.valid_bits[0] = 0b0000_0111;
mp.inner.computations.elems[0].entries[0] = fabricate_comp(10);
mp.inner.computations.elems[1].entries[0] = fabricate_comp(20);
mp.inner.computations.elems[2].entries[0] = fabricate_comp(30);
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw).expect("must parse");
let out = wrapper.computations_with_offset();
assert_eq!(out.len(), 3);
assert_eq!(out[0], (0, fabricate_comp(10)));
assert_eq!(out[1], (1, fabricate_comp(20)));
assert_eq!(out[2], (2, fabricate_comp(30)));
}
#[test]
fn computations_with_offset_preserves_logical_order_when_rotated() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let buf_len = mp.inner.computations.elems.len();
let start = buf_len - 1;
mp.inner.computations.start_index =
u8::try_from(start).expect("fits in u8 for TinyMempool");
mp.inner.computations.length = 2;
let phys_0 = start % buf_len;
let phys_1 = (start + 1) % buf_len;
mp.inner.computations.valid_bits[phys_0 / 8] |= 1 << (phys_0 % 8);
mp.inner.computations.valid_bits[phys_1 / 8] |= 1 << (phys_1 % 8);
mp.inner.computations.elems[phys_0].entries[0] = fabricate_comp(10);
mp.inner.computations.elems[phys_1].entries[0] = fabricate_comp(20);
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw).expect("must parse");
let out = wrapper.computations_with_offset();
assert_eq!(out.len(), 2);
assert_eq!(out[0], (0, fabricate_comp(10)));
assert_eq!(out[1], (1, fabricate_comp(20)));
}
#[test]
fn computations_filters_by_physical_valid_bits_when_rotated() {
let mut mp: Box<TinyMempool> = bytemuck::zeroed_box();
let buf_len = mp.inner.computations.elems.len();
let start = buf_len / 2;
mp.inner.computations.start_index =
u8::try_from(start).expect("fits in u8 for TinyMempool");
mp.inner.computations.length = 2;
let phys_0 = start % buf_len;
let phys_1 = (start + 1) % buf_len;
mp.inner.computations.valid_bits[phys_0 / 8] |= 1 << (phys_0 % 8);
mp.inner.computations.valid_bits[phys_1 / 8] |= 1 << (phys_1 % 8);
mp.inner.computations.elems[phys_0].entries[0] = fabricate_comp(11);
mp.inner.computations.elems[phys_1].entries[0] = fabricate_comp(22);
let raw = encoded_mempool_bytes(&*mp);
let wrapper = MempoolWrapper::from_raw(&raw).expect("must parse");
let out = wrapper.computations();
assert_eq!(out.len(), 2);
assert!(out.iter().any(|c| c == &fabricate_comp(11)));
assert!(out.iter().any(|c| c == &fabricate_comp(22)));
}
#[test]
fn mempool_wrapper_reports_tier_and_slot_capacity() {
let tiny = MempoolWrapper::Tiny(bytemuck::zeroed_box());
assert!(matches!(tiny.tier(), MempoolSize::Tiny));
assert_eq!(tiny.slot_capacity(), 1);
let small = MempoolWrapper::Small(bytemuck::zeroed_box());
assert!(matches!(small.tier(), MempoolSize::Small));
assert_eq!(small.slot_capacity(), 3);
let medium = MempoolWrapper::Medium(bytemuck::zeroed_box());
assert!(matches!(medium.tier(), MempoolSize::Medium));
assert_eq!(medium.slot_capacity(), 10);
let large = MempoolWrapper::Large(bytemuck::zeroed_box());
assert!(matches!(large.tier(), MempoolSize::Large));
assert_eq!(large.slot_capacity(), 100);
}
#[test]
fn execpool_wrapper_reports_tier_and_max_parallel() {
let tiny = ExecpoolWrapper::Tiny(bytemuck::zeroed_box());
assert!(matches!(tiny.tier(), MempoolSize::Tiny));
assert_eq!(tiny.max_parallel(), 1);
let small = ExecpoolWrapper::Small(bytemuck::zeroed_box());
assert!(matches!(small.tier(), MempoolSize::Small));
assert_eq!(small.max_parallel(), 3);
let medium = ExecpoolWrapper::Medium(bytemuck::zeroed_box());
assert!(matches!(medium.tier(), MempoolSize::Medium));
assert_eq!(medium.max_parallel(), 10);
let large = ExecpoolWrapper::Large(bytemuck::zeroed_box());
assert!(matches!(large.tier(), MempoolSize::Large));
assert_eq!(large.max_parallel(), 100);
}
#[test]
fn mempool_tier_from_discriminator_maps_all_tiers() {
assert!(matches!(
mempool_tier_from_discriminator(TinyMempool::DISCRIMINATOR),
Some(MempoolSize::Tiny)
));
assert!(matches!(
mempool_tier_from_discriminator(SmallMempool::DISCRIMINATOR),
Some(MempoolSize::Small)
));
assert!(matches!(
mempool_tier_from_discriminator(MediumMempool::DISCRIMINATOR),
Some(MempoolSize::Medium)
));
assert!(matches!(
mempool_tier_from_discriminator(LargeMempool::DISCRIMINATOR),
Some(MempoolSize::Large)
));
assert!(mempool_tier_from_discriminator(&[0xFF; 8]).is_none());
}
}