1use crate::{DiscoveryReconnectBackoff, EndpointChangeFuture, EndpointSubscription};
16use futures::StreamExt;
17use k8s_openapi::api::discovery::v1::EndpointSlice;
18use kube::{runtime::watcher, Api, Client, ResourceExt};
19use std::{
20 collections::{hash_map::DefaultHasher, BTreeMap, BTreeSet},
21 error::Error,
22 fmt,
23 hash::{Hash, Hasher},
24 net::IpAddr,
25 sync::Arc,
26 time::{Duration, SystemTime, UNIX_EPOCH},
27};
28use tokio::{sync::watch, task::JoinHandle};
29
30#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct KubernetesDiscoveryConfig {
33 pub namespace: String,
34 pub port_name: Option<String>,
35 pub scheme: String,
36 pub startup_timeout: Duration,
37 pub reconnect_backoff: DiscoveryReconnectBackoff,
38}
39
40impl KubernetesDiscoveryConfig {
41 pub fn new(namespace: impl Into<String>) -> Self {
42 Self {
43 namespace: namespace.into(),
44 port_name: None,
45 scheme: "http".to_owned(),
46 startup_timeout: Duration::from_secs(10),
47 reconnect_backoff: DiscoveryReconnectBackoff::default(),
48 }
49 }
50
51 pub fn with_port_name(mut self, port_name: impl Into<String>) -> Self {
52 self.port_name = Some(port_name.into());
53 self
54 }
55
56 pub fn with_scheme(mut self, scheme: impl Into<String>) -> Self {
57 self.scheme = scheme.into();
58 self
59 }
60
61 pub fn with_startup_timeout(mut self, timeout: Duration) -> Self {
62 assert!(
63 !timeout.is_zero(),
64 "Kubernetes discovery startup timeout must be positive"
65 );
66 self.startup_timeout = timeout;
67 self
68 }
69
70 pub fn with_reconnect_backoff(mut self, backoff: DiscoveryReconnectBackoff) -> Self {
71 self.reconnect_backoff = backoff;
72 self
73 }
74}
75
76#[derive(Debug)]
77pub enum KubernetesDiscoveryError {
78 InvalidNamespace,
79 InvalidService,
80 InvalidPortName,
81 InvalidScheme,
82 Client(kube::Error),
83 StartupTimeout,
84 WatchClosed,
85}
86
87impl fmt::Display for KubernetesDiscoveryError {
88 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
89 match self {
90 Self::InvalidNamespace => formatter.write_str("Kubernetes namespace cannot be empty"),
91 Self::InvalidService => formatter.write_str("Kubernetes service name cannot be empty"),
92 Self::InvalidPortName => formatter.write_str("Kubernetes port name cannot be empty"),
93 Self::InvalidScheme => formatter.write_str("Kubernetes endpoint scheme is invalid"),
94 Self::Client(error) => write!(formatter, "Kubernetes client failed: {error}"),
95 Self::StartupTimeout => {
96 formatter.write_str("Kubernetes endpoint discovery timed out during initial list")
97 }
98 Self::WatchClosed => formatter.write_str("Kubernetes endpoint watch closed"),
99 }
100 }
101}
102
103impl Error for KubernetesDiscoveryError {
104 fn source(&self) -> Option<&(dyn Error + 'static)> {
105 match self {
106 Self::Client(error) => Some(error),
107 _ => None,
108 }
109 }
110}
111
112impl From<kube::Error> for KubernetesDiscoveryError {
113 fn from(error: kube::Error) -> Self {
114 Self::Client(error)
115 }
116}
117
118#[derive(Clone)]
120pub struct KubernetesDiscovery {
121 client: Client,
122 config: Arc<KubernetesDiscoveryConfig>,
123}
124
125impl KubernetesDiscovery {
126 pub async fn infer(
128 config: KubernetesDiscoveryConfig,
129 ) -> Result<Self, KubernetesDiscoveryError> {
130 Self::new(Client::try_default().await?, config)
131 }
132
133 pub fn new(
134 client: Client,
135 config: KubernetesDiscoveryConfig,
136 ) -> Result<Self, KubernetesDiscoveryError> {
137 validate_config(&config)?;
138 Ok(Self {
139 client,
140 config: Arc::new(config),
141 })
142 }
143
144 pub async fn subscribe(
149 &self,
150 service: impl AsRef<str>,
151 ) -> Result<KubernetesServiceSubscription, KubernetesDiscoveryError> {
152 let service = service.as_ref().trim();
153 if !valid_dns_label(service) {
154 return Err(KubernetesDiscoveryError::InvalidService);
155 }
156
157 let api = Api::<EndpointSlice>::namespaced(self.client.clone(), &self.config.namespace);
158 let watcher_config =
159 watcher::Config::default().labels(&format!("kubernetes.io/service-name={service}"));
160 let port_name = self.config.port_name.clone();
161 let scheme: Arc<str> = Arc::from(self.config.scheme.clone());
162 let backoff = self.config.reconnect_backoff;
163 let reconnect_seed = reconnect_seed(service);
164 let (updates, mut receiver) = watch::channel(Vec::new());
165
166 let task = tokio::spawn(async move {
167 let mut slices = BTreeMap::<String, Vec<String>>::new();
168 let mut attempt = 0_u32;
169 loop {
170 let mut pending = BTreeMap::<String, Vec<String>>::new();
171 let mut stream = Box::pin(watcher(api.clone(), watcher_config.clone()));
172 let mut received_event = false;
173 while let Some(result) = stream.next().await {
174 let Ok(event) = result else {
175 break;
176 };
177 received_event = true;
178 match event {
179 watcher::Event::Apply(slice) => {
180 slices.insert(
181 slice.name_any(),
182 endpoints_from_slice(&slice, port_name.as_deref(), &scheme),
183 );
184 updates.send_replace(combined_endpoints(&slices));
185 }
186 watcher::Event::Delete(slice) => {
187 slices.remove(&slice.name_any());
188 updates.send_replace(combined_endpoints(&slices));
189 }
190 watcher::Event::Init => {
191 pending.clear();
192 }
193 watcher::Event::InitApply(slice) => {
194 pending.insert(
195 slice.name_any(),
196 endpoints_from_slice(&slice, port_name.as_deref(), &scheme),
197 );
198 }
199 watcher::Event::InitDone => {
200 slices = std::mem::take(&mut pending);
201 updates.send_replace(combined_endpoints(&slices));
202 }
203 }
204 }
205 if received_event {
206 attempt = 0;
207 }
208 tokio::time::sleep(
209 backoff.delay(attempt, reconnect_seed.wrapping_add(u64::from(attempt))),
210 )
211 .await;
212 attempt = attempt.saturating_add(1);
213 }
214 #[allow(unreachable_code)]
215 Ok(())
216 });
217
218 match tokio::time::timeout(self.config.startup_timeout, receiver.changed()).await {
219 Ok(Ok(())) => Ok(KubernetesServiceSubscription { receiver, task }),
220 Ok(Err(_)) => match task.await {
221 Ok(Err(error)) => Err(error),
222 Ok(Ok(())) => Err(KubernetesDiscoveryError::WatchClosed),
223 Err(_) => Err(KubernetesDiscoveryError::WatchClosed),
224 },
225 Err(_) => {
226 task.abort();
227 Err(KubernetesDiscoveryError::StartupTimeout)
228 }
229 }
230 }
231}
232
233pub struct KubernetesServiceSubscription {
235 receiver: watch::Receiver<Vec<String>>,
236 task: JoinHandle<Result<(), KubernetesDiscoveryError>>,
237}
238
239impl KubernetesServiceSubscription {
240 pub fn endpoints(&self) -> Vec<String> {
241 self.receiver.borrow().clone()
242 }
243
244 pub async fn changed(&mut self) -> Result<Vec<String>, KubernetesDiscoveryError> {
245 self.receiver
246 .changed()
247 .await
248 .map_err(|_| KubernetesDiscoveryError::WatchClosed)?;
249 Ok(self.endpoints())
250 }
251}
252
253impl EndpointSubscription for KubernetesServiceSubscription {
254 type Error = KubernetesDiscoveryError;
255
256 fn endpoints(&self) -> Vec<String> {
257 KubernetesServiceSubscription::endpoints(self)
258 }
259
260 fn changed(&mut self) -> EndpointChangeFuture<'_, Self::Error> {
261 Box::pin(KubernetesServiceSubscription::changed(self))
262 }
263}
264
265impl Drop for KubernetesServiceSubscription {
266 fn drop(&mut self) {
267 self.task.abort();
268 }
269}
270
271fn validate_config(config: &KubernetesDiscoveryConfig) -> Result<(), KubernetesDiscoveryError> {
272 if !valid_dns_label(config.namespace.trim()) {
273 return Err(KubernetesDiscoveryError::InvalidNamespace);
274 }
275 if config
276 .port_name
277 .as_ref()
278 .is_some_and(|name| name.trim().is_empty())
279 {
280 return Err(KubernetesDiscoveryError::InvalidPortName);
281 }
282 let mut characters = config.scheme.chars();
283 if !characters
284 .next()
285 .is_some_and(|value| value.is_ascii_alphabetic())
286 || !characters
287 .all(|value| value.is_ascii_alphanumeric() || matches!(value, '+' | '-' | '.'))
288 {
289 return Err(KubernetesDiscoveryError::InvalidScheme);
290 }
291 Ok(())
292}
293
294fn valid_dns_label(value: &str) -> bool {
295 value.len() <= 63
296 && value
297 .bytes()
298 .next()
299 .is_some_and(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit())
300 && value
301 .bytes()
302 .last()
303 .is_some_and(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit())
304 && value
305 .bytes()
306 .all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'-')
307}
308
309fn reconnect_seed(scope: &str) -> u64 {
310 let mut hasher = DefaultHasher::new();
311 scope.hash(&mut hasher);
312 SystemTime::now()
313 .duration_since(UNIX_EPOCH)
314 .unwrap_or_default()
315 .as_nanos()
316 .hash(&mut hasher);
317 hasher.finish()
318}
319
320fn endpoints_from_slice(
321 slice: &EndpointSlice,
322 port_name: Option<&str>,
323 scheme: &str,
324) -> Vec<String> {
325 let port = slice.ports.as_deref().and_then(|ports| {
326 ports.iter().find(|port| {
327 port.port.is_some()
328 && port
329 .protocol
330 .as_deref()
331 .is_none_or(|protocol| protocol == "TCP")
332 && port_name.is_none_or(|name| port.name.as_deref() == Some(name))
333 })
334 });
335 let Some(port) = port.and_then(|port| port.port) else {
336 return Vec::new();
337 };
338
339 let mut endpoints = BTreeSet::new();
340 for endpoint in slice.endpoints.as_deref().unwrap_or_default() {
341 let ready = endpoint
342 .conditions
343 .as_ref()
344 .and_then(|conditions| conditions.ready)
345 .unwrap_or(true);
346 let terminating = endpoint
347 .conditions
348 .as_ref()
349 .and_then(|conditions| conditions.terminating)
350 .unwrap_or(false);
351 if !ready || terminating {
352 continue;
353 }
354 for address in &endpoint.addresses {
355 let host = match address.parse::<IpAddr>() {
356 Ok(IpAddr::V6(_)) => format!("[{address}]"),
357 _ => address.clone(),
358 };
359 endpoints.insert(format!("{scheme}://{host}:{port}"));
360 }
361 }
362 endpoints.into_iter().collect()
363}
364
365fn combined_endpoints(slices: &BTreeMap<String, Vec<String>>) -> Vec<String> {
366 slices
367 .values()
368 .flatten()
369 .cloned()
370 .collect::<BTreeSet<_>>()
371 .into_iter()
372 .collect()
373}
374
375#[cfg(test)]
376mod tests {
377 use super::*;
378 use k8s_openapi::api::discovery::v1::{Endpoint, EndpointConditions, EndpointPort};
379
380 fn slice(
381 name: &str,
382 address: &str,
383 ready: Option<bool>,
384 terminating: Option<bool>,
385 ) -> EndpointSlice {
386 EndpointSlice {
387 metadata: kube::api::ObjectMeta {
388 name: Some(name.to_owned()),
389 ..Default::default()
390 },
391 address_type: "IPv4".to_owned(),
392 endpoints: Some(vec![Endpoint {
393 addresses: vec![address.to_owned()],
394 conditions: Some(EndpointConditions {
395 ready,
396 serving: None,
397 terminating,
398 }),
399 ..Default::default()
400 }]),
401 ports: Some(vec![EndpointPort {
402 name: Some("grpc".to_owned()),
403 port: Some(8080),
404 protocol: Some("TCP".to_owned()),
405 ..Default::default()
406 }]),
407 }
408 }
409
410 #[test]
411 fn validates_configuration() {
412 assert!(matches!(
413 validate_config(&KubernetesDiscoveryConfig::new(" ")),
414 Err(KubernetesDiscoveryError::InvalidNamespace)
415 ));
416 assert!(matches!(
417 validate_config(&KubernetesDiscoveryConfig::new("default").with_scheme("1http")),
418 Err(KubernetesDiscoveryError::InvalidScheme)
419 ));
420 assert!(!valid_dns_label("users=anything"));
421 assert!(valid_dns_label("users-api"));
422 }
423
424 #[test]
425 fn extracts_only_ready_non_terminating_endpoints() {
426 assert_eq!(
427 endpoints_from_slice(
428 &slice("a", "10.0.0.1", Some(true), None),
429 Some("grpc"),
430 "http"
431 ),
432 vec!["http://10.0.0.1:8080"]
433 );
434 assert!(endpoints_from_slice(
435 &slice("b", "10.0.0.2", Some(false), None),
436 Some("grpc"),
437 "http"
438 )
439 .is_empty());
440 assert!(endpoints_from_slice(
441 &slice("c", "10.0.0.3", Some(true), Some(true)),
442 Some("grpc"),
443 "http"
444 )
445 .is_empty());
446 }
447
448 #[test]
449 fn combines_slices_in_stable_deduplicated_order() {
450 let slices = BTreeMap::from([
451 (
452 "b".to_owned(),
453 vec!["http://b:80".to_owned(), "http://a:80".to_owned()],
454 ),
455 ("a".to_owned(), vec!["http://a:80".to_owned()]),
456 ]);
457 assert_eq!(
458 combined_endpoints(&slices),
459 vec!["http://a:80", "http://b:80"]
460 );
461 }
462
463 #[tokio::test]
464 async fn kubernetes_integration_lists_and_watches_a_service() {
465 let Ok(service) = std::env::var("RUST_ZERO_KUBERNETES_SERVICE") else {
466 return;
467 };
468 let namespace = std::env::var("RUST_ZERO_KUBERNETES_NAMESPACE")
469 .unwrap_or_else(|_| "default".to_owned());
470 let mut config = KubernetesDiscoveryConfig::new(&namespace);
471 if let Ok(port_name) = std::env::var("RUST_ZERO_KUBERNETES_PORT_NAME") {
472 config = config.with_port_name(port_name);
473 }
474 let discovery = KubernetesDiscovery::infer(config).await.unwrap();
475 let api = Api::<EndpointSlice>::namespaced(discovery.client.clone(), &namespace);
476 let mut subscription = discovery.subscribe(&service).await.unwrap();
477 assert!(!subscription.endpoints().is_empty());
478
479 let name = format!("rust-zero-watch-{}", std::process::id());
480 let mut added = slice(&name, "10.0.0.99", Some(true), None);
481 added.metadata.labels = Some(BTreeMap::from([(
482 "kubernetes.io/service-name".to_owned(),
483 service,
484 )]));
485 api.create(&kube::api::PostParams::default(), &added)
486 .await
487 .unwrap();
488 let endpoints = tokio::time::timeout(Duration::from_secs(10), subscription.changed())
489 .await
490 .unwrap()
491 .unwrap();
492 assert!(endpoints
493 .iter()
494 .any(|endpoint| endpoint.contains("10.0.0.99")));
495
496 api.delete(&name, &kube::api::DeleteParams::default())
497 .await
498 .unwrap();
499 let endpoints = tokio::time::timeout(Duration::from_secs(10), subscription.changed())
500 .await
501 .unwrap()
502 .unwrap();
503 assert!(!endpoints
504 .iter()
505 .any(|endpoint| endpoint.contains("10.0.0.99")));
506 }
507}