1use 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
22pub 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#[derive(Debug, Clone)]
32pub struct AdsConfig {
33 pub max_concurrent_streams: usize,
35 pub response_buffer_size: usize,
37 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#[derive(Debug, Clone)]
56pub struct AdsService {
57 cache: Arc<ShardedCache>,
59 registry: Arc<ResourceRegistry>,
61 sotw_handler: Arc<SotwHandler>,
63 delta_handler: Arc<DeltaHandler>,
65 config: AdsConfig,
67}
68
69impl AdsService {
70 pub fn new(cache: Arc<ShardedCache>, registry: Arc<ResourceRegistry>) -> Self {
72 let sotw_handler = Arc::new(SotwHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
73 let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
74 Self {
75 cache,
76 registry,
77 sotw_handler,
78 delta_handler,
79 config: AdsConfig::default(),
80 }
81 }
82
83 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(®istry)));
90 let delta_handler = Arc::new(DeltaHandler::new(Arc::clone(&cache), Arc::clone(®istry)));
91 Self {
92 cache,
93 registry,
94 sotw_handler,
95 delta_handler,
96 config,
97 }
98 }
99
100 pub fn cache(&self) -> &ShardedCache {
102 &self.cache
103 }
104
105 pub fn registry(&self) -> &ResourceRegistry {
107 &self.registry
108 }
109
110 pub fn config(&self) -> &AdsConfig {
112 &self.config
113 }
114
115 pub fn into_service(self) -> AggregatedDiscoveryServiceServer<Self> {
119 AggregatedDiscoveryServiceServer::new(self)
120 }
121
122 #[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 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 } else if !response_nonce.is_empty() {
146 self.sotw_handler
148 .handle_ack(ctx, TypeUrl::new(type_url), version_info, response_nonce);
149 }
150
151 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 #[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
197pub type AdsResponseStream = ReceiverStream<Result<DiscoveryResponse, Status>>;
199
200pub 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 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 let error_detail = request
250 .error_detail
251 .as_ref()
252 .map(|e| e.message.as_str());
253
254 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 }
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 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 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 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}