use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status, Streaming};
use tracing::{debug, error, info, instrument, warn};
use xds_cache::ShardedCache;
use xds_core::{NodeHash, ResourceRegistry, TypeUrl};
use crate::delta::{delta_response_to_proto, ClientResourceState, DeltaHandler};
use crate::sotw::{SotwHandler, SotwResponse};
use crate::stream::StreamContext;
pub use xds_types::envoy::service::discovery::v3::{
DeltaDiscoveryRequest, DeltaDiscoveryResponse, DiscoveryRequest, DiscoveryResponse,
};
pub use xds_types::envoy::service::discovery::v3::aggregated_discovery_service_server::{
AggregatedDiscoveryService, AggregatedDiscoveryServiceServer,
};
#[derive(Debug, Clone)]
pub struct AdsConfig {
pub max_concurrent_streams: usize,
pub response_buffer_size: usize,
pub enable_delta: bool,
}
impl Default for AdsConfig {
fn default() -> Self {
Self {
max_concurrent_streams: 100,
response_buffer_size: 16,
enable_delta: true,
}
}
}
#[derive(Debug, Clone)]
pub struct AdsService {
cache: Arc<ShardedCache>,
registry: Arc<ResourceRegistry>,
sotw_handler: Arc<SotwHandler>,
delta_handler: Arc<DeltaHandler>,
config: AdsConfig,
}
impl AdsService {
pub fn new(cache: Arc<ShardedCache>, registry: Arc<ResourceRegistry>) -> Self {
let sotw_handler = Arc::new(SotwHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
Self {
cache,
registry,
sotw_handler,
delta_handler,
config: AdsConfig::default(),
}
}
pub fn with_config(
cache: Arc<ShardedCache>,
registry: Arc<ResourceRegistry>,
config: AdsConfig,
) -> Self {
let sotw_handler = Arc::new(SotwHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
Self {
cache,
registry,
sotw_handler,
delta_handler,
config,
}
}
pub fn cache(&self) -> &ShardedCache {
&self.cache
}
pub fn registry(&self) -> &ResourceRegistry {
&self.registry
}
pub fn config(&self) -> &AdsConfig {
&self.config
}
pub fn into_service(self) -> AggregatedDiscoveryServiceServer<Self> {
AggregatedDiscoveryServiceServer::new(self)
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip(self, ctx), fields(stream = %ctx.id()))]
pub fn process_sotw_request(
&self,
ctx: &StreamContext,
type_url: &str,
version_info: &str,
resource_names: &[String],
node_hash: NodeHash,
response_nonce: &str,
error_detail: Option<&str>,
) -> Result<Option<DiscoveryResponse>, Status> {
if let Some(error) = error_detail {
self.sotw_handler.handle_nack(
ctx,
TypeUrl::new(type_url),
version_info,
response_nonce,
error,
);
} else if !response_nonce.is_empty() {
self.sotw_handler
.handle_ack(ctx, TypeUrl::new(type_url), version_info, response_nonce);
}
let result = self
.sotw_handler
.process_request(
ctx,
TypeUrl::new(type_url),
version_info,
resource_names,
node_hash,
)
.map_err(|e| Status::internal(format!("Failed to process request: {}", e)))?;
match result {
Some(response) => Ok(Some(self.convert_sotw_response(response)?)),
None => Ok(None),
}
}
#[allow(clippy::result_large_err)]
fn convert_sotw_response(&self, response: SotwResponse) -> Result<DiscoveryResponse, Status> {
use xds_types::google::protobuf::Any;
let resources: Vec<Any> = response
.resources
.iter()
.filter_map(|r| {
r.encode().ok().map(|encoded| Any {
type_url: encoded.type_url.clone(),
value: encoded.value.clone(),
})
})
.collect();
Ok(DiscoveryResponse {
version_info: response.version_info,
resources,
type_url: response.type_url.to_string(),
nonce: response.nonce,
canary: false,
control_plane: None,
resource_errors: vec![],
})
}
}
pub type AdsResponseStream = ReceiverStream<Result<DiscoveryResponse, Status>>;
pub type AdsDeltaResponseStream = ReceiverStream<Result<DeltaDiscoveryResponse, Status>>;
#[async_trait]
impl AggregatedDiscoveryService for AdsService {
type StreamAggregatedResourcesStream = AdsResponseStream;
#[instrument(skip(self, request), name = "ads_stream")]
async fn stream_aggregated_resources(
&self,
request: Request<Streaming<DiscoveryRequest>>,
) -> Result<Response<Self::StreamAggregatedResourcesStream>, Status> {
let mut stream = request.into_inner();
let (tx, rx) = mpsc::channel(self.config.response_buffer_size);
let service = self.clone();
let mut ctx = StreamContext::new();
info!(stream = %ctx.id(), "ADS stream started");
tokio::spawn(async move {
let mut node_hash: Option<NodeHash> = None;
while let Some(result) = tokio_stream::StreamExt::next(&mut stream).await {
match result {
Ok(request) => {
if node_hash.is_none() {
if let Some(ref node) = request.node {
let hash = NodeHash::from_id(&node.id);
ctx.set_node(node.id.clone(), hash);
node_hash = Some(hash);
debug!(
stream = %ctx.id(),
node_id = %node.id,
"node identified"
);
}
}
let hash = match node_hash {
Some(h) => h,
None => {
warn!(stream = %ctx.id(), "request without node info");
continue;
}
};
let error_detail = request
.error_detail
.as_ref()
.map(|e| e.message.as_str());
match service.process_sotw_request(
&ctx,
&request.type_url,
&request.version_info,
&request.resource_names,
hash,
&request.response_nonce,
error_detail,
) {
Ok(Some(response)) => {
if tx.send(Ok(response)).await.is_err() {
debug!(stream = %ctx.id(), "client disconnected");
break;
}
}
Ok(None) => {
}
Err(e) => {
error!(stream = %ctx.id(), error = %e, "request processing failed");
let _ = tx.send(Err(e)).await;
break;
}
}
}
Err(e) => {
error!(stream = %ctx.id(), error = %e, "stream error");
break;
}
}
}
info!(
stream = %ctx.id(),
duration = ?ctx.duration(),
requests = ctx.request_count(),
responses = ctx.response_count(),
"ADS stream ended"
);
});
Ok(Response::new(ReceiverStream::new(rx)))
}
type DeltaAggregatedResourcesStream = AdsDeltaResponseStream;
#[instrument(skip(self, request), name = "ads_delta_stream")]
async fn delta_aggregated_resources(
&self,
request: Request<Streaming<DeltaDiscoveryRequest>>,
) -> Result<Response<Self::DeltaAggregatedResourcesStream>, Status> {
let mut stream = request.into_inner();
let (tx, rx) = mpsc::channel(self.config.response_buffer_size);
let service = self.clone();
let mut ctx = StreamContext::new();
info!(stream = %ctx.id(), "Delta ADS stream started");
tokio::spawn(async move {
let mut node_hash: Option<NodeHash> = None;
let mut client_states: HashMap<String, ClientResourceState> = HashMap::new();
while let Some(result) = tokio_stream::StreamExt::next(&mut stream).await {
match result {
Ok(request) => {
if request.type_url.is_empty() {
warn!(
stream = %ctx.id(),
"delta ADS request missing type_url"
);
continue;
}
if node_hash.is_none() {
if let Some(ref node) = request.node {
let hash = NodeHash::from_id(&node.id);
ctx.set_node(node.id.clone(), hash);
node_hash = Some(hash);
}
}
let hash = match node_hash {
Some(h) => h,
None => {
error!(
stream = %ctx.id(),
"first delta ADS request missing required node information"
);
let _ = tx
.send(Err(Status::invalid_argument(
"first request must include node information",
)))
.await;
break;
}
};
let type_url = TypeUrl::new(request.type_url.clone());
if !request.response_nonce.is_empty() {
if let Some(ref err) = request.error_detail {
service.delta_handler.handle_nack(
&ctx,
type_url.clone(),
&request.response_nonce,
&err.message,
);
} else {
service.delta_handler.handle_ack(
&ctx,
type_url.clone(),
&request.response_nonce,
);
}
}
let client_state = client_states
.entry(request.type_url.clone())
.or_default();
match service.delta_handler.process_request(
&ctx,
type_url,
client_state,
request.resource_names_subscribe,
request.resource_names_unsubscribe,
hash,
) {
Ok(Some(response)) => match delta_response_to_proto(response) {
Ok(proto_response) => {
if tx.send(Ok(proto_response)).await.is_err() {
debug!(stream = %ctx.id(), "client disconnected");
break;
}
}
Err(e) => {
error!(stream = %ctx.id(), error = %e, "failed to encode delta ADS response");
let _ = tx.send(Err(e)).await;
break;
}
},
Ok(None) => {}
Err(e) => {
error!(stream = %ctx.id(), error = %e, "delta ADS request failed");
break;
}
}
}
Err(e) => {
error!(stream = %ctx.id(), error = %e, "delta stream error");
break;
}
}
}
info!(
stream = %ctx.id(),
duration = ?ctx.duration(),
requests = ctx.request_count(),
responses = ctx.response_count(),
"Delta ADS stream ended"
);
drop(tx);
});
Ok(Response::new(ReceiverStream::new(rx)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use xds_cache::{Cache, Snapshot};
fn setup() -> AdsService {
let cache = Arc::new(ShardedCache::new());
let registry = Arc::new(ResourceRegistry::new());
AdsService::new(cache, registry)
}
#[test]
fn ads_service_creation() {
let service = setup();
assert!(service.cache().snapshot_count() == 0);
}
#[test]
fn ads_service_with_config() {
let cache = Arc::new(ShardedCache::new());
let registry = Arc::new(ResourceRegistry::new());
let config = AdsConfig {
max_concurrent_streams: 50,
response_buffer_size: 8,
enable_delta: false,
};
let service = AdsService::with_config(cache, registry, config);
assert!(!service.config.enable_delta);
}
#[test]
fn process_request_no_snapshot() {
let service = setup();
let ctx = StreamContext::new();
let node_hash = NodeHash::from_id("unknown-node");
let result = service
.process_sotw_request(
&ctx,
"type.googleapis.com/test",
"",
&[],
node_hash,
"",
None,
)
.expect("process_sotw_request should not error");
assert!(result.is_none());
}
#[test]
fn process_request_with_snapshot() {
let service = setup();
let ctx = StreamContext::new();
let node_hash = NodeHash::from_id("test-node");
let snapshot = Snapshot::builder()
.version("v1")
.resources(TypeUrl::CLUSTER.into(), vec![])
.build();
service.cache().set_snapshot(node_hash, snapshot);
let result = service
.process_sotw_request(&ctx, TypeUrl::CLUSTER, "", &[], node_hash, "", None)
.expect("process_sotw_request should not error");
assert!(result.is_some());
let response = result.expect("response should be Some");
assert_eq!(response.version_info, "v1");
}
}