use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use synta_certificate::{
default_block_cipher_provider, default_hmac_provider, default_pbkdf2_provider,
default_secure_random, default_streaming_hmac_provider, hkdf_expand, hkdf_extract,
hmac_output_len, BlockCipherProvider, HmacProvider, HmacState, Pbkdf2Provider, SecureRandom,
StreamingHmacProvider,
};
fn urlsafe_b64encode(data: &[u8]) -> String {
synta_certificate::encode_base64(data)
.replace('+', "-")
.replace('/', "_")
}
fn urlsafe_b64decode(data: &[u8]) -> PyResult<Vec<u8>> {
let s = std::str::from_utf8(data)
.map_err(|_| PyValueError::new_err("Fernet key/token must be valid ASCII"))?;
let standard = s.replace('-', "+").replace('_', "/");
synta_certificate::decode_base64(standard.as_bytes())
.ok_or_else(|| PyValueError::new_err("base64 decode error"))
}
#[pyfunction]
pub fn hmac_digest<'py>(
py: Python<'py>,
algorithm: &str,
key: &[u8],
data: &[u8],
) -> PyResult<Bound<'py, PyBytes>> {
let mac = default_hmac_provider()
.hmac_compute(algorithm, key, data)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &mac))
}
#[pyfunction]
pub fn hmac_verify(algorithm: &str, key: &[u8], data: &[u8], expected: &[u8]) -> PyResult<()> {
default_hmac_provider()
.hmac_verify(algorithm, key, data, expected)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(())
}
#[pyfunction]
pub fn pbkdf2_hmac<'py>(
py: Python<'py>,
algorithm: &str,
password: &[u8],
salt: &[u8],
iterations: u32,
length: usize,
) -> PyResult<Bound<'py, PyBytes>> {
let out = default_pbkdf2_provider()
.pbkdf2_hmac(algorithm, password, salt, iterations as usize, length)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &out))
}
#[pyfunction]
#[pyo3(signature = (key, iv, plaintext, pad = true))]
pub fn aes_cbc_encrypt<'py>(
py: Python<'py>,
key: &[u8],
iv: &[u8],
plaintext: &[u8],
pad: bool,
) -> PyResult<Bound<'py, PyBytes>> {
let ct = default_block_cipher_provider()
.aes_cbc_encrypt(key, iv, plaintext, pad)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &ct))
}
#[pyfunction]
#[pyo3(signature = (key, iv, ciphertext, unpad = true))]
pub fn aes_cbc_decrypt<'py>(
py: Python<'py>,
key: &[u8],
iv: &[u8],
ciphertext: &[u8],
unpad: bool,
) -> PyResult<Bound<'py, PyBytes>> {
let pt = default_block_cipher_provider()
.aes_cbc_decrypt(key, iv, ciphertext, unpad)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &pt))
}
#[pyfunction]
#[pyo3(signature = (key, nonce, plaintext, aad = None))]
pub fn aes_gcm_encrypt<'py>(
py: Python<'py>,
key: &[u8],
nonce: &[u8],
plaintext: &[u8],
aad: Option<&[u8]>,
) -> PyResult<Bound<'py, PyBytes>> {
let ct = default_block_cipher_provider()
.aes_gcm_encrypt(key, nonce, plaintext, aad.unwrap_or(&[]))
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &ct))
}
#[pyfunction]
#[pyo3(signature = (key, nonce, ciphertext_with_tag, aad = None))]
pub fn aes_gcm_decrypt<'py>(
py: Python<'py>,
key: &[u8],
nonce: &[u8],
ciphertext_with_tag: &[u8],
aad: Option<&[u8]>,
) -> PyResult<Bound<'py, PyBytes>> {
let pt = default_block_cipher_provider()
.aes_gcm_decrypt(key, nonce, ciphertext_with_tag, aad.unwrap_or(&[]))
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &pt))
}
#[pyfunction]
#[pyo3(signature = (key, iv, plaintext, pad = true))]
pub fn des3_cbc_encrypt<'py>(
py: Python<'py>,
key: &[u8],
iv: &[u8],
plaintext: &[u8],
pad: bool,
) -> PyResult<Bound<'py, PyBytes>> {
let ct = default_block_cipher_provider()
.des3_cbc_encrypt(key, iv, plaintext, pad)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &ct))
}
#[pyfunction]
#[pyo3(signature = (key, iv, ciphertext, unpad = true))]
pub fn des3_cbc_decrypt<'py>(
py: Python<'py>,
key: &[u8],
iv: &[u8],
ciphertext: &[u8],
unpad: bool,
) -> PyResult<Bound<'py, PyBytes>> {
let pt = default_block_cipher_provider()
.des3_cbc_decrypt(key, iv, ciphertext, unpad)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &pt))
}
#[pyfunction]
pub fn pkcs7_pad<'py>(
py: Python<'py>,
data: &[u8],
block_size: usize,
) -> PyResult<Bound<'py, PyBytes>> {
if block_size == 0 || block_size > 255 {
return Err(PyValueError::new_err(
"block_size must be between 1 and 255",
));
}
let pad_len = block_size - (data.len() % block_size);
let mut padded = Vec::with_capacity(data.len() + pad_len);
padded.extend_from_slice(data);
padded.extend(std::iter::repeat_n(pad_len as u8, pad_len));
Ok(PyBytes::new(py, &padded))
}
#[pyfunction]
pub fn pkcs7_unpad<'py>(
py: Python<'py>,
data: &[u8],
block_size: usize,
) -> PyResult<Bound<'py, PyBytes>> {
if data.is_empty() {
return Err(PyValueError::new_err("data is empty"));
}
let pad_byte = *data.last().unwrap() as usize;
if pad_byte == 0 || pad_byte > block_size || pad_byte > data.len() {
return Err(PyValueError::new_err("invalid PKCS#7 padding"));
}
let pad_start = data.len() - pad_byte;
if data[pad_start..].iter().any(|&b| b as usize != pad_byte) {
return Err(PyValueError::new_err("inconsistent PKCS#7 padding bytes"));
}
Ok(PyBytes::new(py, &data[..pad_start]))
}
#[pyclass(frozen, name = "Fernet")]
pub struct PyFernet {
signing_key: [u8; 16],
encryption_key: [u8; 16],
}
#[pymethods]
impl PyFernet {
#[new]
fn new(key: &[u8]) -> PyResult<Self> {
let raw = urlsafe_b64decode(key)?;
if raw.len() != 32 {
return Err(PyValueError::new_err(
"Fernet key must decode to exactly 32 bytes",
));
}
let mut signing_key = [0u8; 16];
let mut encryption_key = [0u8; 16];
signing_key.copy_from_slice(&raw[0..16]);
encryption_key.copy_from_slice(&raw[16..32]);
Ok(Self {
signing_key,
encryption_key,
})
}
#[staticmethod]
fn generate_key<'py>(py: Python<'py>) -> PyResult<Bound<'py, PyBytes>> {
let mut raw = [0u8; 32];
default_secure_random()
.rand_bytes(&mut raw)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, urlsafe_b64encode(&raw).as_bytes()))
}
fn encrypt<'py>(&self, py: Python<'py>, data: &[u8]) -> PyResult<Bound<'py, PyBytes>> {
let rng = default_secure_random();
let cipher = default_block_cipher_provider();
let hmac = default_hmac_provider();
let mut iv = [0u8; 16];
rng.rand_bytes(&mut iv)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let ts = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|_| PyValueError::new_err("system clock error"))?
.as_secs();
let ciphertext = cipher
.aes_cbc_encrypt(&self.encryption_key, &iv, data, true)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let mut pre_mac = Vec::with_capacity(1 + 8 + 16 + ciphertext.len());
pre_mac.push(0x80u8);
pre_mac.extend_from_slice(&ts.to_be_bytes());
pre_mac.extend_from_slice(&iv);
pre_mac.extend_from_slice(&ciphertext);
let mac = hmac
.hmac_compute("sha256", &self.signing_key, &pre_mac)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
let mut token_bytes = pre_mac;
token_bytes.extend_from_slice(&mac);
Ok(PyBytes::new(py, urlsafe_b64encode(&token_bytes).as_bytes()))
}
#[pyo3(signature = (token, ttl = None))]
fn decrypt<'py>(
&self,
py: Python<'py>,
token: &[u8],
ttl: Option<u64>,
) -> PyResult<Bound<'py, PyBytes>> {
let hmac = default_hmac_provider();
let cipher = default_block_cipher_provider();
let raw = urlsafe_b64decode(token)?;
if raw.len() < 1 + 8 + 16 + 16 + 32 {
return Err(PyValueError::new_err("invalid Fernet token: too short"));
}
if raw[0] != 0x80 {
return Err(PyValueError::new_err(
"invalid Fernet token: unknown version byte",
));
}
let ts = u64::from_be_bytes(
raw[1..9]
.try_into()
.map_err(|_| PyValueError::new_err("internal error: Fernet timestamp slice"))?,
);
let iv = &raw[9..25];
let ciphertext = &raw[25..raw.len() - 32];
let token_mac = &raw[raw.len() - 32..];
let pre_mac = &raw[..raw.len() - 32];
let computed_mac = hmac
.hmac_compute("sha256", &self.signing_key, pre_mac)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
if !synta_certificate::constant_time_eq(&computed_mac, token_mac) {
return Err(PyValueError::new_err(
"invalid Fernet token: HMAC verification failed",
));
}
if let Some(ttl_secs) = ttl {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|_| PyValueError::new_err("system clock error"))?
.as_secs();
const MAX_CLOCK_SKEW: u64 = 60;
if ts > now.saturating_add(MAX_CLOCK_SKEW) {
return Err(PyValueError::new_err(
"Fernet token timestamp is in the future",
));
}
if now.saturating_sub(ts) > ttl_secs {
return Err(PyValueError::new_err("Fernet token has expired"));
}
}
let plaintext = cipher
.aes_cbc_decrypt(&self.encryption_key, iv, ciphertext, true)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &plaintext))
}
}
#[pyfunction]
#[pyo3(signature = (algorithm, salt, ikm))]
pub fn hkdf_extract_py<'py>(
py: Python<'py>,
algorithm: &str,
salt: Option<&[u8]>,
ikm: &[u8],
) -> PyResult<Bound<'py, PyBytes>> {
let provider = default_hmac_provider();
let prk = hkdf_extract(&provider, algorithm, salt, ikm)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &prk))
}
#[pyfunction]
pub fn hkdf_expand_py<'py>(
py: Python<'py>,
algorithm: &str,
prk: &[u8],
info: &[u8],
length: usize,
) -> PyResult<Bound<'py, PyBytes>> {
if length == 0 {
return Ok(PyBytes::new(py, b""));
}
let hash_len = hmac_output_len(algorithm).ok_or_else(|| {
PyValueError::new_err(format!(
"unknown HKDF algorithm: {algorithm:?} \
(accepted: md5, sha1, sha224, sha256, sha384, sha512)"
))
})?;
if length > 255 * hash_len {
return Err(PyValueError::new_err(format!(
"HKDF-Expand length {length} exceeds maximum {} for {algorithm}",
255 * hash_len
)));
}
let provider = default_hmac_provider();
let okm = hkdf_expand(&provider, algorithm, prk, info, length)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyBytes::new(py, &okm))
}
#[pyclass(name = "HmacDigest")]
pub struct PyHmacDigest {
state: std::sync::Mutex<Option<Box<dyn HmacState>>>,
algorithm: String,
}
#[pymethods]
impl PyHmacDigest {
#[new]
fn new(algorithm: &str, key: &[u8]) -> PyResult<Self> {
let provider = default_streaming_hmac_provider();
let state = provider
.new_hmac(algorithm, key)
.map_err(|e| PyValueError::new_err(format!("{e}")))?;
Ok(PyHmacDigest {
state: std::sync::Mutex::new(Some(state)),
algorithm: algorithm.to_string(),
})
}
fn update(&self, data: &[u8]) -> PyResult<()> {
let mut guard = self.state.lock().unwrap();
match guard.as_mut() {
Some(s) => {
s.update(data);
Ok(())
}
None => Err(PyValueError::new_err(
"HmacDigest.update() called after finalize()",
)),
}
}
fn finalize<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyBytes>> {
let mut guard = self.state.lock().unwrap();
match guard.take() {
Some(s) => Ok(PyBytes::new(py, &s.finalize_boxed())),
None => Err(PyValueError::new_err(
"HmacDigest.finalize() called more than once",
)),
}
}
fn __repr__(&self) -> String {
let finalized = self.state.lock().unwrap().is_none();
format!(
"HmacDigest(algorithm={:?}, finalized={finalized})",
self.algorithm,
)
}
}
pub fn register_crypto_module(parent: &Bound<'_, PyModule>) -> PyResult<()> {
use pyo3::prelude::*;
let py = parent.py();
let m = PyModule::new(py, "crypto")?;
m.add_class::<PyFernet>()?;
m.add_class::<PyHmacDigest>()?;
m.add_function(wrap_pyfunction!(hmac_digest, &m)?)?;
m.add_function(wrap_pyfunction!(hmac_verify, &m)?)?;
m.add_function(wrap_pyfunction!(hkdf_extract_py, &m)?)?;
m.add_function(wrap_pyfunction!(hkdf_expand_py, &m)?)?;
m.add_function(wrap_pyfunction!(pbkdf2_hmac, &m)?)?;
m.add_function(wrap_pyfunction!(aes_cbc_encrypt, &m)?)?;
m.add_function(wrap_pyfunction!(aes_cbc_decrypt, &m)?)?;
m.add_function(wrap_pyfunction!(aes_gcm_encrypt, &m)?)?;
m.add_function(wrap_pyfunction!(aes_gcm_decrypt, &m)?)?;
m.add_function(wrap_pyfunction!(des3_cbc_encrypt, &m)?)?;
m.add_function(wrap_pyfunction!(des3_cbc_decrypt, &m)?)?;
m.add_function(wrap_pyfunction!(pkcs7_pad, &m)?)?;
m.add_function(wrap_pyfunction!(pkcs7_unpad, &m)?)?;
m.add_class::<crate::otp::PyHOTP>()?;
m.add_class::<crate::otp::PyTOTP>()?;
crate::install_submodule(
parent,
&m,
"synta.crypto",
Some(concat!(
"Symmetric cryptography primitives: HMAC, PBKDF2, AES-CBC, AES-GCM, ",
"3DES-CBC, PKCS#7 padding, and Fernet authenticated encryption ",
"(RFC-compatible with the Python ``cryptography`` library).",
)),
)
}