staircase 0.0.7

Kubernetes Step-based Operator
Documentation
//! Various utils, that are not necessarily required for a step-based
//! operator, but handy in some situations.

use k8s_openapi::apiextensions_apiserver::pkg::apis::apiextensions::v1::CustomResourceDefinition;
use kube::{
    api::{Api, ResourceExt},
    client::Client,
    core::params::PostParams,
};
use thiserror::Error;

/// Error returned by [`update_crd`].
#[derive(Error, Debug)]
pub enum UpdateCrdError {
    /// Kubernetes API request failed while creating or replacing the CRD.
    #[error("Kubernetes API request failed while updating CRD: `{0}`")]
    Kube(#[from] kube::Error),
}

#[expect(rustdoc::broken_intra_doc_links)]
/// Tries to update the existing crd with the provided `crd`.
/// Note: this is especially usefull, if you generate the `crd` via
/// [`kube::CustomResource`] derive macro.
/// Use the [`kube::core::crd::merge_crds`] function to generate the `crd`.
/// If no crd exists, it just creates it.
/// It returns weather a change was made.
pub async fn update_crd(crd: CustomResourceDefinition, client: &Client) -> Result<bool, UpdateCrdError> {
    let crds: Api<CustomResourceDefinition> = Api::all(client.clone());
    let crd_name = crd.name_unchecked();

    match crds.get_opt(&crd_name).await? {
        None => {
            #[cfg(feature = "controller_trace")]
            tracing::debug!(?crd_name, "no CRD detected, trying to create it");

            crds.create(&PostParams::default(), &crd).await?;

            #[cfg(feature = "controller_trace")]
            if let Ok(json) = serde_json::to_string(&crd) {
                tracing::info!(?crd_name, ?json, "deployed CRD");
            }

            Ok(true)
        },
        Some(old_crd) => {
            if old_crd.spec.versions != crd.spec.versions
                || old_crd.spec.scope != crd.spec.scope
                || old_crd.spec.group != crd.spec.group
                || old_crd.spec.names.kind != crd.spec.names.kind
                || old_crd.spec.names.plural != crd.spec.names.plural
                || old_crd.spec.names.singular != crd.spec.names.singular
            {
                #[cfg(feature = "controller_trace")]
                tracing::debug!(?crd_name, "detected and outdated crd. trying to patch it");

                let params = &PostParams::default();
                let mut new_crd = crd;
                new_crd.metadata.resource_version = old_crd.metadata.resource_version.clone();
                crds.replace(&crd_name, params, &new_crd).await?;

                #[cfg(feature = "controller_trace")]
                if let (Ok(old_crd), Ok(new_crd)) = (serde_json::to_string(&old_crd), serde_json::to_string(&new_crd)) {
                    tracing::info!(?crd_name, ?old_crd, ?new_crd, "updated CRD");
                }

                Ok(true)
            } else {
                #[cfg(feature = "controller_trace")]
                tracing::debug!(
                    ?crd_name,
                    "detected existing crd spec matches expected crd. No change necessary"
                );

                Ok(false)
            }
        },
    }
}

#[cfg(test)]
mod tests {
    use std::{
        collections::VecDeque,
        convert::Infallible,
        sync::{Arc, Mutex},
    };

    use http::{Method, StatusCode, Uri};
    use kube::{CustomResource, CustomResourceExt, client::Body};
    use schemars::JsonSchema;
    use serde::{Deserialize, Serialize};
    use serde_json::{Value, json};

    use super::*;

    #[derive(CustomResource, Deserialize, Serialize, Clone, Debug, JsonSchema)]
    #[kube(group = "tests.staircase.dev", version = "v1", kind = "CrdUpdateTest", namespaced)]
    struct CrdUpdateTestSpec {
        value: String,
    }

    #[derive(Clone, Debug)]
    struct CapturedRequest {
        method: Method,
        uri:    Uri,
        body:   Vec<u8>,
    }

    #[derive(Clone, Debug)]
    struct QueuedResponse {
        status: StatusCode,
        body:   Value,
    }

    fn not_found_response() -> QueuedResponse {
        QueuedResponse {
            status: StatusCode::NOT_FOUND,
            body:   json!({
                "kind": "Status",
                "apiVersion": "v1",
                "metadata": {},
                "status": "Failure",
                "message": "customresourcedefinitions.apiextensions.k8s.io \"crdupdatetests.tests.staircase.dev\" not found",
                "reason": "NotFound",
                "code": 404
            }),
        }
    }

    fn server_error_response() -> QueuedResponse {
        QueuedResponse {
            status: StatusCode::INTERNAL_SERVER_ERROR,
            body:   json!({
                "kind": "Status",
                "apiVersion": "v1",
                "metadata": {},
                "status": "Failure",
                "message": "mock server error",
                "reason": "InternalError",
                "code": 500
            }),
        }
    }

    fn json_response(body: impl Serialize) -> QueuedResponse {
        QueuedResponse {
            status: StatusCode::OK,
            body:   serde_json::to_value(body).unwrap(),
        }
    }

    fn client_with_responses(responses: Vec<QueuedResponse>) -> (Client, Arc<Mutex<Vec<CapturedRequest>>>) {
        let captured = Arc::new(Mutex::new(Vec::new()));
        let responses = Arc::new(Mutex::new(VecDeque::from(responses)));
        let service = tower::service_fn({
            let captured = captured.clone();
            let responses = responses.clone();
            move |request: http::Request<Body>| {
                let captured = captured.clone();
                let responses = responses.clone();
                async move {
                    let (parts, body) = request.into_parts();
                    let body = body.collect_bytes().await.unwrap().to_vec();
                    captured.lock().unwrap().push(CapturedRequest {
                        method: parts.method,
                        uri: parts.uri,
                        body,
                    });

                    let response = responses.lock().unwrap().pop_front().expect("unexpected request");
                    let body = serde_json::to_vec(&response.body).unwrap();
                    Ok::<_, Infallible>(
                        http::Response::builder()
                            .status(response.status)
                            .header("content-type", "application/json")
                            .body(Body::from(body))
                            .unwrap(),
                    )
                }
            }
        });

        (Client::new(service, "default"), captured)
    }

    fn request_path(request: &CapturedRequest) -> &str { request.uri.path() }

    fn request_json(request: &CapturedRequest) -> Value { serde_json::from_slice(&request.body).unwrap() }

    fn crd_path(crd: &CustomResourceDefinition) -> String {
        format!(
            "/apis/apiextensions.k8s.io/v1/customresourcedefinitions/{}",
            crd.name_unchecked()
        )
    }

    #[test]
    fn creates_crd_when_it_does_not_exist() {
        let rt = tokio::runtime::Builder::new_current_thread().build().unwrap();
        let changed = rt.block_on(async {
            let crd = CrdUpdateTest::crd();
            let path = crd_path(&crd);
            let (client, captured) = client_with_responses(vec![not_found_response(), json_response(&crd)]);

            let changed = update_crd(crd.clone(), &client).await.unwrap();
            let captured = captured.lock().unwrap();

            assert_eq!(captured.len(), 2);
            assert_eq!(captured[0].method, Method::GET);
            assert_eq!(request_path(&captured[0]), path);
            assert_eq!(captured[1].method, Method::POST);
            assert_eq!(
                request_path(&captured[1]),
                "/apis/apiextensions.k8s.io/v1/customresourcedefinitions"
            );
            assert_eq!(
                request_json(&captured[1])["spec"],
                serde_json::to_value(&crd.spec).unwrap()
            );

            changed
        });

        assert!(changed);
    }

    #[test]
    fn returns_false_without_update_when_existing_crd_matches() {
        let rt = tokio::runtime::Builder::new_current_thread().build().unwrap();
        let changed = rt.block_on(async {
            let crd = CrdUpdateTest::crd();
            let path = crd_path(&crd);
            let (client, captured) = client_with_responses(vec![json_response(&crd)]);

            let changed = update_crd(crd, &client).await.unwrap();
            let captured = captured.lock().unwrap();

            assert_eq!(captured.len(), 1);
            assert_eq!(captured[0].method, Method::GET);
            assert_eq!(request_path(&captured[0]), path);

            changed
        });

        assert!(!changed);
    }

    #[test]
    fn replaces_changed_crd_with_existing_resource_version() {
        let rt = tokio::runtime::Builder::new_current_thread().build().unwrap();
        let changed = rt.block_on(async {
            let mut old_crd = CrdUpdateTest::crd();
            old_crd.metadata.resource_version = Some("rv-1".to_string());

            let mut new_crd = old_crd.clone();
            new_crd.spec.versions[0].served = !new_crd.spec.versions[0].served;

            let path = crd_path(&new_crd);
            let (client, captured) = client_with_responses(vec![json_response(&old_crd), json_response(&new_crd)]);

            let changed = update_crd(new_crd.clone(), &client).await.unwrap();
            let captured = captured.lock().unwrap();
            let body = request_json(&captured[1]);

            assert_eq!(captured.len(), 2);
            assert_eq!(captured[0].method, Method::GET);
            assert_eq!(request_path(&captured[0]), path);
            assert_eq!(captured[1].method, Method::PUT);
            assert_eq!(request_path(&captured[1]), path);
            assert_eq!(body["metadata"]["resourceVersion"], "rv-1");
            assert_eq!(
                body["spec"]["versions"],
                serde_json::to_value(&new_crd.spec.versions).unwrap()
            );

            changed
        });

        assert!(changed);
    }

    #[test]
    fn returns_error_when_create_or_replace_request_fails() {
        let rt = tokio::runtime::Builder::new_current_thread().build().unwrap();

        rt.block_on(async {
            let crd = CrdUpdateTest::crd();
            let (client, captured) = client_with_responses(vec![not_found_response(), server_error_response()]);

            let error = update_crd(crd, &client).await.unwrap_err();

            assert!(matches!(error, UpdateCrdError::Kube(kube::Error::Api(_))));
            let captured = captured.lock().unwrap();
            assert_eq!(captured.len(), 2);
            assert_eq!(captured[1].method, Method::POST);
        });

        rt.block_on(async {
            let mut old_crd = CrdUpdateTest::crd();
            old_crd.metadata.resource_version = Some("rv-1".to_string());

            let mut new_crd = old_crd.clone();
            new_crd.spec.versions[0].served = !new_crd.spec.versions[0].served;

            let (client, captured) = client_with_responses(vec![json_response(&old_crd), server_error_response()]);

            let error = update_crd(new_crd, &client).await.unwrap_err();

            assert!(matches!(error, UpdateCrdError::Kube(kube::Error::Api(_))));
            let captured = captured.lock().unwrap();
            assert_eq!(captured.len(), 2);
            assert_eq!(captured[1].method, Method::PUT);
        });
    }
}