From 019b5b35ebf07022e37d6b8715820b439a18f32b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=81rni=20Dagur?= Date: Sun, 21 Jun 2020 18:20:20 -0400 Subject: Correct return type for the serialize method --- src/lib.rs | 11 +++++++++-- tests/test_engine.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 2 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index bbaa08c..588f4e9 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -17,6 +17,7 @@ use adblock::engine::Engine as RustEngine; use pyo3::class::PyObjectProtocol; use pyo3::exceptions::ValueError as PyValueError; use pyo3::prelude::*; +use pyo3::types::PyBytes; use pyo3::PyErr; use std::collections::HashMap; @@ -348,7 +349,13 @@ impl Engine { /// Serialize this blocking engine to bytes. They can then be deserialized /// using `deserialize()` to get the same engine again. #[text_signature = "($self)"] - pub fn serialize(&mut self) -> PyResult> { + pub fn serialize<'p>(&mut self, py: Python<'p>) -> PyResult<&'p PyBytes> { + let bytes = self.serialize_inner()?; + let py_bytes = PyBytes::new(py, &bytes); + Ok(py_bytes) + } + + fn serialize_inner(&mut self) -> PyResult> { let result = self.engine.serialize(); match result { Ok(x) => Ok(x), @@ -364,7 +371,7 @@ impl Engine { /// again. #[text_signature = "($self, file)"] pub fn serialize_to_file(&mut self, file: &str) -> PyResult<()> { - let data = self.serialize()?; + let data = self.serialize_inner()?; let mut fd = fs::OpenOptions::new() .create(true) .truncate(true) diff --git a/tests/test_engine.py b/tests/test_engine.py index af15b6d..28e25e6 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -1,4 +1,5 @@ import adblock +import pytest def test_engine_arguments(): @@ -9,3 +10,31 @@ def test_engine_arguments(): adblock.Engine(network_filters=[]) adblock.Engine(load_network=False, load_cosmetic=True, debug=False) adblock.Engine(debug=True) + + +def test_serde_file(tmpdir): + path = str(tmpdir / "cache.dat") + + engine0 = adblock.Engine() + with pytest.raises(FileNotFoundError): + # We haven't created the cache.dat file, so we should get an exception + # when attempting to deserialize. + engine0.deserialize_from_file(path) + + engine1 = adblock.Engine() + serialization_result = engine1.serialize_to_file(path) + assert serialization_result is None + + engine2 = adblock.Engine() + deserialization_result = engine2.deserialize_from_file(path) + assert deserialization_result is None + + +def test_serde(): + engine = adblock.Engine() + serialization_result = engine.serialize() + assert isinstance(serialization_result, bytes) + + engine2 = adblock.Engine() + deserialization_result = engine2.deserialize(serialization_result) + assert deserialization_result is None -- cgit v1.2.3