mirror of
https://github.com/glomatico/gamdl.git
synced 2026-08-03 13:27:13 +03:00
Use persistent WV2D decrypt sessions
This commit is contained in:
@@ -32,14 +32,14 @@ A command-line app for downloading Apple Music songs, music videos and post vide
|
||||
|
||||
#### Wrapper
|
||||
|
||||
Run the [Wrapper v2](https://github.com/glomatico/wrapper-v2) server for wrapper-backed account, playback, and decryption requests. Enable it with `--use-wrapper` or `use_wrapper = true`. Configure wrapper HTTP account/playback calls with `--wrapper-url` or `wrapper_url`, and configure raw TCP decrypt with `--wrapper-decrypt-host` / `--wrapper-decrypt-port`.
|
||||
Run the [Wrapper v2](https://github.com/glomatico/wrapper-v2) server for wrapper-backed account, playback, and decryption requests. Enable it with `--use-wrapper` or `use_wrapper = true`. Configure wrapper HTTP account/playback calls with `--wrapper-url` or `wrapper_url`, and configure WV2D batch TCP decrypt with `--wrapper-decrypt-host` / `--wrapper-decrypt-port`.
|
||||
|
||||
The wrapper is recommended when using the `alac` song codec. ALAC can be attempted without wrapper, but it probably won't work due to API limitations.
|
||||
|
||||
**Note:**
|
||||
|
||||
- When using the Wrapper, you'll be asked to insert your credentials to login if you haven't already.
|
||||
- Newer wrapper-v2 builds use HTTP JSON for account/playback and TCP port `10020` for decrypt.
|
||||
- Newer wrapper-v2 builds use HTTP JSON for account/playback and WV2D batch TCP port `10020` for decrypt.
|
||||
- Song codecs other than `alac` do not require the wrapper.
|
||||
- Cookies can be skipped when using the wrapper.
|
||||
|
||||
|
||||
@@ -1029,6 +1029,10 @@ async def decrypt_samples(
|
||||
segment_uri: Optional[str] = None
|
||||
# Pending (sample, aligned_cbc, tail) for one SKD segment, flushed in batches.
|
||||
crypto_batch: List[tuple] = []
|
||||
wrapper_decrypt_session = _amdecrypt.WrapperDecryptSession(
|
||||
wrapper_api.decrypt_host,
|
||||
wrapper_api.decrypt_port,
|
||||
)
|
||||
|
||||
def emit(data: bytes) -> None:
|
||||
nonlocal decrypted_bytes
|
||||
@@ -1048,9 +1052,7 @@ async def decrypt_samples(
|
||||
for sample, aligned, tail in crypto_batch
|
||||
]
|
||||
reassembled = await asyncio.to_thread(
|
||||
_amdecrypt.wrapper_decrypt_reassemble,
|
||||
wrapper_api.decrypt_host,
|
||||
wrapper_api.decrypt_port,
|
||||
wrapper_decrypt_session.decrypt_reassemble,
|
||||
segment_adam,
|
||||
segment_uri,
|
||||
native_items,
|
||||
@@ -1144,6 +1146,7 @@ async def decrypt_samples(
|
||||
|
||||
await flush_crypto_batch()
|
||||
finally:
|
||||
wrapper_decrypt_session.close()
|
||||
if decrypted_output:
|
||||
decrypted_output.close()
|
||||
|
||||
|
||||
+217
-64
@@ -1,9 +1,16 @@
|
||||
use pyo3::exceptions::{PyIOError, PyValueError};
|
||||
use pyo3::exceptions::{PyIOError, PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpStream;
|
||||
use std::time::Duration;
|
||||
|
||||
const DECRYPT_MAGIC: u32 = 0x57563244; // WV2D
|
||||
const DECRYPT_VERSION: u16 = 1;
|
||||
const DECRYPT_KIND_BATCH: u16 = 1;
|
||||
const DECRYPT_KIND_OK: u16 = 2;
|
||||
const DECRYPT_KIND_ERROR: u16 = 3;
|
||||
const DECRYPT_KIND_CLOSE: u16 = 9;
|
||||
|
||||
#[pyfunction]
|
||||
fn native_available() -> bool {
|
||||
true
|
||||
@@ -24,7 +31,7 @@ fn validate_label(name: &str, value: &str) -> PyResult<Vec<u8>> {
|
||||
if bytes.is_empty() {
|
||||
return Err(value_error(format!("wrapper-v2: {name} must not be empty")));
|
||||
}
|
||||
if bytes.len() > u8::MAX as usize {
|
||||
if bytes.len() > u16::MAX as usize {
|
||||
return Err(value_error(format!(
|
||||
"wrapper-v2: {name} is too long for TCP decrypt protocol"
|
||||
)));
|
||||
@@ -113,45 +120,30 @@ fn reassemble_sample(
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn tcp_decrypt_reassemble(
|
||||
host: &str,
|
||||
port: u16,
|
||||
fn build_decrypt_batch_payload(
|
||||
adam_id: &[u8],
|
||||
skd_uri: &[u8],
|
||||
items: Vec<BatchItem>,
|
||||
) -> PyResult<Vec<Vec<u8>>> {
|
||||
items: &[BatchItem],
|
||||
) -> PyResult<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 items.len() > u32::MAX as usize {
|
||||
return Err(value_error("wrapper-v2: ciphertext batch is too large"));
|
||||
}
|
||||
let mut size = 8usize
|
||||
.checked_add(
|
||||
items
|
||||
.len()
|
||||
.checked_mul(4)
|
||||
.ok_or_else(|| value_error("wrapper-v2: decrypt batch size overflow"))?,
|
||||
)
|
||||
.and_then(|n| n.checked_add(adam_id.len()))
|
||||
.and_then(|n| n.checked_add(skd_uri.len()))
|
||||
.ok_or_else(|| value_error("wrapper-v2: decrypt batch size overflow"))?;
|
||||
for (idx, (_, aligned, _, _)) in items.iter().enumerate() {
|
||||
if aligned.is_empty() {
|
||||
return Err(value_error(format!(
|
||||
"wrapper-v2: ciphertext sample {idx} must not be empty"
|
||||
@@ -162,45 +154,206 @@ fn tcp_decrypt_reassemble(
|
||||
"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)?);
|
||||
size = size
|
||||
.checked_add(aligned.len())
|
||||
.ok_or_else(|| value_error("wrapper-v2: decrypt batch size overflow"))?;
|
||||
}
|
||||
|
||||
stream.write_all(&0u32.to_ne_bytes()).map_err(|e| {
|
||||
io_error(format!(
|
||||
"wrapper-v2: TCP decrypt terminator write failed: {e}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let mut out = Vec::with_capacity(size);
|
||||
out.extend_from_slice(&(adam_id.len() as u16).to_be_bytes());
|
||||
out.extend_from_slice(&(skd_uri.len() as u16).to_be_bytes());
|
||||
out.extend_from_slice(&(items.len() as u32).to_be_bytes());
|
||||
for (_, aligned, _, _) in items {
|
||||
out.extend_from_slice(&(aligned.len() as u32).to_be_bytes());
|
||||
}
|
||||
out.extend_from_slice(adam_id);
|
||||
out.extend_from_slice(skd_uri);
|
||||
for (_, aligned, _, _) in items {
|
||||
out.extend_from_slice(aligned);
|
||||
}
|
||||
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))
|
||||
fn read_decrypt_samples_payload(data: &[u8]) -> PyResult<Vec<Vec<u8>>> {
|
||||
if data.len() < 4 {
|
||||
return Err(io_error("wrapper-v2: decrypt response too short"));
|
||||
}
|
||||
let sample_count = u32::from_be_bytes([data[0], data[1], data[2], data[3]]) as usize;
|
||||
let table_end = 4usize
|
||||
.checked_add(
|
||||
sample_count
|
||||
.checked_mul(4)
|
||||
.ok_or_else(|| io_error("wrapper-v2: decrypt response overflow"))?,
|
||||
)
|
||||
.ok_or_else(|| io_error("wrapper-v2: decrypt response overflow"))?;
|
||||
if data.len() < table_end {
|
||||
return Err(io_error("wrapper-v2: truncated decrypt length table"));
|
||||
}
|
||||
let mut lengths = Vec::with_capacity(sample_count);
|
||||
for i in 0..sample_count {
|
||||
let off = 4 + i * 4;
|
||||
lengths.push(
|
||||
u32::from_be_bytes([data[off], data[off + 1], data[off + 2], data[off + 3]]) as usize,
|
||||
);
|
||||
}
|
||||
let mut offset = table_end;
|
||||
let mut out = Vec::with_capacity(sample_count);
|
||||
for len in lengths {
|
||||
let end = offset
|
||||
.checked_add(len)
|
||||
.ok_or_else(|| io_error("wrapper-v2: decrypt response overflow"))?;
|
||||
if end > data.len() {
|
||||
return Err(io_error("wrapper-v2: truncated plaintext sample"));
|
||||
}
|
||||
out.push(data[offset..end].to_vec());
|
||||
offset = end;
|
||||
}
|
||||
if offset != data.len() {
|
||||
return Err(io_error("wrapper-v2: trailing decrypt response bytes"));
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn read_frame(stream: &mut TcpStream) -> PyResult<(u16, u32, Vec<u8>)> {
|
||||
let mut h = [0u8; 16];
|
||||
stream.read_exact(&mut h).map_err(|e| {
|
||||
io_error(format!(
|
||||
"wrapper-v2: TCP decrypt truncated frame header: {e}"
|
||||
))
|
||||
})?;
|
||||
let magic = u32::from_be_bytes([h[0], h[1], h[2], h[3]]);
|
||||
let version = u16::from_be_bytes([h[4], h[5]]);
|
||||
if magic != DECRYPT_MAGIC {
|
||||
return Err(io_error("wrapper-v2: bad decrypt response magic"));
|
||||
}
|
||||
if version != DECRYPT_VERSION {
|
||||
return Err(io_error("wrapper-v2: bad decrypt response version"));
|
||||
}
|
||||
let kind = u16::from_be_bytes([h[6], h[7]]);
|
||||
let request_id = u32::from_be_bytes([h[8], h[9], h[10], h[11]]);
|
||||
let payload_len = u32::from_be_bytes([h[12], h[13], h[14], h[15]]) as usize;
|
||||
let mut payload = vec![0u8; payload_len];
|
||||
stream.read_exact(&mut payload).map_err(|e| {
|
||||
io_error(format!(
|
||||
"wrapper-v2: TCP decrypt truncated frame payload: {e}"
|
||||
))
|
||||
})?;
|
||||
Ok((kind, request_id, payload))
|
||||
}
|
||||
|
||||
fn write_frame(stream: &mut TcpStream, kind: u16, request_id: u32, payload: &[u8]) -> PyResult<()> {
|
||||
if payload.len() > u32::MAX as usize {
|
||||
return Err(value_error("wrapper-v2: decrypt frame is too large"));
|
||||
}
|
||||
stream
|
||||
.write_all(&DECRYPT_MAGIC.to_be_bytes())
|
||||
.and_then(|_| stream.write_all(&DECRYPT_VERSION.to_be_bytes()))
|
||||
.and_then(|_| stream.write_all(&kind.to_be_bytes()))
|
||||
.and_then(|_| stream.write_all(&request_id.to_be_bytes()))
|
||||
.and_then(|_| stream.write_all(&(payload.len() as u32).to_be_bytes()))
|
||||
.and_then(|_| stream.write_all(payload))
|
||||
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt frame write failed: {e}")))
|
||||
}
|
||||
|
||||
#[pyclass]
|
||||
struct WrapperDecryptSession {
|
||||
stream: Option<TcpStream>,
|
||||
next_request_id: u32,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl WrapperDecryptSession {
|
||||
#[new]
|
||||
fn new(host: String, port: u16) -> PyResult<Self> {
|
||||
let stream = TcpStream::connect((host.as_str(), port))
|
||||
.map_err(|e| io_error(format!("wrapper-v2: TCP decrypt connect failed: {e}")))?;
|
||||
stream
|
||||
.set_nodelay(true)
|
||||
.map_err(|e| io_error(format!("wrapper-v2: TCP_NODELAY setup 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}"
|
||||
))
|
||||
})?;
|
||||
Ok(Self {
|
||||
stream: Some(stream),
|
||||
next_request_id: 1,
|
||||
})
|
||||
}
|
||||
|
||||
fn decrypt_reassemble(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
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)?;
|
||||
let request_id = self.next_request_id;
|
||||
self.next_request_id = self.next_request_id.wrapping_add(1).max(1);
|
||||
let stream = self
|
||||
.stream
|
||||
.as_mut()
|
||||
.ok_or_else(|| PyRuntimeError::new_err("wrapper-v2: decrypt session is closed"))?;
|
||||
py.allow_threads(move || {
|
||||
let payload = build_decrypt_batch_payload(&adam_id, &skd_uri, &items)?;
|
||||
write_frame(stream, DECRYPT_KIND_BATCH, request_id, &payload)?;
|
||||
let (kind, response_id, response_payload) = read_frame(stream)?;
|
||||
if response_id != request_id {
|
||||
return Err(io_error("wrapper-v2: mismatched decrypt response id"));
|
||||
}
|
||||
if kind == DECRYPT_KIND_ERROR {
|
||||
return Err(io_error(format!(
|
||||
"wrapper-v2: decrypt failed: {}",
|
||||
String::from_utf8_lossy(&response_payload)
|
||||
)));
|
||||
}
|
||||
if kind != DECRYPT_KIND_OK {
|
||||
return Err(io_error("wrapper-v2: unexpected decrypt response kind"));
|
||||
}
|
||||
let plains = read_decrypt_samples_payload(&response_payload)?;
|
||||
if plains.len() != items.len() {
|
||||
return Err(io_error(format!(
|
||||
"wrapper-v2: expected {} plaintexts, got {}",
|
||||
items.len(),
|
||||
plains.len()
|
||||
)));
|
||||
}
|
||||
let mut out = Vec::with_capacity(items.len());
|
||||
for ((data, _, tail, subsamples), plain) in items.into_iter().zip(plains) {
|
||||
out.push(reassemble_sample(&data, &plain, &tail, &subsamples)?);
|
||||
}
|
||||
Ok(out)
|
||||
})
|
||||
}
|
||||
|
||||
fn close(&mut self) -> PyResult<()> {
|
||||
if let Some(mut stream) = self.stream.take() {
|
||||
let _ = write_frame(&mut stream, DECRYPT_KIND_CLOSE, 0, &[]);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WrapperDecryptSession {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.close();
|
||||
}
|
||||
}
|
||||
|
||||
#[pymodule]
|
||||
fn _amdecrypt(m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_function(wrap_pyfunction!(native_available, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(wrapper_decrypt_reassemble, m)?)?;
|
||||
m.add_class::<WrapperDecryptSession>()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+128
-65
@@ -6,6 +6,13 @@ import pytest
|
||||
|
||||
from gamdl import _amdecrypt
|
||||
|
||||
MAGIC = b"WV2D"
|
||||
VERSION = 1
|
||||
KIND_BATCH = 1
|
||||
KIND_OK = 2
|
||||
KIND_ERROR = 3
|
||||
KIND_CLOSE = 9
|
||||
|
||||
|
||||
def _recv_exact(conn: socket.socket, size: int) -> bytes:
|
||||
out = bytearray()
|
||||
@@ -17,14 +24,48 @@ def _recv_exact(conn: socket.socket, size: int) -> bytes:
|
||||
return bytes(out)
|
||||
|
||||
|
||||
def _start_fake_wrapper(
|
||||
expected_prefix: bytes,
|
||||
plaintexts: list[bytes],
|
||||
*,
|
||||
close_after_plaintexts: bool = False,
|
||||
):
|
||||
def _frame(kind: int, request_id: int, payload: bytes) -> bytes:
|
||||
return (
|
||||
MAGIC
|
||||
+ struct.pack(">HHII", VERSION, kind, request_id, len(payload))
|
||||
+ payload
|
||||
)
|
||||
|
||||
|
||||
def _parse_batch_payload(payload: bytes):
|
||||
adam_len, uri_len, sample_count = struct.unpack_from(">HHI", payload, 0)
|
||||
off = 8
|
||||
lengths = [
|
||||
struct.unpack_from(">I", payload, off + i * 4)[0]
|
||||
for i in range(sample_count)
|
||||
]
|
||||
off += sample_count * 4
|
||||
adam = payload[off : off + adam_len]
|
||||
off += adam_len
|
||||
uri = payload[off : off + uri_len]
|
||||
off += uri_len
|
||||
samples = []
|
||||
for length in lengths:
|
||||
samples.append(payload[off : off + length])
|
||||
off += length
|
||||
assert off == len(payload)
|
||||
return adam, uri, samples
|
||||
|
||||
|
||||
def _ok_payload(plaintexts: list[bytes]) -> bytes:
|
||||
out = bytearray()
|
||||
out += struct.pack(">I", len(plaintexts))
|
||||
for plain in plaintexts:
|
||||
out += struct.pack(">I", len(plain))
|
||||
for plain in plaintexts:
|
||||
out += plain
|
||||
return bytes(out)
|
||||
|
||||
|
||||
def _start_fake_wrapper(responses: list[list[bytes]]):
|
||||
ready = threading.Event()
|
||||
seen = bytearray()
|
||||
seen: list[tuple[int, int, bytes, bytes, list[bytes]]] = []
|
||||
close_seen = threading.Event()
|
||||
errors: list[BaseException] = []
|
||||
|
||||
def run(server: socket.socket):
|
||||
@@ -33,17 +74,33 @@ def _start_fake_wrapper(
|
||||
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)
|
||||
for plaintexts in responses:
|
||||
header = _recv_exact(conn, 16)
|
||||
magic = header[:4]
|
||||
version, kind, request_id, payload_len = struct.unpack(
|
||||
">HHII", header[4:]
|
||||
)
|
||||
payload = _recv_exact(conn, payload_len)
|
||||
assert magic == MAGIC
|
||||
assert version == VERSION
|
||||
assert kind == KIND_BATCH
|
||||
adam, uri, samples = _parse_batch_payload(payload)
|
||||
seen.append((kind, request_id, adam, uri, samples))
|
||||
conn.sendall(_frame(KIND_OK, request_id, _ok_payload(plaintexts)))
|
||||
|
||||
header = _recv_exact(conn, 16)
|
||||
if header:
|
||||
magic = header[:4]
|
||||
version, kind, request_id, payload_len = struct.unpack(
|
||||
">HHII", header[4:]
|
||||
)
|
||||
payload = _recv_exact(conn, payload_len)
|
||||
assert magic == MAGIC
|
||||
assert version == VERSION
|
||||
assert kind == KIND_CLOSE
|
||||
assert request_id == 0
|
||||
assert payload == b""
|
||||
close_seen.set()
|
||||
except BaseException as exc:
|
||||
errors.append(exc)
|
||||
finally:
|
||||
@@ -55,38 +112,39 @@ def _start_fake_wrapper(
|
||||
thread = threading.Thread(target=run, args=(server,), daemon=True)
|
||||
thread.start()
|
||||
assert ready.wait(5)
|
||||
return port, seen, errors, thread
|
||||
return port, seen, close_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)
|
||||
def test_wrapper_decrypt_session_wire_reassemble_and_persist():
|
||||
port, seen, close_seen, errors, thread = _start_fake_wrapper(
|
||||
[[b"AAAA", b"BBBBCCCC"], [b"ZZZZ"]]
|
||||
)
|
||||
|
||||
out = _amdecrypt.wrapper_decrypt_reassemble(
|
||||
"127.0.0.1",
|
||||
port,
|
||||
adam,
|
||||
uri,
|
||||
session = _amdecrypt.WrapperDecryptSession("127.0.0.1", port)
|
||||
out1 = session.decrypt_reassemble(
|
||||
"12345",
|
||||
"skd://example",
|
||||
[
|
||||
(b"xxxx", b"1111", b"", []),
|
||||
(b"aa22222222zz", b"22222222", b"", [(2, 8)]),
|
||||
],
|
||||
)
|
||||
out2 = session.decrypt_reassemble(
|
||||
"12345",
|
||||
"skd://example",
|
||||
[(b"yyyy", b"3333", b"", [])],
|
||||
)
|
||||
session.close()
|
||||
|
||||
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)
|
||||
)
|
||||
assert close_seen.is_set()
|
||||
assert out1 == [b"AAAA", b"aaBBBBCCCCzz"]
|
||||
assert out2 == [b"ZZZZ"]
|
||||
assert seen == [
|
||||
(KIND_BATCH, 1, b"12345", b"skd://example", [b"1111", b"22222222"]),
|
||||
(KIND_BATCH, 2, b"12345", b"skd://example", [b"3333"]),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -94,38 +152,43 @@ def test_wrapper_decrypt_reassemble_wire_and_order():
|
||||
[
|
||||
("", "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):
|
||||
def test_wrapper_decrypt_session_rejects_bad_inputs(adam, uri, items):
|
||||
port, _, _, _, thread = _start_fake_wrapper([])
|
||||
session = _amdecrypt.WrapperDecryptSession("127.0.0.1", port)
|
||||
with pytest.raises((OSError, ValueError)):
|
||||
_amdecrypt.wrapper_decrypt_reassemble(
|
||||
"127.0.0.1",
|
||||
9,
|
||||
adam,
|
||||
uri,
|
||||
items,
|
||||
)
|
||||
session.decrypt_reassemble(adam, uri, items)
|
||||
session.close()
|
||||
thread.join(5)
|
||||
|
||||
|
||||
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,
|
||||
)
|
||||
def test_wrapper_decrypt_session_rejects_error_frame():
|
||||
ready = threading.Event()
|
||||
|
||||
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"", [])],
|
||||
)
|
||||
def run(server: socket.socket):
|
||||
server.listen(1)
|
||||
ready.set()
|
||||
conn, _ = server.accept()
|
||||
with conn:
|
||||
header = _recv_exact(conn, 16)
|
||||
request_id = struct.unpack(">I", header[8:12])[0]
|
||||
payload_len = struct.unpack(">I", header[12:16])[0]
|
||||
_recv_exact(conn, payload_len)
|
||||
conn.sendall(_frame(KIND_ERROR, request_id, b"nope"))
|
||||
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)
|
||||
|
||||
session = _amdecrypt.WrapperDecryptSession("127.0.0.1", port)
|
||||
with pytest.raises(OSError, match="nope"):
|
||||
session.decrypt_reassemble("1", "skd://x", [(b"x" * 16, b"x" * 16, b"", [])])
|
||||
session.close()
|
||||
thread.join(5)
|
||||
|
||||
Reference in New Issue
Block a user