use std::future::Future; use crate::errors::{ChainGatewayError, ChainGatewayOp}; use crate::primitives::{IsSyncing, QueryViewFunction}; use crate::types::ObservedState; use near_account_id::AccountId; use serde::{Serialize, de::DeserializeOwned}; use super::subscription::ContractMethodSubscription; /// Provides a subscribe-and-poll interface for observing contract state changes. /// Polls the view method every 200 ms and emits change notifications only when /// the returned bytes differ. /// /// # Example /// /// ``` /// use chain_gateway::mock::{MockChainState, Call}; /// use chain_gateway::state_viewer::{WatchContractState, SubscribeToContractMethod}; /// use chain_gateway::types::ObservedState; /// /// #[tokio::main] /// async fn main() { /// let viewer = MockChainState::builder() /// .with_syncing_status(Ok(false)) /// .with_query_view_function_response(Ok(ObservedState { /// observed_at: 1.into(), /// value: br#""hello""#.to_vec(), /// })) /// .build(); /// /// let mut stream = viewer /// .subscribe_to_contract_method::("contract.near".parse().unwrap(), "get_greeting") /// .await; /// /// let state = stream.latest().unwrap(); /// assert_eq!(state.value, "hello"); /// } /// ``` pub trait SubscribeToContractMethod { /// Subscribes to a contract view method and returns a stream of state updates. /// /// The returned stream polls the contract every 200 ms. /// /// # Type Parameter /// /// `T` is the deserialized return type of the contract method. fn subscribe_to_contract_method( &self, contract: AccountId, view_method: &str, ) -> impl Future + Send> + Send where T: DeserializeOwned + Send + Clone; } /// Performs a typed view call: serializes `args` as JSON, calls the /// contract and deserializes the response. /// /// # Example /// /// ``` /// use chain_gateway::mock::{MockChainState, Call}; /// use chain_gateway::state_viewer::ViewMethod; /// use chain_gateway::types::{NoArgs, ObservedState}; /// /// #[tokio::main] /// async fn main() { /// let viewer = MockChainState::builder() /// .with_syncing_status(Ok(false)) /// .with_query_view_function_response(Ok(ObservedState { /// observed_at: 1.into(), /// value: br#""hello""#.to_vec(), /// })) /// .build(); /// /// let result: ObservedState = viewer /// .view_method("contract.near".parse().unwrap(), "get_greeting", &NoArgs {}) /// .await /// .unwrap(); /// /// assert_eq!(result.value, "hello"); /// assert_eq!(result.observed_at, 1.into()); /// } /// ``` pub trait ViewMethod { fn view_method( &self, contract_id: AccountId, method_name: &str, args: &Arg, ) -> impl Future, ChainGatewayError>> + Send where Arg: Serialize + Sync, Res: DeserializeOwned + Send + Clone; } /// All other viewer traits are derived from this one pub(crate) trait ViewRaw: IsSyncing + QueryViewFunction { // waits until self is synced and then queries the view function fn view_raw( &self, contract_id: &AccountId, method_name: &str, args: &[u8], ) -> impl Future> + Send; } /// A watch-like stream of contract state changes. /// /// Call [`latest()`](WatchContractState::latest) to get the most recent value, /// and [`changed()`](WatchContractState::changed) to wait for the next update. /// Only actual value changes (different bytes) trigger a notification (block /// height increases alone do not). pub trait WatchContractState { /// Returns the last value observed on chain and the block height at which it was first /// observed. fn latest(&mut self) -> Result, ChainGatewayError>; /// Waits until the observed value changes. fn changed(&mut self) -> impl Future> + Send; } impl ViewRaw for T { async fn view_raw( &self, contract_id: &AccountId, method_name: &str, args: &[u8], ) -> Result { self.wait_for_full_sync().await; self.query_view_function(contract_id, method_name, args) .await .map_err(|err| ChainGatewayError::ViewError { op: ChainGatewayOp::ViewQuery { account_id: contract_id.to_string(), method_name: method_name.to_string(), }, message: err.to_string(), }) } } impl SubscribeToContractMethod for V { fn subscribe_to_contract_method( &self, contract: AccountId, view_method: &str, ) -> impl Future + Send> + Send where T: DeserializeOwned + Send + Clone, { ContractMethodSubscription::new(self.clone(), contract, view_method, b"{}".to_vec()) } } impl ViewMethod for T { async fn view_method( &self, contract_id: AccountId, method_name: &str, args: &Arg, ) -> Result, ChainGatewayError> where Arg: Serialize + Sync, Res: DeserializeOwned + Send + Clone, { let args: Vec = serde_json::to_vec(args).map_err(|err| ChainGatewayError::Serialization { op: ChainGatewayOp::ViewQuery { account_id: contract_id.to_string(), method_name: method_name.to_string(), }, message: err.to_string(), })?; let res = self.view_raw(&contract_id, method_name, &args).await?; let value = serde_json::from_slice::(&res.value).map_err(|err| { ChainGatewayError::Deserialization { message: err.to_string(), } })?; Ok(ObservedState { observed_at: res.observed_at, value, }) } } #[cfg(test)] mod tests { use super::ViewRaw; use crate::errors::{ChainGatewayError, ChainGatewayOp}; use crate::mock::{Call, MockChainState, MockError}; use crate::state_viewer::{SubscribeToContractMethod, ViewMethod, WatchContractState}; use crate::types::{NoArgs, ObservedState}; use assert_matches::assert_matches; use near_account_id::AccountId; use rand::distributions::{Alphanumeric, DistString}; use rand::rngs::StdRng; use rand::{Rng, SeedableRng}; /// Produces a deterministic `(Call, ObservedState)` pair from the given RNG /// so every test uses unique but reproducible data. fn random_view_params(rng: &mut StdRng) -> (Call, ObservedState) { let contract_id: AccountId = format!( "{}.testnet", Alphanumeric.sample_string(rng, 8).to_lowercase() ) .parse() .unwrap(); let method_name = Alphanumeric.sample_string(rng, 10); let args: Vec = (0..rng.gen_range(1..16)).map(|_| rng.r#gen()).collect(); let block_height: u64 = rng.gen_range(1..1_000_000); let payload: Vec = (0..rng.gen_range(1..32)).map(|_| rng.r#gen()).collect(); ( Call { contract_id, method_name, args, }, ObservedState { observed_at: block_height.into(), value: payload, }, ) } #[tokio::test] async fn test_view_raw_returns_ok_on_success() { let mut rng = StdRng::seed_from_u64(1); let (call, response) = random_view_params(&mut rng); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(response.clone())) .build(); let state = viewer .view_raw(&call.contract_id, &call.method_name, &call.args) .await .unwrap(); assert_eq!(state.observed_at, response.observed_at); assert_eq!(state.value, response.value); } #[tokio::test] async fn test_view_raw_queries_correct_arguments() { let mut rng = StdRng::seed_from_u64(2); let (call, response) = random_view_params(&mut rng); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(response)) .build(); viewer .view_raw(&call.contract_id, &call.method_name, &call.args) .await .unwrap(); assert_eq!(viewer.view_calls().await, vec![call]); } #[tokio::test] async fn test_view_raw_wraps_error_in_view_client() { let mut rng = StdRng::seed_from_u64(3); let (call, _response) = random_view_params(&mut rng); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Err(MockError::SyncError)) .build(); let err = viewer .view_raw(&call.contract_id, &call.method_name, b"{}") .await .unwrap_err(); assert_eq!( err, ChainGatewayError::ViewError { op: ChainGatewayOp::ViewQuery { account_id: call.contract_id.to_string(), method_name: call.method_name, }, message: MockError::SyncError.to_string(), } ); } #[tokio::test(start_paused = true)] async fn test_view_raw_blocks_until_synced() { let mut rng = StdRng::seed_from_u64(4); let (call, response) = random_view_params(&mut rng); let viewer = MockChainState::builder() .with_syncing_status(Ok(true)) .with_query_view_function_response(Ok(response)) .build(); let v = viewer.clone(); let cid = call.contract_id.clone(); let mn = call.method_name.clone(); let a = call.args.clone(); let handle = tokio::spawn(async move { v.view_raw(&cid, &mn, &a).await }); // wait_for_full_sync polls every 500ms; advance past one interval tokio::time::sleep(std::time::Duration::from_millis(1000)).await; assert!(!handle.is_finished(), "should block while syncing"); viewer.set_sync_response(Ok(false)); // wait_for_full_sync polls every 500ms; advance past one interval tokio::time::sleep(std::time::Duration::from_millis(600)).await; handle.await.unwrap().unwrap(); } #[tokio::test] async fn test_view_method_deserializes_response() { let mut rng = StdRng::seed_from_u64(5); let block_height: u64 = rng.gen_range(1..1_000_000); let value = Alphanumeric.sample_string(&mut rng, 12); let json_bytes = serde_json::to_vec(&value).unwrap(); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(ObservedState { observed_at: block_height.into(), value: json_bytes, })) .build(); let result = viewer .view_method::("a.testnet".parse().unwrap(), "m", &NoArgs {}) .await .unwrap(); assert_eq!(result.value, value); assert_eq!(result.observed_at, block_height.into()); } #[tokio::test] async fn test_view_method_propagates_view_error() { let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Err(MockError::ViewClientError)) .build(); let account_id: AccountId = "a.testnet".parse().unwrap(); let method_name = "m".to_string(); let err = viewer .view_method::(account_id.clone(), &method_name, &NoArgs {}) .await .unwrap_err(); assert_eq!( err, ChainGatewayError::ViewError { op: ChainGatewayOp::ViewQuery { account_id: account_id.to_string(), method_name }, message: MockError::ViewClientError.to_string() } ); } #[tokio::test] async fn test_view_returns_deserialization_error_on_bad_bytes() { let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(ObservedState { observed_at: 1.into(), value: b"not valid json".to_vec(), })) .build(); let contract_id: AccountId = "a.testnet".parse().unwrap(); let method_name: String = "m".into(); let err = viewer .view_method::(contract_id, &method_name, &NoArgs {}) .await .unwrap_err(); assert_matches!(err, ChainGatewayError::Deserialization { .. }); } #[tokio::test(start_paused = true)] async fn test_subscribe_latest_returns_initial_value() { let mut rng = StdRng::seed_from_u64(8); let block_height: u64 = rng.gen_range(1..1_000_000); let value = Alphanumeric.sample_string(&mut rng, 12); let json_bytes = serde_json::to_vec(&value).unwrap(); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(ObservedState { observed_at: block_height.into(), value: json_bytes, })) .build(); let mut sub = viewer .subscribe_to_contract_method::("a.testnet".parse().unwrap(), "m") .await; let state = sub.latest().unwrap(); assert_eq!(state.value, value); assert_eq!(state.observed_at, block_height.into()); } #[tokio::test(start_paused = true)] async fn test_subscribe_latest_returns_deserialization_error() { let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(ObservedState { observed_at: 1.into(), value: b"not json".to_vec(), })) .build(); let contract_id: AccountId = "a.testnet".parse().unwrap(); let method_name: String = "m".into(); let err = { let mut sub = viewer .subscribe_to_contract_method::(contract_id.clone(), &method_name) .await; sub.latest().unwrap_err() }; assert_matches!(err, ChainGatewayError::Deserialization { .. }); } #[tokio::test(start_paused = true)] async fn test_subscribe_changed_fires_on_value_change() { let mut rng = StdRng::seed_from_u64(10); let initial = Alphanumeric.sample_string(&mut rng, 10); let updated = Alphanumeric.sample_string(&mut rng, 10); let viewer = MockChainState::builder() .with_syncing_status(Ok(false)) .with_query_view_function_response(Ok(ObservedState { observed_at: 1.into(), value: serde_json::to_vec(&initial).unwrap(), })) .build(); let mut sub = viewer .subscribe_to_contract_method::("a.testnet".parse().unwrap(), "m") .await; assert_eq!(sub.latest().unwrap().value, initial); viewer .set_view_response(Ok(ObservedState { observed_at: 2.into(), value: serde_json::to_vec(&updated).unwrap(), })) .await; // Wait for change to be propagated tokio::time::timeout(std::time::Duration::from_secs(2), sub.changed()) .await .expect("changed() should resolve") .unwrap(); assert_eq!(sub.latest().unwrap().value, updated); } }