Skip to main content

rust_zero_core/
kubernetes.rs

1//! Kubernetes EndpointSlice service discovery.
2//!
3//! ```no_run
4//! # async fn run() -> Result<(), Box<dyn std::error::Error>> {
5//! use rust_zero_core::{KubernetesDiscovery, KubernetesDiscoveryConfig};
6//! let discovery = KubernetesDiscovery::infer(
7//!     KubernetesDiscoveryConfig::new("production").with_port_name("grpc"),
8//! ).await?;
9//! let subscription = discovery.subscribe("users").await?;
10//! println!("{:?}", subscription.endpoints());
11//! # Ok(())
12//! # }
13//! ```
14
15use 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/// Namespace, port, and URI settings for Kubernetes EndpointSlice discovery.
31#[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/// A Kubernetes client that discovers ready addresses from EndpointSlices.
119#[derive(Clone)]
120pub struct KubernetesDiscovery {
121    client: Client,
122    config: Arc<KubernetesDiscoveryConfig>,
123}
124
125impl KubernetesDiscovery {
126    /// Creates a client from the local kubeconfig or in-cluster service account.
127    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    /// Starts a self-healing EndpointSlice watch and returns its first complete snapshot.
145    ///
146    /// The caller needs `list` and `watch` access to `discovery.k8s.io/v1` EndpointSlices in the
147    /// configured namespace. Only ready, non-terminating endpoints are returned.
148    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
233/// A live, complete endpoint snapshot for one Kubernetes Service.
234pub 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}