use alloc::collections::BTreeMap;
use alloc::format;
use alloc::string::String;
use alloc::vec::Vec;
use deser_core::State;
use deser_core::ext::{BigInt, ExtValue, Number};
use deser_core::ser::{self, EventSink, SerializeDriver, SerializeRef};
use deser_core::{Atom, Error, ErrorKind, Event, Serialize};
use crate::compat;
use crate::types::{ClassData, Form, FormData, Global, Kind, KindData, Reference, SharedIdData};
use crate::vm::HIGHEST_PROTOCOL;
const DEFAULT_PROTOCOL: u8 = 4;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SerializerConfig {
context: deser_core::Context,
protocol: u8,
}
impl Default for SerializerConfig {
fn default() -> SerializerConfig {
SerializerConfig::new()
}
}
impl SerializerConfig {
pub const fn new() -> SerializerConfig {
SerializerConfig {
context: deser_core::Context::new(),
protocol: DEFAULT_PROTOCOL,
}
}
pub const fn builder() -> SerializerConfigBuilder {
SerializerConfigBuilder::new()
}
pub const fn into_builder(self) -> SerializerConfigBuilder {
SerializerConfigBuilder { value: self }
}
pub fn set_context(&mut self, context: deser_core::Context) {
self.context = context;
}
pub fn context(&self) -> &deser_core::Context {
&self.context
}
pub fn set_protocol(&mut self, protocol: u8) {
self.protocol = protocol;
}
pub fn protocol(&self) -> u8 {
self.protocol
}
#[inline]
fn apply_context(&self, driver: &mut SerializeDriver<'_>) {
if !self.context.is_empty() {
driver.set_default_context(self.context.clone());
}
}
pub fn to_vec<T: Serialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>, Error> {
self.to_vec_ref(SerializeRef::new(&value))
}
pub fn to_vec_with<F, T: Serialize + ?Sized>(
&self,
value: &T,
setup: F,
) -> Result<Vec<u8>, Error>
where
F: FnOnce(&mut SerializeDriver<'_>),
{
let mut driver = SerializeDriver::new(&value);
setup(&mut driver);
self.apply_context(&mut driver);
serialize_driver(&mut driver, self.protocol)
}
fn to_vec_ref(&self, value: SerializeRef<'_>) -> Result<Vec<u8>, Error> {
let mut driver = SerializeDriver::from_ref(value);
self.apply_context(&mut driver);
serialize_driver(&mut driver, self.protocol)
}
}
#[derive(Debug, Clone)]
#[must_use]
pub struct SerializerConfigBuilder {
value: SerializerConfig,
}
impl SerializerConfigBuilder {
pub const fn new() -> SerializerConfigBuilder {
SerializerConfigBuilder {
value: SerializerConfig::new(),
}
}
pub fn context(mut self, context: deser_core::Context) -> SerializerConfigBuilder {
self.value.set_context(context);
self
}
pub const fn protocol(mut self, protocol: u8) -> SerializerConfigBuilder {
self.value.protocol = protocol;
self
}
pub const fn build(self) -> SerializerConfig {
let value = unsafe { core::ptr::read(&self.value) };
core::mem::forget(self);
value
}
}
impl Default for SerializerConfigBuilder {
fn default() -> SerializerConfigBuilder {
SerializerConfigBuilder::new()
}
}
#[derive(Debug, Clone, Default)]
pub struct Serializer {
config: SerializerConfig,
out: Vec<u8>,
}
impl Serializer {
pub fn new() -> Serializer {
Serializer::with_config(SerializerConfig::new())
}
pub fn with_config(config: SerializerConfig) -> Serializer {
Serializer {
config,
out: Vec::new(),
}
}
pub fn config(&self) -> &SerializerConfig {
&self.config
}
pub fn serialize<T: Serialize + ?Sized>(&mut self, value: &T) -> Result<(), Error> {
ser::Serializer::serialize(self, value)
}
pub fn serialize_with<F, T: Serialize + ?Sized>(
&mut self,
value: &T,
setup: F,
) -> Result<(), Error>
where
F: FnOnce(&mut SerializeDriver<'_>),
{
ser::Serializer::serialize_with(self, value, setup)
}
pub fn output(&self) -> &[u8] {
&self.out
}
pub fn finish(self) -> Vec<u8> {
self.out
}
}
impl ser::Serializer for Serializer {
fn drive(&mut self, driver: &mut SerializeDriver<'_>) -> Result<(), Error> {
self.config.apply_context(driver);
let bytes = serialize_driver(driver, self.config.protocol)?;
self.out.extend_from_slice(&bytes);
Ok(())
}
}
impl ser::StreamSerializer for Serializer {
fn output(&self) -> &[u8] {
&self.out
}
fn clear_output(&mut self) {
self.out.clear();
}
}
#[cfg(feature = "io")]
impl SerializerConfig {
pub fn writer<W: std::io::Write>(&self, writer: W) -> deser_core::io::Writer<W, Serializer> {
deser_core::io::Writer::new(writer, Serializer::with_config(self.clone()))
}
pub fn to_writer<W: std::io::Write, T: Serialize + ?Sized>(
&self,
writer: W,
value: &T,
) -> Result<(), Error> {
self.writer(writer).write(value)
}
}
#[cfg(feature = "io")]
pub fn to_writer<W: std::io::Write, T: Serialize + ?Sized>(
writer: W,
value: &T,
) -> Result<(), Error> {
SerializerConfig::new().to_writer(writer, value)
}
pub fn to_vec<T: Serialize + ?Sized>(value: &T) -> Result<Vec<u8>, Error> {
SerializerConfig::new().to_vec(value)
}
fn serialize_driver(driver: &mut SerializeDriver<'_>, protocol: u8) -> Result<Vec<u8>, Error> {
if !(2..=HIGHEST_PROTOCOL).contains(&protocol) {
return Err(Error::new(
ErrorKind::Configuration,
format!("unsupported pickle protocol {}", protocol),
));
}
let mut writer = Writer {
out: alloc::vec![0x80, protocol],
proto: protocol,
stack: Vec::new(),
memo: BTreeMap::new(),
memo_len: 0,
skip: 0,
started: false,
};
driver.drive_sink(&mut writer)?;
writer.finish()
}
#[derive(Clone, Copy, PartialEq)]
enum End {
Appends,
SetItems,
AddItems,
Tuple,
FrozenSet,
List,
Empty,
}
#[derive(Clone, Copy, Default)]
struct Suffix {
buf: [u8; 4],
len: u8,
}
impl Suffix {
fn new(bytes: &[u8]) -> Suffix {
Suffix::default().with(bytes)
}
fn with(mut self, bytes: &[u8]) -> Suffix {
let len = usize::from(self.len);
self.buf[len..len + bytes.len()].copy_from_slice(bytes);
self.len += bytes.len() as u8;
self
}
fn as_slice(&self) -> &[u8] {
&self.buf[..usize::from(self.len)]
}
}
struct Open {
end: End,
suffix: Suffix,
memoize: Option<u64>,
is_map: bool,
expects_key: bool,
in_key: bool,
}
struct Writer {
out: Vec<u8>,
proto: u8,
stack: Vec<Open>,
memo: BTreeMap<u64, u32>,
memo_len: u32,
skip: usize,
started: bool,
}
impl EventSink for Writer {
fn event(
&mut self,
event: Event<'_>,
_value: SerializeRef<'_>,
state: &mut State,
) -> Result<(), Error> {
if self.skip > 0 {
match event {
Event::MapStart(_) | Event::SeqStart(_) => self.skip += 1,
Event::MapEnd | Event::SeqEnd => {
self.skip -= 1;
if self.skip == 0 {
self.value_done();
}
}
Event::Atom(_) => {}
}
return Ok(());
}
let Some(open) = self.stack.last() else {
if self.started {
return Err(Error::new(ErrorKind::InvalidState, "unexpected event"));
}
self.started = true;
return self.value(event, state, false);
};
let in_key = open.in_key || open.is_map && open.expects_key;
match event {
Event::MapEnd if open.is_map && open.expects_key => self.close(),
Event::SeqEnd if !open.is_map => self.close(),
Event::MapEnd | Event::SeqEnd => {
Err(Error::new(ErrorKind::InvalidState, "unexpected end event"))
}
_ if open.end == End::Empty => Err(Error::new(
ErrorKind::UnsupportedType,
"keyword arguments need protocol 4",
)),
event => self.value(event, state, in_key),
}
}
}
impl Writer {
fn finish(mut self) -> Result<Vec<u8>, Error> {
if !self.started || !self.stack.is_empty() || self.skip > 0 {
return Err(Error::new(ErrorKind::InvalidState, "incomplete value"));
}
self.out.push(b'.');
Ok(self.out)
}
fn value_done(&mut self) {
if let Some(open) = self.stack.last_mut()
&& open.is_map
{
open.expects_key = !open.expects_key;
}
}
fn value(&mut self, event: Event<'_>, state: &mut State, in_key: bool) -> Result<(), Error> {
let class = state.event::<ClassData>().and_then(|x| x.0.clone());
let form = state.event::<FormData>().and_then(|x| x.0);
let kind = state.event::<KindData>().and_then(|x| x.0);
let shared = state.event::<SharedIdData>().and_then(|x| x.0);
if let Some(idx) = shared.and_then(|id| self.memo.get(&id).copied()) {
self.get(idx);
match event {
Event::MapStart(_) | Event::SeqStart(_) => self.skip = 1,
_ => self.value_done(),
}
return Ok(());
}
let Some(class) = class else {
return match event {
Event::Atom(atom) => {
self.atom(atom, kind)?;
self.value_done();
Ok(())
}
Event::MapStart(_) => {
if in_key {
return Err(Error::new(
ErrorKind::UnsupportedType,
"maps cannot be keys of dicts",
));
}
self.out.push(b'}');
self.memoize(shared);
self.out.push(b'(');
self.open(End::SetItems, Suffix::default(), None, true, false);
Ok(())
}
Event::SeqStart(_) => self.open_seq(kind, in_key, shared, Suffix::default(), None),
Event::MapEnd | Event::SeqEnd => {
Err(Error::new(ErrorKind::InvalidState, "unexpected end event"))
}
};
};
self.global(&class)?;
let form = form.unwrap_or(match event {
Event::MapStart(_) => Form::State,
Event::SeqStart(_) => match kind {
None => Form::Items,
Some(Kind::Tuple) => Form::Arguments,
Some(_) => Form::Argument,
},
_ => Form::Argument,
});
match event {
Event::Atom(atom) => {
match form {
Form::State | Form::Slots => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
self.atom(atom, kind)?;
self.out.push(b'b');
}
_ => {
self.atom(atom, kind)?;
self.out.extend_from_slice(b"\x85R");
self.memoize(shared);
}
}
self.value_done();
Ok(())
}
Event::MapStart(_) => {
let (end, suffix, late) = match form {
Form::State => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
self.out.push(b'}');
(End::SetItems, Suffix::new(b"b"), None)
}
Form::Slots => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
self.out.extend_from_slice(b"N}");
(End::SetItems, Suffix::new(b"\x86b"), None)
}
Form::Items => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
(End::SetItems, Suffix::default(), None)
}
Form::Arguments if self.proto >= 4 => {
self.out.extend_from_slice(b")}");
(End::SetItems, Suffix::new(b"\x92"), shared)
}
Form::Arguments => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
(End::Empty, Suffix::default(), None)
}
Form::Argument => {
self.out.push(b'}');
(End::SetItems, Suffix::new(b"\x85R"), shared)
}
};
if end != End::Empty {
self.out.push(b'(');
}
self.open(end, suffix, late, true, false);
Ok(())
}
Event::SeqStart(_) => match form {
Form::State | Form::Slots => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
self.open_seq(kind, false, None, Suffix::new(b"b"), None)
}
Form::Items => {
self.out.extend_from_slice(b")\x81");
self.memoize(shared);
self.out.push(b'(');
let end = match kind {
Some(Kind::Set | Kind::FrozenSet) => End::AddItems,
_ => End::Appends,
};
self.open(end, Suffix::default(), None, false, false);
Ok(())
}
Form::Arguments => {
self.out.push(b'(');
self.open(End::Tuple, Suffix::new(b"R"), shared, false, false);
Ok(())
}
Form::Argument => self.open_seq(kind, false, None, Suffix::new(b"\x85R"), shared),
},
Event::MapEnd | Event::SeqEnd => {
Err(Error::new(ErrorKind::InvalidState, "unexpected end event"))
}
}
}
fn open_seq(
&mut self,
kind: Option<Kind>,
in_key: bool,
memo: Option<u64>,
suffix: Suffix,
late: Option<u64>,
) -> Result<(), Error> {
let (end, suffix, late) = match (kind, in_key) {
(Some(Kind::Tuple), _) | (None, true) => (End::Tuple, suffix, memo.or(late)),
(Some(Kind::FrozenSet), _) | (Some(Kind::Set), true) => {
if self.proto >= 4 {
(End::FrozenSet, suffix, memo.or(late))
} else {
self.global(&Global::new(self.builtins(), "frozenset"))?;
(
End::List,
Suffix::new(b"\x85R").with(suffix.as_slice()),
memo.or(late),
)
}
}
(Some(Kind::Set), false) => {
if self.proto >= 4 {
self.out.push(0x8f);
self.memoize(memo);
(End::AddItems, suffix, late)
} else {
self.global(&Global::new(self.builtins(), "set"))?;
(
End::List,
Suffix::new(b"\x85R").with(suffix.as_slice()),
memo.or(late),
)
}
}
_ => {
self.out.push(b']');
self.memoize(memo);
(End::Appends, suffix, late)
}
};
self.out.push(b'(');
self.open(end, suffix, late, false, in_key);
Ok(())
}
fn open(&mut self, end: End, suffix: Suffix, memoize: Option<u64>, is_map: bool, in_key: bool) {
self.stack.push(Open {
end,
suffix,
memoize,
is_map,
expects_key: true,
in_key,
});
}
fn close(&mut self) -> Result<(), Error> {
let open = self.stack.pop().unwrap();
self.out.extend_from_slice(match open.end {
End::Appends => b"e",
End::SetItems => b"u",
End::AddItems => b"\x90",
End::Tuple => b"t",
End::FrozenSet => b"\x91",
End::List => b"l",
End::Empty => b"",
});
self.out.extend_from_slice(open.suffix.as_slice());
self.memoize(open.memoize);
self.value_done();
Ok(())
}
fn builtins(&self) -> &'static str {
match self.proto {
2 => "__builtin__",
_ => "builtins",
}
}
fn memoize(&mut self, id: Option<u64>) {
let Some(id) = id else {
return;
};
let idx = self.memo_len;
self.memo_len += 1;
self.memo.insert(id, idx);
if self.proto >= 4 {
self.out.push(0x94);
} else if idx < 256 {
self.out.push(b'q');
self.out.push(idx as u8);
} else {
self.out.push(b'r');
self.out.extend_from_slice(&idx.to_le_bytes());
}
}
fn get(&mut self, idx: u32) {
if idx < 256 {
self.out.push(b'h');
self.out.push(idx as u8);
} else {
self.out.push(b'j');
self.out.extend_from_slice(&idx.to_le_bytes());
}
}
fn global(&mut self, global: &Global) -> Result<(), Error> {
let (mut module, mut name) = (global.module(), global.name());
if self.proto < 3
&& let Some((new_module, new_name)) = compat::reverse_fix_import(module, name)
{
module = new_module;
name = new_name.unwrap_or(name);
}
if self.proto >= 4 {
self.string(module)?;
self.string(name)?;
self.out.push(0x93);
} else {
let valid = |x: &str| !x.is_empty() && !x.contains('\n');
if !valid(module) || !valid(name) {
return Err(Error::new(
ErrorKind::InvalidValue,
"names of globals must not be empty or contain newlines before protocol 4",
));
}
self.out.push(b'c');
self.out.extend_from_slice(module.as_bytes());
self.out.push(b'\n');
self.out.extend_from_slice(name.as_bytes());
self.out.push(b'\n');
}
Ok(())
}
fn atom(&mut self, atom: Atom<'_>, kind: Option<Kind>) -> Result<(), Error> {
match atom {
Atom::Null => self.out.push(b'N'),
Atom::Bool(value) => self.out.push(if value { 0x88 } else { 0x89 }),
Atom::U64(value) => self.int(value.into()),
Atom::I64(value) => self.int(value.into()),
Atom::F32(value) => self.float(value.into()),
Atom::F64(value) => self.float(value),
Atom::Str(value) | Atom::Lexical(value) => self.string(value.as_str())?,
Atom::Char(value) => self.string(value.encode_utf8(&mut [0; 4]))?,
Atom::Bytes(value) => match kind {
Some(Kind::ByteArray) => self.bytearray(value.data())?,
_ => self.bytes(value.data())?,
},
Atom::Implicit(value) => return self.atom(value.value().to_atom(), kind),
Atom::Ext(ref ext) => return self.ext(ext, kind),
_ => return Err(Error::new(ErrorKind::UnsupportedType, "unknown atom")),
}
Ok(())
}
#[cold]
fn ext(&mut self, ext: &ExtValue<'_>, kind: Option<Kind>) -> Result<(), Error> {
if let Some(reference) = ext.downcast_ref::<Reference>() {
return match self.memo.get(&reference.id()) {
Some(&idx) => {
self.get(idx);
Ok(())
}
None => Err(Error::new(
ErrorKind::InvalidValue,
"reference to a value that was not written or a tuple that contains itself",
)),
};
}
if let Some(global) = ext.downcast_ref::<Global>() {
return self.global(global);
}
if let Some(&value) = ext.downcast_ref::<u128>() {
self.long(&BigInt {
negative: false,
magnitude: value.to_be_bytes().to_vec(),
});
return Ok(());
}
if let Some(&value) = ext.downcast_ref::<i128>() {
self.long(&BigInt {
negative: value < 0,
magnitude: value.unsigned_abs().to_be_bytes().to_vec(),
});
return Ok(());
}
if let Some(value) = ext.downcast_ref::<BigInt>() {
self.long(value);
return Ok(());
}
if let Some(value) = ext.downcast_value_ref::<Number>() {
return match value.as_str().parse::<i64>() {
Ok(value) => {
self.int(value.into());
Ok(())
}
Err(_) => match value.as_str().parse::<BigInt>() {
Ok(value) => {
self.long(&value);
Ok(())
}
Err(_) => {
self.float(value.value());
Ok(())
}
},
};
}
match ext.fallback() {
Atom::Ext(_) => Err(Error::new(
ErrorKind::UnsupportedType,
format!("pickle does not support {}", ext.name()),
)),
fallback => self.atom(fallback, kind),
}
}
fn int(&mut self, value: i128) {
if let Ok(value) = u8::try_from(value) {
self.out.push(b'K');
self.out.push(value);
} else if let Ok(value) = u16::try_from(value) {
self.out.push(b'M');
self.out.extend_from_slice(&value.to_le_bytes());
} else if let Ok(value) = i32::try_from(value) {
self.out.push(b'J');
self.out.extend_from_slice(&value.to_le_bytes());
} else {
self.long(&BigInt {
negative: value < 0,
magnitude: value.unsigned_abs().to_be_bytes().to_vec(),
});
}
}
fn long(&mut self, value: &BigInt) {
let magnitude = value.significant_magnitude();
let mut bytes: Vec<u8> = magnitude.iter().rev().copied().collect();
bytes.push(0);
if value.is_negative() {
for byte in bytes.iter_mut() {
*byte = !*byte;
}
for byte in bytes.iter_mut() {
let (sum, overflow) = byte.overflowing_add(1);
*byte = sum;
if !overflow {
break;
}
}
}
while bytes.len() > 1 {
let last = bytes[bytes.len() - 1];
let before = bytes[bytes.len() - 2];
if (last == 0 && before & 0x80 == 0) || (last == 0xff && before & 0x80 != 0) {
bytes.pop();
} else {
break;
}
}
if value.is_zero() {
bytes.clear();
}
if bytes.len() < 256 {
self.out.push(0x8a);
self.out.push(bytes.len() as u8);
} else {
self.out.push(0x8b);
self.out
.extend_from_slice(&(bytes.len() as u32).to_le_bytes());
}
self.out.extend_from_slice(&bytes);
}
fn float(&mut self, value: f64) {
self.out.push(b'G');
self.out.extend_from_slice(&value.to_be_bytes());
}
fn string(&mut self, value: &str) -> Result<(), Error> {
let len = value.len();
if self.proto >= 4 && len < 256 {
self.out.push(0x8c);
self.out.push(len as u8);
} else if let Ok(len) = u32::try_from(len) {
self.out.push(b'X');
self.out.extend_from_slice(&len.to_le_bytes());
} else if self.proto >= 4 {
self.out.push(0x8d);
self.out.extend_from_slice(&(len as u64).to_le_bytes());
} else {
return Err(too_large());
}
self.out.extend_from_slice(value.as_bytes());
Ok(())
}
fn bytes(&mut self, value: &[u8]) -> Result<(), Error> {
if self.proto < 3 {
if value.is_empty() {
self.out.extend_from_slice(b"c__builtin__\nbytes\n)R");
} else {
self.out.extend_from_slice(b"c_codecs\nencode\n");
self.string(&latin1(value))?;
self.string("latin1")?;
self.out.extend_from_slice(b"\x86R");
}
return Ok(());
}
let len = value.len();
if len < 256 {
self.out.push(b'C');
self.out.push(len as u8);
} else if let Ok(len) = u32::try_from(len) {
self.out.push(b'B');
self.out.extend_from_slice(&len.to_le_bytes());
} else if self.proto >= 4 {
self.out.push(0x8e);
self.out.extend_from_slice(&(len as u64).to_le_bytes());
} else {
return Err(too_large());
}
self.out.extend_from_slice(value);
Ok(())
}
fn bytearray(&mut self, value: &[u8]) -> Result<(), Error> {
match self.proto {
5.. => {
self.out.push(0x96);
self.out
.extend_from_slice(&(value.len() as u64).to_le_bytes());
self.out.extend_from_slice(value);
}
3 | 4 => {
self.global(&Global::new("builtins", "bytearray"))?;
self.bytes(value)?;
self.out.extend_from_slice(b"\x85R");
}
_ => {
self.out.extend_from_slice(b"c__builtin__\nbytearray\n");
self.string(&latin1(value))?;
self.string("latin-1")?;
self.out.extend_from_slice(b"\x86R");
}
}
Ok(())
}
}
fn latin1(bytes: &[u8]) -> String {
bytes.iter().map(|&c| char::from(c)).collect()
}
#[cold]
fn too_large() -> Error {
Error::new(
ErrorKind::OutOfRange,
"value too large for the pickle protocol",
)
}