Skip to main content

xds_server/services/
ads.rs

1//! Aggregated Discovery Service (ADS) implementation.
2//!
3//! ADS multiplexes all xDS resource types over a single gRPC stream,
4//! ensuring consistent ordering of configuration updates.
5
6use std::collections::HashMap;
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use tokio::sync::mpsc;
11use tokio_stream::wrappers::ReceiverStream;
12use tonic::{Request, Response, Status, Streaming};
13use tracing::{debug, error, info, instrument, warn};
14
15use xds_cache::ShardedCache;
16use xds_core::{NodeHash, ResourceRegistry, TypeUrl};
17
18use crate::delta::{delta_response_to_proto, ClientResourceState, DeltaHandler};
19use crate::sotw::{SotwHandler, SotwResponse};
20use crate::stream::StreamContext;
21
22// Re-export the data-plane-api types for external use
23pub use xds_types::envoy::service::discovery::v3::{
24    DeltaDiscoveryRequest, DeltaDiscoveryResponse, DiscoveryRequest, DiscoveryResponse,
25};
26pub use xds_types::envoy::service::discovery::v3::aggregated_discovery_service_server::{
27    AggregatedDiscoveryService, AggregatedDiscoveryServiceServer,
28};
29
30/// Configuration for the ADS service.
31#[derive(Debug, Clone)]
32pub struct AdsConfig {
33    /// Maximum concurrent streams per connection.
34    pub max_concurrent_streams: usize,
35    /// Response buffer size per stream.
36    pub response_buffer_size: usize,
37    /// Enable delta protocol support.
38    pub enable_delta: bool,
39}
40
41impl Default for AdsConfig {
42    fn default() -> Self {
43        Self {
44            max_concurrent_streams: 100,
45            response_buffer_size: 16,
46            enable_delta: true,
47        }
48    }
49}
50
51/// Aggregated Discovery Service.
52///
53/// Implements the ADS gRPC service, multiplexing CDS, EDS, LDS, RDS, and SDS
54/// over a single bidirectional stream.
55#[derive(Debug, Clone)]
56pub struct AdsService {
57    /// Shared cache.
58    cache: Arc<ShardedCache>,
59    /// Resource registry.
60    registry: Arc<ResourceRegistry>,
61    /// SotW handler.
62    sotw_handler: Arc<SotwHandler>,
63    /// Delta handler.
64    delta_handler: Arc<DeltaHandler>,
65    /// Configuration.
66    config: AdsConfig,
67}
68
69impl AdsService {
70    /// Create a new ADS service.
71    pub fn new(cache: Arc<ShardedCache>, registry: Arc<ResourceRegistry>) -> Self {
72        let sotw_handler = Arc::new(SotwHandler::new(Arc::clone(&cache), Arc::clone(&registry)));
73        let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(&registry)));
74        Self {
75            cache,
76            registry,
77            sotw_handler,
78            delta_handler,
79            config: AdsConfig::default(),
80        }
81    }
82
83    /// Create with custom configuration.
84    pub fn with_config(
85        cache: Arc<ShardedCache>,
86        registry: Arc<ResourceRegistry>,
87        config: AdsConfig,
88    ) -> Self {
89        let sotw_handler = Arc::new(SotwHandler::new(Arc::clone(&cache), Arc::clone(&registry)));
90        let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(&registry)));
91        Self {
92            cache,
93            registry,
94            sotw_handler,
95            delta_handler,
96            config,
97        }
98    }
99
100    /// Get a reference to the cache.
101    pub fn cache(&self) -> &ShardedCache {
102        &self.cache
103    }
104
105    /// Get a reference to the registry.
106    pub fn registry(&self) -> &ResourceRegistry {
107        &self.registry
108    }
109
110    /// Get a reference to the configuration.
111    pub fn config(&self) -> &AdsConfig {
112        &self.config
113    }
114
115    /// Convert this service into a tonic service for use with Server::add_service.
116    ///
117    /// This creates a properly typed gRPC service using the data-plane-api generated server.
118    pub fn into_service(self) -> AggregatedDiscoveryServiceServer<Self> {
119        AggregatedDiscoveryServiceServer::new(self)
120    }
121
122    /// Process an incoming SotW discovery request.
123    #[allow(clippy::too_many_arguments)]
124    #[instrument(skip(self, ctx), fields(stream = %ctx.id()))]
125    pub fn process_sotw_request(
126        &self,
127        ctx: &StreamContext,
128        type_url: &str,
129        version_info: &str,
130        resource_names: &[String],
131        node_hash: NodeHash,
132        response_nonce: &str,
133        error_detail: Option<&str>,
134    ) -> Result<Option<DiscoveryResponse>, Status> {
135        // Check for NACK
136        if let Some(error) = error_detail {
137            self.sotw_handler.handle_nack(
138                ctx,
139                TypeUrl::new(type_url),
140                version_info,
141                response_nonce,
142                error,
143            );
144            // On NACK, we don't send a new response unless there's new data
145        } else if !response_nonce.is_empty() {
146            // ACK
147            self.sotw_handler
148                .handle_ack(ctx, TypeUrl::new(type_url), version_info, response_nonce);
149        }
150
151        // Process the request
152        let result = self
153            .sotw_handler
154            .process_request(
155                ctx,
156                TypeUrl::new(type_url),
157                version_info,
158                resource_names,
159                node_hash,
160            )
161            .map_err(|e| Status::internal(format!("Failed to process request: {}", e)))?;
162
163        match result {
164            Some(response) => Ok(Some(self.convert_sotw_response(response)?)),
165            None => Ok(None),
166        }
167    }
168
169    /// Convert internal SotW response to proto DiscoveryResponse.
170    #[allow(clippy::result_large_err)]
171    fn convert_sotw_response(&self, response: SotwResponse) -> Result<DiscoveryResponse, Status> {
172        use xds_types::google::protobuf::Any;
173
174        let resources: Vec<Any> = response
175            .resources
176            .iter()
177            .filter_map(|r| {
178                r.encode().ok().map(|encoded| Any {
179                    type_url: encoded.type_url.clone(),
180                    value: encoded.value.clone(),
181                })
182            })
183            .collect();
184
185        Ok(DiscoveryResponse {
186            version_info: response.version_info,
187            resources,
188            type_url: response.type_url.to_string(),
189            nonce: response.nonce,
190            canary: false,
191            control_plane: None,
192            resource_errors: vec![],
193        })
194    }
195}
196
197/// Response stream type for ADS.
198pub type AdsResponseStream = ReceiverStream<Result<DiscoveryResponse, Status>>;
199
200/// Delta response stream type for ADS.
201pub type AdsDeltaResponseStream = ReceiverStream<Result<DeltaDiscoveryResponse, Status>>;
202
203#[async_trait]
204impl AggregatedDiscoveryService for AdsService {
205    type StreamAggregatedResourcesStream = AdsResponseStream;
206
207    #[instrument(skip(self, request), name = "ads_stream")]
208    async fn stream_aggregated_resources(
209        &self,
210        request: Request<Streaming<DiscoveryRequest>>,
211    ) -> Result<Response<Self::StreamAggregatedResourcesStream>, Status> {
212        let mut stream = request.into_inner();
213        let (tx, rx) = mpsc::channel(self.config.response_buffer_size);
214
215        let service = self.clone();
216        let mut ctx = StreamContext::new();
217
218        info!(stream = %ctx.id(), "ADS stream started");
219
220        tokio::spawn(async move {
221            let mut node_hash: Option<NodeHash> = None;
222
223            while let Some(result) = tokio_stream::StreamExt::next(&mut stream).await {
224                match result {
225                    Ok(request) => {
226                        // Extract node info on first request
227                        if node_hash.is_none() {
228                            if let Some(ref node) = request.node {
229                                let hash = NodeHash::from_id(&node.id);
230                                ctx.set_node(node.id.clone(), hash);
231                                node_hash = Some(hash);
232                                debug!(
233                                    stream = %ctx.id(),
234                                    node_id = %node.id,
235                                    "node identified"
236                                );
237                            }
238                        }
239
240                        let hash = match node_hash {
241                            Some(h) => h,
242                            None => {
243                                warn!(stream = %ctx.id(), "request without node info");
244                                continue;
245                            }
246                        };
247
248                        // Extract error detail from the proto message
249                        let error_detail = request
250                            .error_detail
251                            .as_ref()
252                            .map(|e| e.message.as_str());
253
254                        // Process the request
255                        match service.process_sotw_request(
256                            &ctx,
257                            &request.type_url,
258                            &request.version_info,
259                            &request.resource_names,
260                            hash,
261                            &request.response_nonce,
262                            error_detail,
263                        ) {
264                            Ok(Some(response)) => {
265                                if tx.send(Ok(response)).await.is_err() {
266                                    debug!(stream = %ctx.id(), "client disconnected");
267                                    break;
268                                }
269                            }
270                            Ok(None) => {
271                                // No update needed
272                            }
273                            Err(e) => {
274                                error!(stream = %ctx.id(), error = %e, "request processing failed");
275                                let _ = tx.send(Err(e)).await;
276                                break;
277                            }
278                        }
279                    }
280                    Err(e) => {
281                        error!(stream = %ctx.id(), error = %e, "stream error");
282                        break;
283                    }
284                }
285            }
286
287            info!(
288                stream = %ctx.id(),
289                duration = ?ctx.duration(),
290                requests = ctx.request_count(),
291                responses = ctx.response_count(),
292                "ADS stream ended"
293            );
294        });
295
296        Ok(Response::new(ReceiverStream::new(rx)))
297    }
298
299    type DeltaAggregatedResourcesStream = AdsDeltaResponseStream;
300
301    #[instrument(skip(self, request), name = "ads_delta_stream")]
302    async fn delta_aggregated_resources(
303        &self,
304        request: Request<Streaming<DeltaDiscoveryRequest>>,
305    ) -> Result<Response<Self::DeltaAggregatedResourcesStream>, Status> {
306        let mut stream = request.into_inner();
307        let (tx, rx) = mpsc::channel(self.config.response_buffer_size);
308
309        let service = self.clone();
310        let mut ctx = StreamContext::new();
311        info!(stream = %ctx.id(), "Delta ADS stream started");
312
313        tokio::spawn(async move {
314            let mut node_hash: Option<NodeHash> = None;
315            // ADS multiplexes types over one stream — track client state per type URL.
316            let mut client_states: HashMap<String, ClientResourceState> = HashMap::new();
317
318            while let Some(result) = tokio_stream::StreamExt::next(&mut stream).await {
319                match result {
320                    Ok(request) => {
321                        if request.type_url.is_empty() {
322                            warn!(
323                                stream = %ctx.id(),
324                                "delta ADS request missing type_url"
325                            );
326                            continue;
327                        }
328
329                        if node_hash.is_none() {
330                            if let Some(ref node) = request.node {
331                                let hash = NodeHash::from_id(&node.id);
332                                ctx.set_node(node.id.clone(), hash);
333                                node_hash = Some(hash);
334                            }
335                        }
336
337                        let hash = match node_hash {
338                            Some(h) => h,
339                            None => {
340                                error!(
341                                    stream = %ctx.id(),
342                                    "first delta ADS request missing required node information"
343                                );
344                                let _ = tx
345                                    .send(Err(Status::invalid_argument(
346                                        "first request must include node information",
347                                    )))
348                                    .await;
349                                break;
350                            }
351                        };
352
353                        let type_url = TypeUrl::new(request.type_url.clone());
354
355                        if !request.response_nonce.is_empty() {
356                            if let Some(ref err) = request.error_detail {
357                                service.delta_handler.handle_nack(
358                                    &ctx,
359                                    type_url.clone(),
360                                    &request.response_nonce,
361                                    &err.message,
362                                );
363                            } else {
364                                service.delta_handler.handle_ack(
365                                    &ctx,
366                                    type_url.clone(),
367                                    &request.response_nonce,
368                                );
369                            }
370                        }
371
372                        let client_state = client_states
373                            .entry(request.type_url.clone())
374                            .or_default();
375
376                        match service.delta_handler.process_request(
377                            &ctx,
378                            type_url,
379                            client_state,
380                            request.resource_names_subscribe,
381                            request.resource_names_unsubscribe,
382                            hash,
383                        ) {
384                            Ok(Some(response)) => match delta_response_to_proto(response) {
385                                Ok(proto_response) => {
386                                    if tx.send(Ok(proto_response)).await.is_err() {
387                                        debug!(stream = %ctx.id(), "client disconnected");
388                                        break;
389                                    }
390                                }
391                                Err(e) => {
392                                    error!(stream = %ctx.id(), error = %e, "failed to encode delta ADS response");
393                                    let _ = tx.send(Err(e)).await;
394                                    break;
395                                }
396                            },
397                            Ok(None) => {}
398                            Err(e) => {
399                                error!(stream = %ctx.id(), error = %e, "delta ADS request failed");
400                                break;
401                            }
402                        }
403                    }
404                    Err(e) => {
405                        error!(stream = %ctx.id(), error = %e, "delta stream error");
406                        break;
407                    }
408                }
409            }
410
411            info!(
412                stream = %ctx.id(),
413                duration = ?ctx.duration(),
414                requests = ctx.request_count(),
415                responses = ctx.response_count(),
416                "Delta ADS stream ended"
417            );
418            drop(tx);
419        });
420
421        Ok(Response::new(ReceiverStream::new(rx)))
422    }
423}
424
425#[cfg(test)]
426mod tests {
427    use super::*;
428    use xds_cache::{Cache, Snapshot};
429
430    fn setup() -> AdsService {
431        let cache = Arc::new(ShardedCache::new());
432        let registry = Arc::new(ResourceRegistry::new());
433        AdsService::new(cache, registry)
434    }
435
436    #[test]
437    fn ads_service_creation() {
438        let service = setup();
439        assert!(service.cache().snapshot_count() == 0);
440    }
441
442    #[test]
443    fn ads_service_with_config() {
444        let cache = Arc::new(ShardedCache::new());
445        let registry = Arc::new(ResourceRegistry::new());
446        let config = AdsConfig {
447            max_concurrent_streams: 50,
448            response_buffer_size: 8,
449            enable_delta: false,
450        };
451
452        let service = AdsService::with_config(cache, registry, config);
453        assert!(!service.config.enable_delta);
454    }
455
456    #[test]
457    fn process_request_no_snapshot() {
458        let service = setup();
459        let ctx = StreamContext::new();
460        let node_hash = NodeHash::from_id("unknown-node");
461
462        let result = service
463            .process_sotw_request(
464                &ctx,
465                "type.googleapis.com/test",
466                "",
467                &[],
468                node_hash,
469                "",
470                None,
471            )
472            .expect("process_sotw_request should not error");
473
474        assert!(result.is_none());
475    }
476
477    #[test]
478    fn process_request_with_snapshot() {
479        let service = setup();
480        let ctx = StreamContext::new();
481        let node_hash = NodeHash::from_id("test-node");
482
483        // Add a snapshot
484        let snapshot = Snapshot::builder()
485            .version("v1")
486            .resources(TypeUrl::CLUSTER.into(), vec![])
487            .build();
488        service.cache().set_snapshot(node_hash, snapshot);
489
490        // Request should return a response (empty resources but valid)
491        let result = service
492            .process_sotw_request(&ctx, TypeUrl::CLUSTER, "", &[], node_hash, "", None)
493            .expect("process_sotw_request should not error");
494
495        assert!(result.is_some());
496        let response = result.expect("response should be Some");
497        assert_eq!(response.version_info, "v1");
498    }
499}