use k8s_openapi::apiextensions_apiserver::pkg::apis::apiextensions::v1::CustomResourceDefinition;
use kube::{
api::{Api, ResourceExt},
client::Client,
core::params::PostParams,
};
use thiserror::Error;
#[derive(Error, Debug)]
pub enum UpdateCrdError {
#[error("Kubernetes API request failed while updating CRD: `{0}`")]
Kube(#[from] kube::Error),
}
#[expect(rustdoc::broken_intra_doc_links)]
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);
});
}
}