use std::sync::Arc;
use async_trait::async_trait;
use dynamo_kv_router::{
indexer::{LocalKvIndexer, WorkerKvQueryRequest, WorkerKvQueryResponse},
protocols::DpRank,
};
use dynamo_runtime::{
component::{Component, StartedEndpoint},
pipeline::{
AsyncEngine, AsyncEngineContextProvider, ManyOut, ResponseStream, SingleIn,
network::Ingress,
},
stream,
traits::DistributedRuntimeProvider,
};
use tokio::sync::Semaphore;
pub(crate) async fn start_worker_kv_query_endpoint(
component: Component,
publisher_id: u64,
worker_id: u64,
dp_rank: DpRank,
local_indexer: Arc<LocalKvIndexer>,
) -> anyhow::Result<StartedEndpoint> {
let engine = Arc::new(WorkerKvQueryEngine {
worker_id,
dp_rank,
local_indexer,
processing_semaphore: Semaphore::new(1),
});
let ingress = Ingress::for_engine(engine)?;
let route_worker_id = component.drt().connection_id();
let endpoint_name = format!("worker_kv_query_source_{publisher_id:x}");
tracing::info!(
"WorkerKvQuery endpoint starting for worker {worker_id} dp_rank {dp_rank} \
routed by instance {route_worker_id} on endpoint '{endpoint_name}'"
);
component
.endpoint(&endpoint_name)
.endpoint_builder()
.handler(ingress)
.graceful_shutdown(true)
.start_with_registration()
.await
}
pub(super) struct WorkerKvQueryEngine {
pub(super) worker_id: u64,
pub(super) dp_rank: DpRank,
pub(super) local_indexer: Arc<LocalKvIndexer>,
pub(super) processing_semaphore: Semaphore,
}
#[async_trait]
impl AsyncEngine<SingleIn<WorkerKvQueryRequest>, ManyOut<WorkerKvQueryResponse>, anyhow::Error>
for WorkerKvQueryEngine
{
async fn generate(
&self,
request: SingleIn<WorkerKvQueryRequest>,
) -> anyhow::Result<ManyOut<WorkerKvQueryResponse>> {
let (request, ctx) = request.into_parts();
tracing::debug!(
"Received query request for worker {}: {:?}",
self.worker_id,
request
);
if request.worker_id != self.worker_id {
let error_message = format!(
"WorkerKvQueryEngine::generate worker_id mismatch: request.worker_id={} this.worker_id={}",
request.worker_id, self.worker_id
);
let response = WorkerKvQueryResponse::Error(error_message);
return Ok(ResponseStream::new(
Box::pin(stream::iter(vec![response])),
ctx.context(),
));
}
if request.dp_rank != self.dp_rank {
let error_message = format!(
"WorkerKvQueryEngine::generate dp_rank mismatch: request.dp_rank={} this.dp_rank={}",
request.dp_rank, self.dp_rank
);
let response = WorkerKvQueryResponse::Error(error_message);
return Ok(ResponseStream::new(
Box::pin(stream::iter(vec![response])),
ctx.context(),
));
}
let likely_buffer_read = self
.local_indexer
.likely_served_from_buffer(request.start_event_id);
let _maybe_permit = if !likely_buffer_read {
let engine_ctx = ctx.context();
let permit = tokio::select! {
result = self.processing_semaphore.acquire() => {
result.map_err(|_| anyhow::anyhow!("Worker KV query semaphore closed"))?
}
_ = futures::future::select(engine_ctx.stopped(), engine_ctx.killed()) => {
tracing::warn!("Worker<>Router KV query request cancelled while waiting for semaphore");
return Ok(ResponseStream::new(
Box::pin(stream::iter(vec![WorkerKvQueryResponse::Error(
"Request cancelled by client".to_string(),
)])),
ctx.context(),
));
}
};
Some(permit)
} else {
None
};
let _slow_query_guard = if !likely_buffer_read {
Some(SlowQueryGuard::spawn(self.worker_id))
} else {
None
};
let response = self
.local_indexer
.get_events_in_id_range(request.start_event_id, request.end_event_id)
.await;
let response = negotiate_tree_dump_failure(response, request.supports_tree_dump_failed);
Ok(ResponseStream::new(
Box::pin(stream::iter(vec![response])),
ctx.context(),
))
}
}
fn negotiate_tree_dump_failure(
response: WorkerKvQueryResponse,
supports_tree_dump_failed: bool,
) -> WorkerKvQueryResponse {
match response {
WorkerKvQueryResponse::TreeDumpFailed { last_event_id, .. }
if !supports_tree_dump_failed =>
{
WorkerKvQueryResponse::TreeDump {
events: Vec::new(),
last_event_id,
}
}
response => response,
}
}
struct SlowQueryGuard(tokio::task::JoinHandle<()>);
impl SlowQueryGuard {
fn spawn(worker_id: u64) -> Self {
Self(tokio::spawn(async move {
let mut elapsed_secs = 0u64;
loop {
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
elapsed_secs += 5;
tracing::warn!(
worker_id,
elapsed_secs,
"Worker KV query still running - possible slow tree dump",
);
}
}))
}
}
impl Drop for SlowQueryGuard {
fn drop(&mut self) {
self.0.abort();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn explicit_dump_failure_is_capability_negotiated() {
let failure = || WorkerKvQueryResponse::TreeDumpFailed {
last_event_id: 17,
message: "offline".to_string(),
};
assert!(matches!(
negotiate_tree_dump_failure(failure(), true),
WorkerKvQueryResponse::TreeDumpFailed {
last_event_id: 17,
..
}
));
assert!(matches!(
negotiate_tree_dump_failure(failure(), false),
WorkerKvQueryResponse::TreeDump {
events,
last_event_id: 17,
} if events.is_empty()
));
}
}