Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 27 additions & 9 deletions hack/scripts/verify-custom-ca-e2e.sh
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ set -euo pipefail
NAMESPACE="wandb-ca-e2e"
NAME="wandb"
API_APP="api"
MYSQL_INSTANCE="default"
TIMEOUT="20m"
POLL_SECONDS=10

Expand Down Expand Up @@ -79,6 +80,18 @@ if actual != expected:
PY
}

url_param() {
local url="$1"
local key="$2"
python3 - "${url}" "${key}" <<'PY'
import sys
from urllib.parse import parse_qs, urlparse

url, key = sys.argv[1:]
print(parse_qs(urlparse(url).query).get(key, [""])[0])
PY
}

json_has_env() {
local name="$1"
jq -e --arg name "${name}" '[.spec.containers[]?.env[]? | select(.name == $name)] | length > 0' >/dev/null
Expand Down Expand Up @@ -203,7 +216,7 @@ wait_until "WeightsAndBiases ${NAMESPACE}/${NAME} status.ready=true" check_wandb
wandb_json="$(kubectl -n "${NAMESPACE}" get weightsandbiases.apps.wandb.com "${NAME}" -o json)"
inline_ca_count="$(echo "${wandb_json}" | jq -r '(.spec.global.customCACerts // []) | length')"
USER_CONFIGMAP="$(echo "${wandb_json}" | jq -r '.spec.global.caCertsConfigMap // ""')"
mysql_ca_enabled="$(echo "${wandb_json}" | jq -r '(((.spec.mysql.externalMysql.sslCa.name // "") | length) > 0 and ((.spec.mysql.externalMysql.sslCa.key // "") | length) > 0)')"
mysql_ca_enabled="$(echo "${wandb_json}" | jq -r --arg instance "${MYSQL_INSTANCE}" '(((.spec.mysql[$instance].externalMysql.sslCa.name // "") | length) > 0 and ((.spec.mysql[$instance].externalMysql.sslCa.key // "") | length) > 0)')"
redis_ca_enabled="$(echo "${wandb_json}" | jq -r '(((.spec.redis.externalRedis.sslCa.name // "") | length) > 0 and ((.spec.redis.externalRedis.sslCa.key // "") | length) > 0)')"

if [[ "${inline_ca_count}" == "0" && -z "${USER_CONFIGMAP}" ]]; then
Expand All @@ -219,9 +232,16 @@ if [[ -n "${USER_CONFIGMAP}" ]]; then
fi

if [[ "${mysql_ca_enabled}" == "true" ]]; then
mysql_url="$(secret_value wandb-mysql-connection url)"
# The connection bundle Secret is named per instance; take it from status
# rather than reconstructing the instance fingerprint here.
mysql_secret="$(echo "${wandb_json}" | jq -r --arg instance "${MYSQL_INSTANCE}" '.status.mysqlStatus[$instance].connection.url.name // ""')"
[[ -n "${mysql_secret}" ]] || fail "no MySQL connection Secret in status for instance ${MYSQL_INSTANCE}"
mysql_url="$(secret_value "${mysql_secret}" url)"
assert_url_param "${mysql_url}" "tls" "custom"
assert_url_param "${mysql_url}" "ssl-ca" "/etc/ssl/certs/mysql_ca.pem"
mysql_ca_path="$(url_param "${mysql_url}" "ssl-ca")"
[[ "${mysql_ca_path}" == /*/ca.pem ]] || fail "unexpected ssl-ca path ${mysql_ca_path} in MySQL URL"
mysql_ca_dir="$(dirname "${mysql_ca_path}")"
mysql_ca_volume="mysql-$(basename "${mysql_ca_dir}")"
log "MySQL connection URL includes expected CA parameters"
fi

Expand All @@ -238,9 +258,6 @@ workload_ref="${WORKLOAD_KIND} ${WORKLOAD_NAME}"
for env_name in SSL_CERT_FILE SSL_CERT_DIR REQUESTS_CA_BUNDLE; do
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_env "${env_name}" || fail "missing env ${env_name} on ${workload_ref}"
done
if [[ "${mysql_ca_enabled}" == "true" ]]; then
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_env MYSQL_CA_CERT_PATH || fail "missing env MYSQL_CA_CERT_PATH on ${workload_ref}"
fi

echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_volume wandb-ca-certs-root || fail "missing volume wandb-ca-certs-root on ${workload_ref}"
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_mount wandb-ca-certs-root /usr/local/share/ca-certificates/ ||
Expand All @@ -257,8 +274,9 @@ if [[ -n "${USER_CONFIGMAP}" ]]; then
fail "missing user CA ConfigMap mount"
fi
if [[ "${mysql_ca_enabled}" == "true" ]]; then
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_volume mysql-ca || fail "missing volume mysql-ca on ${workload_ref}"
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_mount mysql-ca /etc/ssl/certs/mysql_ca.pem ||
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_volume "${mysql_ca_volume}" ||
fail "missing volume ${mysql_ca_volume} on ${workload_ref}"
echo "${WORKLOAD_TEMPLATE_JSON}" | json_has_mount "${mysql_ca_volume}" "${mysql_ca_dir}" ||
fail "missing MySQL CA mount"
fi
if [[ "${redis_ca_enabled}" == "true" ]]; then
Expand All @@ -281,7 +299,7 @@ if [[ -n "${workload_pod}" ]]; then
pod_checks+=("test -d /usr/local/share/ca-certificates/configmap")
fi
if [[ "${mysql_ca_enabled}" == "true" ]]; then
pod_checks+=("test -s /etc/ssl/certs/mysql_ca.pem")
pod_checks+=("test -s ${mysql_ca_path}")
fi
if [[ "${redis_ca_enabled}" == "true" ]]; then
pod_checks+=("test -s /etc/ssl/certs/redis_ca.pem")
Expand Down
155 changes: 85 additions & 70 deletions internal/controller/infra/external/mysql/mysql.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,34 +2,16 @@ package mysql

import (
"context"
"fmt"
"net/url"

apiv2 "github.com/wandb/operator/api/v2"
"github.com/wandb/operator/internal/controller/infra/external"
"github.com/wandb/operator/internal/controller/infra/mysqlconnection"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/types"
"k8s.io/utils/ptr"
ctrl "sigs.k8s.io/controller-runtime"
"sigs.k8s.io/controller-runtime/pkg/client"
)

const ConnectionSecretName = "wandb-mysql-connection"
const caCertPath = "/etc/ssl/certs/mysql_ca.pem"
const sslCertPath = "/etc/ssl/certs/mysql_ssl_cert.pem"
const sslKeyPath = "/etc/ssl/certs/mysql_ssl_key.pem"

// connectionSecretName returns the connection secret name for an instance. The
// reserved default instance keeps the historical name for backward
// compatibility; other instances are suffixed with their key.
func connectionSecretName(key string) string {
if key == "" || key == apiv2.DefaultInstanceName {
return ConnectionSecretName
}
return fmt.Sprintf("%s-%s", ConnectionSecretName, key)
}

func WriteState(
ctx context.Context,
c client.Client,
Expand All @@ -54,41 +36,87 @@ func WriteState(
data, err := external.ResolveFields(ctx, c, wandb.Namespace, fields)
if err != nil {
logger.Error(err, "failed to resolve external mysql fields")
return []metav1.Condition{{
Type: "Reconciled",
Status: metav1.ConditionFalse,
Reason: "ApiError",
}}
return []metav1.Condition{
{
Type: mysqlconnection.ProviderReadyType,
Status: metav1.ConditionFalse,
Reason: "SourceSecretsUnavailable",
Message: err.Error(),
},
{
Type: mysqlconnection.ConnectionResolvedType,
Status: metav1.ConditionFalse,
Reason: "SourceSecretsUnavailable",
Message: err.Error(),
},
{
Type: mysqlconnection.BundleReadyType,
Status: metav1.ConditionFalse,
Reason: "ConnectionNotResolved",
},
{
Type: "Reconciled",
Status: metav1.ConditionFalse,
Reason: "ApiError",
},
}
}

dbUrl := url.URL{
Scheme: "mysql",
Host: fmt.Sprintf("%s:%s", data["Host"], data["Port"]),
User: url.UserPassword(data["Username"], data["Password"]),
Path: data["Database"],
material := mysqlconnection.Material{
Host: data["Host"],
Port: data["Port"],
Database: data["Database"],
Username: data["Username"],
Password: data["Password"],
TLS: data["Tls"],
CACert: []byte(data["SslCa"]),
ClientCert: []byte(data["SslCert"]),
ClientKey: []byte(data["SslKey"]),
}
values := dbUrl.Query()
if tls, ok := data["Tls"]; ok {
values.Set("tls", tls)
}
if _, ok := data["SslCa"]; ok {
if values.Get("tls") == "" {
values.Set("tls", "custom")
if _, err := mysqlconnection.Write(ctx, c, wandb, key, material); err != nil {
logger.Error(err, "failed to write external mysql connection bundle")
return []metav1.Condition{
{
Type: mysqlconnection.ProviderReadyType,
Status: metav1.ConditionTrue,
Reason: "SourceSecretsResolved",
},
{
Type: mysqlconnection.ConnectionResolvedType,
Status: metav1.ConditionFalse,
Reason: "InvalidConnection",
Message: err.Error(),
},
{
Type: mysqlconnection.BundleReadyType,
Status: metav1.ConditionFalse,
Reason: "InvalidConnection",
},
{
Type: "Reconciled",
Status: metav1.ConditionFalse,
Reason: "InvalidConnection",
},
}
values.Set("ssl-ca", caCertPath)
}
if _, ok := data["SslCert"]; ok {
values.Set("ssl-cert", sslCertPath)
}
if _, ok := data["SslKey"]; ok {
values.Set("ssl-key", sslKeyPath)
}
dbUrl.RawQuery = values.Encode()

data["url"] = dbUrl.String()

nsName := types.NamespacedName{Namespace: wandb.Namespace, Name: connectionSecretName(key)}
return external.WriteConnectionSecret(ctx, c, wandb, nsName, data)
return []metav1.Condition{
{
Type: mysqlconnection.ProviderReadyType,
Status: metav1.ConditionTrue,
Reason: "SourceSecretsResolved",
},
{
Type: mysqlconnection.ConnectionResolvedType,
Status: metav1.ConditionTrue,
Reason: "SourceSecretsResolved",
},
{
Type: mysqlconnection.BundleReadyType,
Status: metav1.ConditionTrue,
Reason: "BundleWritten",
},
}
}

func ReadState(
Expand All @@ -98,30 +126,17 @@ func ReadState(
key string,
newConditions []metav1.Condition,
) ([]metav1.Condition, *apiv2.MysqlConnection) {
nsName := types.NamespacedName{Namespace: wandb.Namespace, Name: connectionSecretName(key)}
_, conditions, found := external.ReadConnectionSecret(ctx, c, nsName, newConditions)
if !found {
return conditions, nil
}

localRef := corev1.LocalObjectReference{Name: nsName.Name}
return conditions, &apiv2.MysqlConnection{
URL: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "url", Optional: ptr.To(false)},
Host: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Host", Optional: ptr.To(false)},
Port: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Port", Optional: ptr.To(false)},
Database: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Database", Optional: ptr.To(false)},
Username: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Username", Optional: ptr.To(false)},
Password: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Password", Optional: ptr.To(false)},
Tls: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "Tls", Optional: ptr.To(true)},
SslCa: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "SslCa", Optional: ptr.To(true)},
SslCert: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "SslCert", Optional: ptr.To(true)},
SslKey: corev1.SecretKeySelector{LocalObjectReference: localRef, Key: "SslKey", Optional: ptr.To(true)},
connection, err := mysqlconnection.Read(ctx, c, wandb, key)
if err != nil {
return append(newConditions, metav1.Condition{
Type: "Reconciled",
Status: metav1.ConditionFalse,
Reason: "ApiError",
}), nil
}
return newConditions, connection
}

func DeleteConnectionSecret(ctx context.Context, c client.Client, wandb *apiv2.WeightsAndBiases, key string) error {
return external.DeleteConnectionSecret(ctx, c, types.NamespacedName{
Namespace: wandb.Namespace,
Name: connectionSecretName(key),
})
return mysqlconnection.Delete(ctx, c, wandb, key)
}
62 changes: 52 additions & 10 deletions internal/controller/infra/external/mysql/mysql_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,19 @@ package mysql

import (
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net/url"
"testing"
"time"

"github.com/stretchr/testify/require"
apiv2 "github.com/wandb/operator/api/v2"
"github.com/wandb/operator/internal/controller/infra/mysqlconnection"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
Expand Down Expand Up @@ -36,9 +44,7 @@ func TestWriteStateAddsCustomTLSParamsWhenCACertPresent(t *testing.T) {
"Database": []byte("wandb"),
"Username": []byte("wandb"),
"Password": []byte("secret"),
"SslCa": []byte("---ca---"),
"SslCert": []byte("---cert---"),
"SslKey": []byte("---key---"),
"SslCa": testCACertificate(t),
},
}
wandb := &apiv2.WeightsAndBiases{
Expand All @@ -53,29 +59,48 @@ func TestWriteStateAddsCustomTLSParamsWhenCACertPresent(t *testing.T) {
Username: mysqlSel("Username"),
Password: mysqlSel("Password"),
SslCa: mysqlSel("SslCa"),
SslCert: mysqlSel("SslCert"),
SslKey: mysqlSel("SslKey"),
},
}},
},
}
client := fake.NewClientBuilder().WithScheme(scheme).WithObjects(wandb, source).Build()

conditions := WriteState(context.Background(), client, wandb, apiv2.DefaultInstanceName, wandb.Spec.MySQL[apiv2.DefaultInstanceName].ExternalMysql)
require.Nil(t, conditions)
require.Len(t, conditions, 3)
for _, conditionType := range []string{
mysqlconnection.ProviderReadyType,
mysqlconnection.ConnectionResolvedType,
mysqlconnection.BundleReadyType,
} {
require.Condition(t, func() bool {
for _, condition := range conditions {
if condition.Type == conditionType {
return condition.Status == metav1.ConditionTrue
}
}
return false
})
}

written := &corev1.Secret{}
require.NoError(t, client.Get(context.Background(), types.NamespacedName{Name: ConnectionSecretName, Namespace: "default"}, written))
require.NoError(t, client.Get(context.Background(), types.NamespacedName{
Name: mysqlconnection.SecretName(wandb.Name, apiv2.DefaultInstanceName),
Namespace: "default",
}, written))
data := mysqlConnectionData(written)
parsed, err := url.Parse(data["url"])
require.NoError(t, err)
require.Equal(t, "mysql", parsed.Scheme)
require.Equal(t, "mysql.example.com:3306", parsed.Host)
require.Equal(t, "/wandb", parsed.Path)
require.Equal(t, "custom", parsed.Query().Get("tls"))
require.Equal(t, caCertPath, parsed.Query().Get("ssl-ca"))
require.Equal(t, sslCertPath, parsed.Query().Get("ssl-cert"))
require.Equal(t, sslKeyPath, parsed.Query().Get("ssl-key"))
require.Equal(
t,
mysqlconnection.MountPath(apiv2.DefaultInstanceName)+"/"+mysqlconnection.CACertFile,
parsed.Query().Get("ssl-ca"),
)
require.Empty(t, parsed.Query().Get("ssl-cert"))
require.Equal(t, mysqlconnection.BundleVersion, written.Annotations[mysqlconnection.BundleVersionAnnotation])
}

func mysqlConnectionData(secret *corev1.Secret) map[string]string {
Expand All @@ -88,3 +113,20 @@ func mysqlConnectionData(secret *corev1.Secret) map[string]string {
}
return out
}

func testCACertificate(t *testing.T) []byte {
t.Helper()
publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader)
require.NoError(t, err)
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "test CA"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
IsCA: true,
KeyUsage: x509.KeyUsageCertSign,
}
certificate, err := x509.CreateCertificate(rand.Reader, template, template, publicKey, privateKey)
require.NoError(t, err)
return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate})
}
Loading
Loading