Add native wrapper TCP decrypt helper

This commit is contained in:
glomatico
2026-07-08 11:50:08 -03:00
parent 2075c40333
commit 2e55b841ff
3 changed files with 315 additions and 0 deletions
+2
View File
@@ -9,5 +9,7 @@ __pycache__
!README.md
!src
!src/**
!tests
!tests/**
!uv.lock
gamdl/*.so
+182
View File
@@ -1,12 +1,194 @@
use pyo3::exceptions::{PyIOError, PyValueError};
use pyo3::prelude::*;
use std::io::{Read, Write};
use std::net::TcpStream;
use std::time::Duration;
#[pyfunction]
fn native_available() -> bool {
true
}
type BatchItem = (Vec<u8>, Vec<u8>, Vec<u8>, Vec<(usize, usize)>);
fn value_error(message: impl Into<String>) -> PyErr {
PyValueError::new_err(message.into())
}
fn io_error(message: impl Into<String>) -> PyErr {
PyIOError::new_err(message.into())
}
fn validate_label(name: &str, value: &str) -> PyResult<Vec<u8>> {
let bytes = value.as_bytes();
if bytes.is_empty() {
return Err(value_error(format!("wrapper-v2: {name} must not be empty")));
}
if bytes.len() > u8::MAX as usize {
return Err(value_error(format!(
"wrapper-v2: {name} is too long for TCP decrypt protocol"
)));
}
Ok(bytes.to_vec())
}
fn reassemble_sample(
data: &[u8],
plain: &[u8],
tail: &[u8],
subsamples: &[(usize, usize)],
) -> PyResult<Vec<u8>> {
let mut full_dec = Vec::with_capacity(plain.len() + tail.len());
full_dec.extend_from_slice(plain);
full_dec.extend_from_slice(tail);
if subsamples.is_empty() {
if full_dec.len() != data.len() {
return Err(io_error(format!(
"decrypted sample length mismatch: expected {}, got {}",
data.len(),
full_dec.len()
)));
}
return Ok(full_dec);
}
let encrypted_total = subsamples
.iter()
.try_fold(0usize, |acc, (_, enc)| acc.checked_add(*enc))
.ok_or_else(|| value_error("subsample encrypted byte count overflow"))?;
if full_dec.len() != encrypted_total {
return Err(io_error(format!(
"decrypted subsample length mismatch: expected {}, got {}",
encrypted_total,
full_dec.len()
)));
}
let mut out = Vec::with_capacity(data.len());
let mut dec_off = 0usize;
let mut offset = 0usize;
for (clear_b, enc_b) in subsamples {
let clear_end = offset
.checked_add(*clear_b)
.ok_or_else(|| value_error("subsample clear byte offset overflow"))?;
if clear_end > data.len() {
return Err(value_error("subsample clear range exceeds sample size"));
}
if *clear_b > 0 {
out.extend_from_slice(&data[offset..clear_end]);
}
offset = clear_end;
let dec_end = dec_off
.checked_add(*enc_b)
.ok_or_else(|| value_error("subsample encrypted byte offset overflow"))?;
let enc_end = offset
.checked_add(*enc_b)
.ok_or_else(|| value_error("subsample encrypted range overflow"))?;
if dec_end > full_dec.len() {
return Err(value_error("subsample decrypt range exceeds plaintext size"));
}
if enc_end > data.len() {
return Err(value_error("subsample encrypted range exceeds sample size"));
}
if *enc_b > 0 {
out.extend_from_slice(&full_dec[dec_off..dec_end]);
}
dec_off = dec_end;
offset = enc_end;
}
if offset < data.len() {
out.extend_from_slice(&data[offset..]);
}
if out.len() != data.len() {
return Err(io_error(format!(
"reassembled sample length mismatch: expected {}, got {}",
data.len(),
out.len()
)));
}
Ok(out)
}
fn tcp_decrypt_reassemble(
host: &str,
port: u16,
adam_id: &[u8],
skd_uri: &[u8],
items: Vec<BatchItem>,
) -> PyResult<Vec<Vec<u8>>> {
if items.is_empty() {
return Err(value_error(
"wrapper-v2: ciphertext batch must not be empty",
));
}
let mut stream = TcpStream::connect((host, port))
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt connect failed: {e}")))?;
stream
.set_read_timeout(Some(Duration::from_secs(600)))
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt read timeout setup failed: {e}")))?;
stream
.set_write_timeout(Some(Duration::from_secs(600)))
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt write timeout setup failed: {e}")))?;
stream
.write_all(&[adam_id.len() as u8])
.and_then(|_| stream.write_all(adam_id))
.and_then(|_| stream.write_all(&[skd_uri.len() as u8]))
.and_then(|_| stream.write_all(skd_uri))
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt header write failed: {e}")))?;
let mut out = Vec::with_capacity(items.len());
for (idx, (data, aligned, tail, subsamples)) in items.into_iter().enumerate() {
if aligned.is_empty() {
return Err(value_error(format!(
"wrapper-v2: ciphertext sample {idx} must not be empty"
)));
}
if aligned.len() > u32::MAX as usize {
return Err(value_error(format!(
"wrapper-v2: ciphertext sample {idx} is too large"
)));
}
stream
.write_all(&(aligned.len() as u32).to_ne_bytes())
.and_then(|_| stream.write_all(&aligned))
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt sample write failed: {e}")))?;
let mut plain = vec![0u8; aligned.len()];
stream
.read_exact(&mut plain)
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt truncated plaintext: {e}")))?;
out.push(reassemble_sample(&data, &plain, &tail, &subsamples)?);
}
stream
.write_all(&0u32.to_ne_bytes())
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt terminator write failed: {e}")))?;
Ok(out)
}
#[pyfunction]
fn wrapper_decrypt_reassemble(
py: Python<'_>,
host: String,
port: u16,
adam_id: String,
skd_uri: String,
items: Vec<BatchItem>,
) -> PyResult<Vec<Vec<u8>>> {
let adam_id = validate_label("adam_id", &adam_id)?;
let skd_uri = validate_label("skd_uri", &skd_uri)?;
py.allow_threads(move || tcp_decrypt_reassemble(&host, port, &adam_id, &skd_uri, items))
}
#[pymodule]
fn _amdecrypt(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(native_available, m)?)?;
m.add_function(wrap_pyfunction!(wrapper_decrypt_reassemble, m)?)?;
Ok(())
}
+131
View File
@@ -0,0 +1,131 @@
import socket
import struct
import threading
import pytest
from gamdl import _amdecrypt
def _recv_exact(conn: socket.socket, size: int) -> bytes:
out = bytearray()
while len(out) < size:
chunk = conn.recv(size - len(out))
if not chunk:
break
out.extend(chunk)
return bytes(out)
def _start_fake_wrapper(
expected_prefix: bytes,
plaintexts: list[bytes],
*,
close_after_plaintexts: bool = False,
):
ready = threading.Event()
seen = bytearray()
errors: list[BaseException] = []
def run(server: socket.socket):
try:
server.listen(1)
ready.set()
conn, _ = server.accept()
with conn:
seen.extend(_recv_exact(conn, len(expected_prefix)))
for plain in plaintexts:
size = struct.unpack("=I", _recv_exact(conn, 4))[0]
seen.extend(struct.pack("=I", size))
ciphertext = _recv_exact(conn, size)
seen.extend(ciphertext)
conn.sendall(plain)
if close_after_plaintexts:
conn.shutdown(socket.SHUT_WR)
terminator = _recv_exact(conn, 4)
seen.extend(terminator)
except BaseException as exc:
errors.append(exc)
finally:
server.close()
server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
server.bind(("127.0.0.1", 0))
port = server.getsockname()[1]
thread = threading.Thread(target=run, args=(server,), daemon=True)
thread.start()
assert ready.wait(5)
return port, seen, errors, thread
def test_wrapper_decrypt_reassemble_wire_and_order():
adam = "12345"
uri = "skd://example"
expected_prefix = bytes([len(adam)]) + adam.encode() + bytes([len(uri)]) + uri.encode()
plaintexts = [b"AAAA", b"BBBBCCCC"]
port, seen, errors, thread = _start_fake_wrapper(expected_prefix, plaintexts)
out = _amdecrypt.wrapper_decrypt_reassemble(
"127.0.0.1",
port,
adam,
uri,
[
(b"xxxx", b"1111", b"", []),
(b"aa22222222zz", b"22222222", b"", [(2, 8)]),
],
)
thread.join(5)
assert not errors
assert out == [b"AAAA", b"aaBBBBCCCCzz"]
assert bytes(seen) == (
expected_prefix
+ struct.pack("=I", 4)
+ b"1111"
+ struct.pack("=I", 8)
+ b"22222222"
+ struct.pack("=I", 0)
)
@pytest.mark.parametrize(
("adam", "uri", "items"),
[
("", "skd://x", [(b"x" * 16, b"x" * 16, b"", [])]),
("1", "", [(b"x" * 16, b"x" * 16, b"", [])]),
("1" * 256, "skd://x", [(b"x" * 16, b"x" * 16, b"", [])]),
("1", "skd://x", []),
("1", "skd://x", [(b"x", b"", b"", [])]),
],
)
def test_wrapper_decrypt_reassemble_rejects_bad_inputs(adam, uri, items):
with pytest.raises((OSError, ValueError)):
_amdecrypt.wrapper_decrypt_reassemble(
"127.0.0.1",
9,
adam,
uri,
items,
)
def test_wrapper_decrypt_reassemble_rejects_truncated_plaintext():
adam = "1"
uri = "skd://x"
expected_prefix = bytes([len(adam)]) + adam.encode() + bytes([len(uri)]) + uri.encode()
port, _, _, thread = _start_fake_wrapper(
expected_prefix,
[b"short"],
close_after_plaintexts=True,
)
with pytest.raises(OSError, match="truncated plaintext"):
_amdecrypt.wrapper_decrypt_reassemble(
"127.0.0.1",
port,
adam,
uri,
[(b"x" * 16, b"x" * 16, b"", [])],
)
thread.join(5)