use std::{
collections::HashMap,
sync::{Arc, Mutex},
time::Duration,
};
use futures::StreamExt as _;
use linera_base::identifiers::{ApplicationId, ChainId};
use linera_client::chain_listener::ClientContext;
use linera_core::worker::Reason;
use linera_execution::{Query, QueryResponse};
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
use tracing::{debug, warn};
#[derive(Clone, Debug)]
pub struct RegisteredQuery {
pub name: String,
pub query: String,
}
pub fn parse_allowed_subscription(s: &str) -> anyhow::Result<RegisteredQuery> {
let trimmed = s.trim();
let rest = trimmed
.strip_prefix("query")
.ok_or_else(|| anyhow::anyhow!("expected query to start with 'query', got: {s}"))?;
anyhow::ensure!(
rest.starts_with(char::is_whitespace),
"expected whitespace after 'query' keyword"
);
let rest = rest.trim_start();
let name = rest
.split(|c: char| !c.is_alphanumeric() && c != '_')
.next()
.unwrap_or_default();
anyhow::ensure!(
!name.is_empty(),
"expected an operation name after 'query', e.g. 'query MyQuery {{ ... }}'"
);
Ok(RegisteredQuery {
name: name.to_string(),
query: trimmed.to_string(),
})
}
pub fn parse_subscription_ttl(s: &str) -> Result<(String, u64), String> {
let (name, secs) = s
.split_once('=')
.ok_or_else(|| format!("expected format Name=Secs, got: {s}"))?;
let secs: u64 = secs
.parse()
.map_err(|e| format!("invalid seconds value '{secs}': {e}"))?;
Ok((name.to_string(), secs))
}
#[derive(Clone, Debug, Hash, Eq, PartialEq)]
pub struct SubscriptionKey {
pub name: String,
pub chain_id: ChainId,
pub application_id: ApplicationId,
}
struct WatcherState {
sender: watch::Sender<Option<String>>,
}
pub struct QuerySubscriptionManager {
queries: HashMap<String, String>,
ttls: HashMap<String, Duration>,
watchers: Mutex<HashMap<SubscriptionKey, WatcherState>>,
}
impl QuerySubscriptionManager {
pub fn new(registered: Vec<RegisteredQuery>, ttls: HashMap<String, Duration>) -> Self {
let queries = registered
.into_iter()
.map(|rq| (rq.name, rq.query))
.collect();
Self {
queries,
ttls,
watchers: Mutex::new(HashMap::new()),
}
}
pub fn get_query(&self, name: &str) -> Option<&str> {
self.queries.get(name).map(|s| s.as_str())
}
pub fn subscribe<C: ClientContext + 'static>(
self: &Arc<Self>,
key: &SubscriptionKey,
context: Arc<futures::lock::Mutex<C>>,
token: CancellationToken,
) -> anyhow::Result<watch::Receiver<Option<String>>> {
let query_string = self
.get_query(&key.name)
.ok_or_else(|| {
anyhow::anyhow!("no subscription query registered with name '{}'", key.name)
})?
.to_string();
let mut watchers = self.watchers.lock().unwrap();
if let Some(state) = watchers.get(key) {
return Ok(state.sender.subscribe());
}
let (sender, receiver) = watch::channel(None);
watchers.insert(
key.clone(),
WatcherState {
sender: sender.clone(),
},
);
let ttl = self.ttls.get(&key.name).copied();
let manager = Arc::clone(self);
let key_clone = key.clone();
tokio::spawn(run_query_subscription_watcher(
context,
manager,
key_clone,
query_string,
sender,
token,
ttl,
));
Ok(receiver)
}
}
async fn run_query_subscription_watcher<C: ClientContext + 'static>(
context: Arc<futures::lock::Mutex<C>>,
manager: Arc<QuerySubscriptionManager>,
key: SubscriptionKey,
query_string: String,
sender: watch::Sender<Option<String>>,
token: CancellationToken,
ttl: Option<Duration>,
) {
debug!(
name = %key.name,
chain_id = %key.chain_id,
application_id = %key.application_id,
?ttl,
"starting query subscription watcher"
);
let notification_stream = {
let ctx = context.lock().await;
match ctx.make_chain_client(key.chain_id).await {
Ok(client) => match client.subscribe() {
Ok(stream) => stream,
Err(e) => {
warn!("failed to subscribe to chain notifications: {e}");
cleanup_watcher(&manager, &key);
return;
}
},
Err(e) => {
warn!("failed to create chain client: {e}");
cleanup_watcher(&manager, &key);
return;
}
}
};
let mut notification_stream = Box::pin(notification_stream);
let mut last_result: Option<String> = None;
execute_and_maybe_send(&context, &key, &query_string, &sender, &mut last_result).await;
let mut last_execution = tokio::time::Instant::now();
let mut pending_invalidation = false;
loop {
let ttl_sleep = if pending_invalidation {
if let Some(ttl) = ttl {
let elapsed = last_execution.elapsed();
if elapsed < ttl {
tokio::time::sleep(ttl - elapsed)
} else {
tokio::time::sleep(Duration::ZERO)
}
} else {
tokio::time::sleep(Duration::ZERO)
}
} else {
tokio::time::sleep(Duration::MAX)
};
tokio::pin!(ttl_sleep);
tokio::select! {
_ = token.cancelled() => {
debug!(name = %key.name, "watcher cancelled");
break;
}
() = &mut ttl_sleep, if pending_invalidation => {
pending_invalidation = false;
execute_and_maybe_send(
&context,
&key,
&query_string,
&sender,
&mut last_result,
)
.await;
last_execution = tokio::time::Instant::now();
}
notification = notification_stream.next() => {
match notification {
Some(n) => {
if matches!(n.reason, Reason::NewBlock { .. }) {
if ttl.is_some() {
if !pending_invalidation {
debug!(name = %key.name, "deferring invalidation until TTL expires");
}
pending_invalidation = true;
} else {
execute_and_maybe_send(
&context,
&key,
&query_string,
&sender,
&mut last_result,
)
.await;
last_execution = tokio::time::Instant::now();
}
}
}
None => {
debug!(name = %key.name, "notification stream ended");
break;
}
}
if sender.is_closed() {
debug!(name = %key.name, "no more subscribers, stopping watcher");
break;
}
}
}
}
cleanup_watcher(&manager, &key);
}
async fn execute_and_maybe_send<C: ClientContext + 'static>(
context: &Arc<futures::lock::Mutex<C>>,
key: &SubscriptionKey,
query_string: &str,
sender: &watch::Sender<Option<String>>,
last_result: &mut Option<String>,
) {
let json_request = serde_json::json!({ "query": query_string });
let request_bytes = serde_json::to_vec(&json_request).unwrap();
let query = Query::User {
application_id: key.application_id,
bytes: request_bytes,
};
let result = {
let ctx = context.lock().await;
match ctx.make_chain_client(key.chain_id).await {
Ok(client) => client.query_application(query, None).await,
Err(e) => {
warn!(name = %key.name, "failed to create chain client: {e}");
return;
}
}
};
match result {
Ok((outcome, _height)) => {
let response_bytes = match outcome.response {
QueryResponse::User(bytes) => bytes,
QueryResponse::System(_) => {
warn!(name = %key.name, "unexpected system response for user query");
return;
}
};
let json_string = match String::from_utf8(response_bytes) {
Ok(s) => s,
Err(e) => {
warn!(name = %key.name, "response bytes are not valid UTF-8: {e}");
return;
}
};
if last_result.as_ref() == Some(&json_string) {
return;
}
if let Err(e) = sender.send(Some(json_string.clone())) {
debug!(name = %key.name, "Failed to send graphql response: {e}");
}
*last_result = Some(json_string);
}
Err(e) => {
warn!(name = %key.name, "query execution failed: {e}");
}
}
}
fn cleanup_watcher(manager: &QuerySubscriptionManager, key: &SubscriptionKey) {
let mut watchers = manager.watchers.lock().unwrap();
watchers.remove(key);
debug!(name = %key.name, "watcher cleaned up");
}