use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use pyo3::Py;
use synta::ObjectIdentifier;
use synta_certificate::{CertificateSigner as _, CsrBuilder, PrivateKey as _};
use crate::certificate::pkix::PyCsr;
use crate::crypto_keys::PyPrivateKey;
#[pyclass(name = "CsrBuilder")]
pub struct PyCsrBuilder {
subject: Option<Vec<u8>>,
spki: Option<Vec<u8>>,
extensions: Vec<(ObjectIdentifier, bool, Vec<u8>)>,
}
#[pymethods]
impl PyCsrBuilder {
#[new]
fn new() -> Self {
Self {
subject: None,
spki: None,
extensions: Vec::new(),
}
}
fn subject_name<'py>(slf: Bound<'py, Self>, name_der: &[u8]) -> Bound<'py, Self> {
slf.borrow_mut().subject = Some(name_der.to_vec());
slf
}
fn public_key<'py>(
slf: Bound<'py, Self>,
key: &crate::crypto_keys::PyPublicKey,
) -> PyResult<Bound<'py, Self>> {
let der = key
.inner
.to_der()
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
slf.borrow_mut().spki = Some(der);
Ok(slf)
}
fn public_key_der<'py>(slf: Bound<'py, Self>, spki_der: &[u8]) -> Bound<'py, Self> {
slf.borrow_mut().spki = Some(spki_der.to_vec());
slf
}
fn add_extension<'py>(
slf: Bound<'py, Self>,
oid: &str,
critical: bool,
value_der: &[u8],
) -> PyResult<Bound<'py, Self>> {
use std::str::FromStr;
let oid = ObjectIdentifier::from_str(oid)
.map_err(|_| PyValueError::new_err(format!("invalid OID: {oid}")))?;
slf.borrow_mut()
.extensions
.push((oid, critical, value_der.to_vec()));
Ok(slf)
}
#[pyo3(signature = (key, algorithm, context = None))]
fn sign<'py>(
&self,
py: Python<'py>,
key: &PyPrivateKey,
algorithm: &str,
context: Option<&[u8]>,
) -> PyResult<Bound<'py, PyCsr>> {
let mut builder = CsrBuilder::new();
if let Some(ref b) = self.subject {
builder = builder.subject_name(b);
}
if let Some(ref b) = self.spki {
builder = builder.public_key_der(b);
}
for (oid, critical, value_bytes) in &self.extensions {
builder = builder.add_extension(oid.clone(), *critical, value_bytes);
}
let ctx = context.unwrap_or(b"");
let is_ml_dsa = matches!(
key.inner.key_type(),
"ml-dsa-44" | "ml-dsa-65" | "ml-dsa-87"
);
let csr_der = if !ctx.is_empty() && is_ml_dsa {
let signer = key.as_signer(algorithm);
let sig_alg_der = signer
.signature_algorithm_der()
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let cri_der = builder
.build_cri(&sig_alg_der)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let signature = key
.inner
.sign_ml_dsa_with_context(&cri_der, ctx)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
CsrBuilder::assemble(&cri_der, &sig_alg_der, &signature)
.map_err(|e| PyValueError::new_err(format!("{e}")))?
} else {
let signer = key.as_signer(algorithm);
builder
.sign(&signer)
.map_err(|e| PyValueError::new_err(format!("{e}")))?
};
let py_bytes = PyBytes::new(py, &csr_der);
let csr = PyCsr::new_from_der(py, py_bytes)?;
Py::new(py, csr).map(|py_csr| py_csr.into_bound(py))
}
fn __repr__(&self) -> String {
let subject = self
.subject
.as_ref()
.map(|b| format!("{} bytes", b.len()))
.unwrap_or_else(|| "not set".to_string());
format!("CsrBuilder(subject={subject})")
}
}