try-runtime::follow-chain - keep connection (#12167)

* Refactor RPC module

* Add flag to `follow-chain`

* Multithreading remark

* fmt

* O_O

* unused import

* cmon

* accidental removal reverted

* remove RpcHeaderProvider

* mut refs

* fmt

* no mutability

* now?

* now?

* arc mutex

* async mutex

* async mutex

* uhm

* connect in constructor

* remove dep

* old import

* another take

* trigger polkadot pipeline

* trigger pipeline
This commit is contained in:
Piotr Mikołajczyk
2022-09-06 10:01:35 +02:00
committed by GitHub
parent d213e95784
commit 198f94f931
6 changed files with 161 additions and 107 deletions
+3 -2
View File
@@ -24,7 +24,7 @@ use nix::{
unistd::Pid, unistd::Pid,
}; };
use node_primitives::Block; use node_primitives::Block;
use remote_externalities::rpc_api; use remote_externalities::rpc_api::RpcService;
use std::{ use std::{
io::{BufRead, BufReader, Read}, io::{BufRead, BufReader, Read},
ops::{Deref, DerefMut}, ops::{Deref, DerefMut},
@@ -71,9 +71,10 @@ pub async fn wait_n_finalized_blocks(
pub async fn wait_n_finalized_blocks_from(n: usize, url: &str) { pub async fn wait_n_finalized_blocks_from(n: usize, url: &str) {
let mut built_blocks = std::collections::HashSet::new(); let mut built_blocks = std::collections::HashSet::new();
let mut interval = tokio::time::interval(Duration::from_secs(2)); let mut interval = tokio::time::interval(Duration::from_secs(2));
let rpc_service = RpcService::new(url, false).await.unwrap();
loop { loop {
if let Ok(block) = rpc_api::get_finalized_head::<Block, _>(url.to_string()).await { if let Ok(block) = rpc_service.get_finalized_head::<Block>().await {
built_blocks.insert(block); built_blocks.insert(block);
if built_blocks.len() > n { if built_blocks.len() > n {
break break
@@ -19,82 +19,131 @@
// TODO: Consolidate one off RPC calls https://github.com/paritytech/substrate/issues/8988 // TODO: Consolidate one off RPC calls https://github.com/paritytech/substrate/issues/8988
use jsonrpsee::{ use jsonrpsee::{
core::client::ClientT, core::client::{Client, ClientT},
rpc_params, rpc_params,
types::ParamsSer,
ws_client::{WsClient, WsClientBuilder}, ws_client::{WsClient, WsClientBuilder},
}; };
use sp_runtime::{ use serde::de::DeserializeOwned;
generic::SignedBlock, use sp_runtime::{generic::SignedBlock, traits::Block as BlockT};
traits::{Block as BlockT, Header as HeaderT}, use std::sync::Arc;
};
/// Get the header of the block identified by `at` enum RpcCall {
pub async fn get_header<Block, S>(from: S, at: Block::Hash) -> Result<Block::Header, String> GetHeader,
where GetFinalizedHead,
Block: BlockT, GetBlock,
Block::Header: serde::de::DeserializeOwned, GetRuntimeVersion,
S: AsRef<str>, }
{
let client = build_client(from).await?;
impl RpcCall {
fn as_str(&self) -> &'static str {
match self {
RpcCall::GetHeader => "chain_getHeader",
RpcCall::GetFinalizedHead => "chain_getFinalizedHead",
RpcCall::GetBlock => "chain_getBlock",
RpcCall::GetRuntimeVersion => "state_getRuntimeVersion",
}
}
}
/// General purpose method for making RPC calls.
async fn make_request<'a, T: DeserializeOwned>(
client: &Arc<Client>,
call: RpcCall,
params: Option<ParamsSer<'a>>,
) -> Result<T, String> {
client client
.request::<Block::Header>("chain_getHeader", rpc_params!(at)) .request::<T>(call.as_str(), params)
.await .await
.map_err(|e| format!("chain_getHeader request failed: {:?}", e)) .map_err(|e| format!("{} request failed: {:?}", call.as_str(), e))
} }
/// Get the finalized head enum ConnectionPolicy {
pub async fn get_finalized_head<Block, S>(from: S) -> Result<Block::Hash, String> Reuse(Arc<Client>),
where Reconnect,
Block: BlockT,
S: AsRef<str>,
{
let client = build_client(from).await?;
client
.request::<Block::Hash>("chain_getFinalizedHead", None)
.await
.map_err(|e| format!("chain_getFinalizedHead request failed: {:?}", e))
} }
/// Get the signed block identified by `at`. /// Simple RPC service that is capable of keeping the connection.
pub async fn get_block<Block, S>(from: S, at: Block::Hash) -> Result<Block, String> ///
where /// Service will connect to `uri` for the first time already during initialization.
S: AsRef<str>, ///
Block: BlockT + serde::de::DeserializeOwned, /// Be careful with reusing the connection in a multithreaded environment.
Block::Header: HeaderT, pub struct RpcService {
{ uri: String,
let client = build_client(from).await?; policy: ConnectionPolicy,
let signed_block = client
.request::<SignedBlock<Block>>("chain_getBlock", rpc_params!(at))
.await
.map_err(|e| format!("chain_getBlock request failed: {:?}", e))?;
Ok(signed_block.block)
} }
/// Build a websocket client that connects to `from`. impl RpcService {
async fn build_client<S: AsRef<str>>(from: S) -> Result<WsClient, String> { /// Creates a new RPC service. If `keep_connection`, then connects to `uri` right away.
WsClientBuilder::default() pub async fn new<S: AsRef<str>>(uri: S, keep_connection: bool) -> Result<Self, String> {
.max_request_body_size(u32::MAX) let policy = if keep_connection {
.build(from.as_ref()) ConnectionPolicy::Reuse(Arc::new(Self::build_client(uri.as_ref()).await?))
.await } else {
.map_err(|e| format!("`WsClientBuilder` failed to build: {:?}", e)) ConnectionPolicy::Reconnect
} };
Ok(Self { uri: uri.as_ref().to_string(), policy })
}
/// Get the runtime version of a given chain. /// Returns the address at which requests are sent.
pub async fn get_runtime_version<Block, S>( pub fn uri(&self) -> String {
from: S, self.uri.clone()
at: Option<Block::Hash>, }
) -> Result<sp_version::RuntimeVersion, String>
where /// Build a websocket client that connects to `self.uri`.
S: AsRef<str>, async fn build_client<S: AsRef<str>>(uri: S) -> Result<WsClient, String> {
Block: BlockT + serde::de::DeserializeOwned, WsClientBuilder::default()
Block::Header: HeaderT, .max_request_body_size(u32::MAX)
{ .build(uri)
let client = build_client(from).await?; .await
client .map_err(|e| format!("`WsClientBuilder` failed to build: {:?}", e))
.request::<sp_version::RuntimeVersion>("state_getRuntimeVersion", rpc_params!(at)) }
.await
.map_err(|e| format!("state_getRuntimeVersion request failed: {:?}", e)) /// Generic method for making RPC requests.
async fn make_request<'a, T: DeserializeOwned>(
&self,
call: RpcCall,
params: Option<ParamsSer<'a>>,
) -> Result<T, String> {
match self.policy {
// `self.keep_connection` must have been `true`.
ConnectionPolicy::Reuse(ref client) => make_request(client, call, params).await,
ConnectionPolicy::Reconnect => {
let client = Arc::new(Self::build_client(&self.uri).await?);
make_request(&client, call, params).await
},
}
}
/// Get the header of the block identified by `at`.
pub async fn get_header<Block>(&self, at: Block::Hash) -> Result<Block::Header, String>
where
Block: BlockT,
Block::Header: DeserializeOwned,
{
self.make_request(RpcCall::GetHeader, rpc_params!(at)).await
}
/// Get the finalized head.
pub async fn get_finalized_head<Block: BlockT>(&self) -> Result<Block::Hash, String> {
self.make_request(RpcCall::GetFinalizedHead, None).await
}
/// Get the signed block identified by `at`.
pub async fn get_block<Block: BlockT + DeserializeOwned>(
&self,
at: Block::Hash,
) -> Result<Block, String> {
Ok(self
.make_request::<SignedBlock<Block>>(RpcCall::GetBlock, rpc_params!(at))
.await?
.block)
}
/// Get the runtime version of a given chain.
pub async fn get_runtime_version<Block: BlockT + DeserializeOwned>(
&self,
at: Option<Block::Hash>,
) -> Result<sp_version::RuntimeVersion, String> {
self.make_request(RpcCall::GetRuntimeVersion, rpc_params!(at)).await
}
} }
@@ -89,6 +89,8 @@ impl ExecuteBlockCmd {
Block::Hash: FromStr, Block::Hash: FromStr,
<Block::Hash as FromStr>::Err: Debug, <Block::Hash as FromStr>::Err: Debug,
{ {
let rpc_service = rpc_api::RpcService::new(ws_uri, false).await?;
match (&self.block_at, &self.state) { match (&self.block_at, &self.state) {
(Some(block_at), State::Snap { .. }) => hash_of::<Block>(block_at), (Some(block_at), State::Snap { .. }) => hash_of::<Block>(block_at),
(Some(block_at), State::Live { .. }) => { (Some(block_at), State::Live { .. }) => {
@@ -100,9 +102,7 @@ impl ExecuteBlockCmd {
target: LOG_TARGET, target: LOG_TARGET,
"No --block-at or --at provided, using the latest finalized block instead" "No --block-at or --at provided, using the latest finalized block instead"
); );
remote_externalities::rpc_api::get_finalized_head::<Block, _>(ws_uri) rpc_service.get_finalized_head::<Block>().await.map_err(Into::into)
.await
.map_err(Into::into)
}, },
(None, State::Live { at: Some(at), .. }) => hash_of::<Block>(at), (None, State::Live { at: Some(at), .. }) => hash_of::<Block>(at),
_ => { _ => {
@@ -148,7 +148,8 @@ where
let block_ws_uri = command.block_ws_uri::<Block>(); let block_ws_uri = command.block_ws_uri::<Block>();
let block_at = command.block_at::<Block>(block_ws_uri.clone()).await?; let block_at = command.block_at::<Block>(block_ws_uri.clone()).await?;
let block: Block = rpc_api::get_block::<Block, _>(block_ws_uri.clone(), block_at).await?; let rpc_service = rpc_api::RpcService::new(block_ws_uri.clone(), false).await?;
let block: Block = rpc_service.get_block::<Block>(block_at).await?;
let parent_hash = block.header().parent_hash(); let parent_hash = block.header().parent_hash();
log::info!( log::info!(
target: LOG_TARGET, target: LOG_TARGET,
@@ -27,13 +27,13 @@ use jsonrpsee::{
ws_client::WsClientBuilder, ws_client::WsClientBuilder,
}; };
use parity_scale_codec::{Decode, Encode}; use parity_scale_codec::{Decode, Encode};
use remote_externalities::{rpc_api, Builder, Mode, OnlineConfig}; use remote_externalities::{rpc_api::RpcService, Builder, Mode, OnlineConfig};
use sc_executor::NativeExecutionDispatch; use sc_executor::NativeExecutionDispatch;
use sc_service::Configuration; use sc_service::Configuration;
use serde::de::DeserializeOwned; use serde::de::DeserializeOwned;
use sp_core::H256; use sp_core::H256;
use sp_runtime::traits::{Block as BlockT, Header as HeaderT, NumberFor}; use sp_runtime::traits::{Block as BlockT, Header as HeaderT, NumberFor};
use std::{collections::VecDeque, fmt::Debug, marker::PhantomData, str::FromStr}; use std::{collections::VecDeque, fmt::Debug, str::FromStr};
const SUB: &str = "chain_subscribeFinalizedHeads"; const SUB: &str = "chain_subscribeFinalizedHeads";
const UN_SUB: &str = "chain_unsubscribeFinalizedHeads"; const UN_SUB: &str = "chain_unsubscribeFinalizedHeads";
@@ -60,6 +60,10 @@ pub struct FollowChainCmd {
/// round-robin fashion. /// round-robin fashion.
#[clap(long, default_value = "none")] #[clap(long, default_value = "none")]
try_state: frame_try_runtime::TryStateSelect, try_state: frame_try_runtime::TryStateSelect,
/// If present, a single connection to a node will be kept and reused for fetching blocks.
#[clap(long)]
keep_connection: bool,
} }
/// Start listening for with `SUB` at `url`. /// Start listening for with `SUB` at `url`.
@@ -93,21 +97,16 @@ where
Block::Header: HeaderT, Block::Header: HeaderT,
{ {
/// Awaits for the header of the block with hash `hash`. /// Awaits for the header of the block with hash `hash`.
async fn get_header(&mut self, hash: Block::Hash) -> Block::Header; async fn get_header(&self, hash: Block::Hash) -> Block::Header;
}
struct RpcHeaderProvider<Block: BlockT> {
uri: String,
_phantom: PhantomData<Block>,
} }
#[async_trait] #[async_trait]
impl<Block: BlockT> HeaderProvider<Block> for RpcHeaderProvider<Block> impl<Block: BlockT> HeaderProvider<Block> for RpcService
where where
Block::Header: DeserializeOwned, Block::Header: DeserializeOwned,
{ {
async fn get_header(&mut self, hash: Block::Hash) -> Block::Header { async fn get_header(&self, hash: Block::Hash) -> Block::Header {
rpc_api::get_header::<Block, _>(&self.uri, hash).await.unwrap() self.get_header::<Block>(hash).await.unwrap()
} }
} }
@@ -148,19 +147,20 @@ where
/// ///
/// Returned headers are guaranteed to be ordered. There are no missing headers (even if some of /// Returned headers are guaranteed to be ordered. There are no missing headers (even if some of
/// them lack justification). /// them lack justification).
struct FinalizedHeaders<Block: BlockT, HP: HeaderProvider<Block>, HS: HeaderSubscription<Block>> { struct FinalizedHeaders<'a, Block: BlockT, HP: HeaderProvider<Block>, HS: HeaderSubscription<Block>>
header_provider: HP, {
header_provider: &'a HP,
subscription: HS, subscription: HS,
fetched_headers: VecDeque<Block::Header>, fetched_headers: VecDeque<Block::Header>,
last_returned: Option<<Block::Header as HeaderT>::Hash>, last_returned: Option<<Block::Header as HeaderT>::Hash>,
} }
impl<Block: BlockT, HP: HeaderProvider<Block>, HS: HeaderSubscription<Block>> impl<'a, Block: BlockT, HP: HeaderProvider<Block>, HS: HeaderSubscription<Block>>
FinalizedHeaders<Block, HP, HS> FinalizedHeaders<'a, Block, HP, HS>
where where
<Block as BlockT>::Header: DeserializeOwned, <Block as BlockT>::Header: DeserializeOwned,
{ {
pub fn new(header_provider: HP, subscription: HS) -> Self { pub fn new(header_provider: &'a HP, subscription: HS) -> Self {
Self { Self {
header_provider, header_provider,
subscription, subscription,
@@ -229,19 +229,16 @@ where
let executor = build_executor::<ExecDispatch>(&shared, &config); let executor = build_executor::<ExecDispatch>(&shared, &config);
let execution = shared.execution; let execution = shared.execution;
let header_provider: RpcHeaderProvider<Block> = let rpc_service = RpcService::new(&command.uri, command.keep_connection).await?;
RpcHeaderProvider { uri: command.uri.clone(), _phantom: PhantomData {} };
let mut finalized_headers: FinalizedHeaders< let mut finalized_headers: FinalizedHeaders<Block, RpcService, Subscription<Block::Header>> =
Block, FinalizedHeaders::new(&rpc_service, subscription);
RpcHeaderProvider<Block>,
Subscription<Block::Header>,
> = FinalizedHeaders::new(header_provider, subscription);
while let Some(header) = finalized_headers.next().await { while let Some(header) = finalized_headers.next().await {
let hash = header.hash(); let hash = header.hash();
let number = header.number(); let number = header.number();
let block = rpc_api::get_block::<Block, _>(&command.uri, hash).await.unwrap(); let block = rpc_service.get_block::<Block>(hash).await.unwrap();
log::debug!( log::debug!(
target: LOG_TARGET, target: LOG_TARGET,
@@ -333,12 +330,14 @@ where
mod tests { mod tests {
use super::*; use super::*;
use sp_runtime::testing::{Block as TBlock, ExtrinsicWrapper, Header}; use sp_runtime::testing::{Block as TBlock, ExtrinsicWrapper, Header};
use std::sync::Arc;
use tokio::sync::Mutex;
type Block = TBlock<ExtrinsicWrapper<()>>; type Block = TBlock<ExtrinsicWrapper<()>>;
type BlockNumber = u64; type BlockNumber = u64;
type Hash = H256; type Hash = H256;
struct MockHeaderProvider(pub VecDeque<BlockNumber>); struct MockHeaderProvider(pub Arc<Mutex<VecDeque<BlockNumber>>>);
fn headers() -> Vec<Header> { fn headers() -> Vec<Header> {
let mut headers = vec![Header::new_from_number(0)]; let mut headers = vec![Header::new_from_number(0)];
@@ -353,8 +352,8 @@ mod tests {
#[async_trait] #[async_trait]
impl HeaderProvider<Block> for MockHeaderProvider { impl HeaderProvider<Block> for MockHeaderProvider {
async fn get_header(&mut self, _hash: Hash) -> Header { async fn get_header(&self, _hash: Hash) -> Header {
let height = self.0.pop_front().unwrap(); let height = self.0.lock().await.pop_front().unwrap();
headers()[height as usize].clone() headers()[height as usize].clone()
} }
} }
@@ -372,9 +371,9 @@ mod tests {
async fn finalized_headers_works_when_every_block_comes_from_subscription() { async fn finalized_headers_works_when_every_block_comes_from_subscription() {
let heights = vec![4, 5, 6, 7]; let heights = vec![4, 5, 6, 7];
let provider = MockHeaderProvider(vec![].into()); let provider = MockHeaderProvider(Default::default());
let subscription = MockHeaderSubscription(heights.clone().into()); let subscription = MockHeaderSubscription(heights.clone().into());
let mut headers = FinalizedHeaders::new(provider, subscription); let mut headers = FinalizedHeaders::new(&provider, subscription);
for h in heights { for h in heights {
assert_eq!(h, headers.next().await.unwrap().number); assert_eq!(h, headers.next().await.unwrap().number);
@@ -389,9 +388,9 @@ mod tests {
// Consecutive headers will be requested in the reversed order. // Consecutive headers will be requested in the reversed order.
let heights_not_in_subscription = vec![5, 9, 8, 7]; let heights_not_in_subscription = vec![5, 9, 8, 7];
let provider = MockHeaderProvider(heights_not_in_subscription.into()); let provider = MockHeaderProvider(Arc::new(Mutex::new(heights_not_in_subscription.into())));
let subscription = MockHeaderSubscription(heights_in_subscription.into()); let subscription = MockHeaderSubscription(heights_in_subscription.into());
let mut headers = FinalizedHeaders::new(provider, subscription); let mut headers = FinalizedHeaders::new(&provider, subscription);
for h in all_heights { for h in all_heights {
assert_eq!(h, headers.next().await.unwrap().number); assert_eq!(h, headers.next().await.unwrap().number);
@@ -119,7 +119,8 @@ where
let header_at = command.header_at::<Block>()?; let header_at = command.header_at::<Block>()?;
let header_ws_uri = command.header_ws_uri::<Block>(); let header_ws_uri = command.header_ws_uri::<Block>();
let header = rpc_api::get_header::<Block, _>(header_ws_uri.clone(), header_at).await?; let rpc_service = rpc_api::RpcService::new(header_ws_uri.clone(), false).await?;
let header = rpc_service.get_header::<Block>(header_at).await?;
log::info!( log::info!(
target: LOG_TARGET, target: LOG_TARGET,
"fetched header from {:?}, block number: {:?}", "fetched header from {:?}, block number: {:?}",
@@ -267,7 +267,8 @@
use parity_scale_codec::Decode; use parity_scale_codec::Decode;
use remote_externalities::{ use remote_externalities::{
Builder, Mode, OfflineConfig, OnlineConfig, SnapshotConfig, TestExternalities, rpc_api::RpcService, Builder, Mode, OfflineConfig, OnlineConfig, SnapshotConfig,
TestExternalities,
}; };
use sc_chain_spec::ChainSpec; use sc_chain_spec::ChainSpec;
use sc_cli::{ use sc_cli::{
@@ -541,8 +542,8 @@ impl State {
impl TryRuntimeCmd { impl TryRuntimeCmd {
pub async fn run<Block, ExecDispatch>(&self, config: Configuration) -> sc_cli::Result<()> pub async fn run<Block, ExecDispatch>(&self, config: Configuration) -> sc_cli::Result<()>
where where
Block: BlockT<Hash = H256> + serde::de::DeserializeOwned, Block: BlockT<Hash = H256> + DeserializeOwned,
Block::Header: serde::de::DeserializeOwned, Block::Header: DeserializeOwned,
Block::Hash: FromStr, Block::Hash: FromStr,
<Block::Hash as FromStr>::Err: Debug, <Block::Hash as FromStr>::Err: Debug,
NumberFor<Block>: FromStr, NumberFor<Block>: FromStr,
@@ -626,13 +627,15 @@ where
/// ///
/// If the spec names don't match, if `relaxed`, then it emits a warning, else it panics. /// If the spec names don't match, if `relaxed`, then it emits a warning, else it panics.
/// If the spec versions don't match, it only ever emits a warning. /// If the spec versions don't match, it only ever emits a warning.
pub(crate) async fn ensure_matching_spec<Block: BlockT + serde::de::DeserializeOwned>( pub(crate) async fn ensure_matching_spec<Block: BlockT + DeserializeOwned>(
uri: String, uri: String,
expected_spec_name: String, expected_spec_name: String,
expected_spec_version: u32, expected_spec_version: u32,
relaxed: bool, relaxed: bool,
) { ) {
match remote_externalities::rpc_api::get_runtime_version::<Block, _>(uri.clone(), None) let rpc_service = RpcService::new(uri.clone(), false).await.unwrap();
match rpc_service
.get_runtime_version::<Block>(None)
.await .await
.map(|version| (String::from(version.spec_name.clone()), version.spec_version)) .map(|version| (String::from(version.spec_name.clone()), version.spec_version))
.map(|(spec_name, spec_version)| (spec_name.to_lowercase(), spec_version)) .map(|(spec_name, spec_version)| (spec_name.to_lowercase(), spec_version))