use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use pyo3::Py;
use synta::{Integer, ObjectIdentifier};
use synta_certificate::{
CertificateBuilder, CertificateSigner as _, PrivateKey as _, Time, UnsignedCertificateSigner,
};
use crate::certificate::cert::PyCertificate;
use crate::crypto_keys::PyPrivateKey;
fn py_to_synta_time(dt: &Bound<'_, PyAny>) -> PyResult<Time> {
if dt.getattr("tzinfo")?.is_none() {
return Err(PyValueError::new_err(
"datetime must be timezone-aware (tzinfo must not be None); \
use datetime.timezone.utc to specify UTC",
));
}
let year = dt.getattr("year")?.extract::<u32>()?;
let month = dt.getattr("month")?.extract::<u8>()?;
let day = dt.getattr("day")?.extract::<u8>()?;
let hour = dt.getattr("hour")?.extract::<u8>()?;
let minute = dt.getattr("minute")?.extract::<u8>()?;
let second = dt.getattr("second")?.extract::<u8>()?;
if (1950..=2049).contains(&year) {
Ok(Time::UtcTime(
synta::UtcTime::new(year as u16, month, day, hour, minute, second)
.map_err(|e| PyValueError::new_err(format!("invalid UTCTime: {e}")))?,
))
} else {
Ok(Time::GeneralTime(
synta::GeneralizedTime::new(year as u16, month, day, hour, minute, second, None)
.map_err(|e| PyValueError::new_err(format!("invalid GeneralizedTime: {e}")))?,
))
}
}
#[pyclass(name = "CertificateBuilder")]
pub struct PyCertificateBuilder {
issuer: Option<Vec<u8>>,
subject: Option<Vec<u8>>,
spki: Option<Vec<u8>>,
serial: Option<Integer>,
not_before: Option<Time>,
not_after: Option<Time>,
extensions: Vec<(ObjectIdentifier, bool, Vec<u8>)>,
}
#[pymethods]
impl PyCertificateBuilder {
#[new]
fn new() -> Self {
Self {
issuer: None,
subject: None,
spki: None,
serial: None,
not_before: None,
not_after: None,
extensions: Vec::new(),
}
}
fn issuer_name<'py>(slf: Bound<'py, Self>, name_der: &[u8]) -> Bound<'py, Self> {
slf.borrow_mut().issuer = Some(name_der.to_vec());
slf
}
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 serial_number<'py>(
slf: Bound<'py, Self>,
n: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, Self>> {
let serial = if let Ok(i) = n.extract::<i64>() {
Integer::from_i64(i)
} else if let Ok(b) = n.cast::<PyBytes>() {
Integer::from_unsigned_bytes(b.as_bytes())
} else if n.is_instance_of::<pyo3::types::PyInt>() {
let bit_length: usize = n.call_method0("bit_length")?.extract()?;
let byte_length = bit_length.div_ceil(8).max(1);
let bytes_obj = n.call_method1("to_bytes", (byte_length, "big"))?;
let b = bytes_obj.cast::<PyBytes>()?;
Integer::from_unsigned_bytes(b.as_bytes())
} else {
return Err(PyValueError::new_err("serial_number expects int or bytes"));
};
slf.borrow_mut().serial = Some(serial);
Ok(slf)
}
fn not_valid_before_utc<'py>(
slf: Bound<'py, Self>,
dt: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, Self>> {
slf.borrow_mut().not_before = Some(py_to_synta_time(dt)?);
Ok(slf)
}
fn not_valid_after_utc<'py>(
slf: Bound<'py, Self>,
dt: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, Self>> {
slf.borrow_mut().not_after = Some(py_to_synta_time(dt)?);
Ok(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, PyCertificate>> {
let mut builder = CertificateBuilder::new();
if let Some(ref b) = self.issuer {
builder = builder.issuer_name(b);
}
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);
}
if let Some(ref s) = self.serial {
builder = builder.serial_number(s.clone());
}
if let Some(ref t) = self.not_before {
builder = builder.not_valid_before(t.clone());
}
if let Some(ref t) = self.not_after {
builder = builder.not_valid_after(t.clone());
}
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 cert_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 tbs_der = builder
.build_tbs(&sig_alg_der)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let signature = key
.inner
.sign_ml_dsa_with_context(&tbs_der, ctx)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
CertificateBuilder::assemble(&tbs_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, &cert_der);
let cert = PyCertificate::new_from_der(py, py_bytes)?;
Py::new(py, cert).map(|py_cert| py_cert.into_bound(py))
}
fn sign_unsigned<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyCertificate>> {
let mut builder = CertificateBuilder::new();
if let Some(ref b) = self.issuer {
builder = builder.issuer_name(b);
}
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);
}
if let Some(ref s) = self.serial {
builder = builder.serial_number(s.clone());
}
if let Some(ref t) = self.not_before {
builder = builder.not_valid_before(t.clone());
}
if let Some(ref t) = self.not_after {
builder = builder.not_valid_after(t.clone());
}
for (oid, critical, value_bytes) in &self.extensions {
builder = builder.add_extension(oid.clone(), *critical, value_bytes);
}
let cert_der = builder
.sign(&UnsignedCertificateSigner)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let py_bytes = PyBytes::new(py, &cert_der);
let cert = PyCertificate::new_from_der(py, py_bytes)?;
Py::new(py, cert).map(|py_cert| py_cert.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!("CertificateBuilder(subject={subject})")
}
}