use ahash::AHashSet;
use std::marker::PhantomData;
use pyo3::exceptions::{PyTypeError, PyValueError};
use pyo3::ffi;
use pyo3::prelude::*;
use pyo3::types::{PyBool, PyDict, PyList, PyString};
use pyo3::ToPyObject;
use smallvec::SmallVec;
use crate::errors::{json_err, json_error, JsonError, JsonResult, DEFAULT_RECURSION_LIMIT};
use crate::number_decoder::{AbstractNumberDecoder, NumberAny, NumberRange};
use crate::parse::{Parser, Peek};
use crate::py_lossless_float::{get_decimal_type, FloatMode};
use crate::py_string_cache::{StringCacheAll, StringCacheKeys, StringCacheMode, StringMaybeCache, StringNoCache};
use crate::string_decoder::{StringDecoder, Tape};
use crate::{JsonErrorType, LosslessFloat};
#[derive(Default)]
#[allow(clippy::struct_excessive_bools)]
pub struct PythonParse {
pub allow_inf_nan: bool,
pub cache_mode: StringCacheMode,
pub partial_mode: PartialMode,
pub catch_duplicate_keys: bool,
pub float_mode: FloatMode,
}
impl PythonParse {
pub fn python_parse<'py>(self, py: Python<'py>, json_data: &[u8]) -> JsonResult<Bound<'py, PyAny>> {
macro_rules! ppp {
($string_cache:ident, $key_check:ident, $parse_number:ident) => {
PythonParser::<$string_cache, $key_check, $parse_number>::parse(
py,
json_data,
self.allow_inf_nan,
self.partial_mode,
)
};
}
macro_rules! ppp_group {
($string_cache:ident) => {
match (self.catch_duplicate_keys, self.float_mode) {
(true, FloatMode::Float) => ppp!($string_cache, DuplicateKeyCheck, ParseNumberLossy),
(true, FloatMode::Decimal) => ppp!($string_cache, DuplicateKeyCheck, ParseNumberDecimal),
(true, FloatMode::LosslessFloat) => ppp!($string_cache, DuplicateKeyCheck, ParseNumberLossless),
(false, FloatMode::Float) => ppp!($string_cache, NoopKeyCheck, ParseNumberLossy),
(false, FloatMode::Decimal) => ppp!($string_cache, NoopKeyCheck, ParseNumberDecimal),
(false, FloatMode::LosslessFloat) => ppp!($string_cache, NoopKeyCheck, ParseNumberLossless),
}
};
}
match self.cache_mode {
StringCacheMode::All => ppp_group!(StringCacheAll),
StringCacheMode::Keys => ppp_group!(StringCacheKeys),
StringCacheMode::None => ppp_group!(StringNoCache),
}
}
}
pub fn map_json_error(json_data: &[u8], json_error: &JsonError) -> PyErr {
PyValueError::new_err(json_error.description(json_data))
}
struct PythonParser<'j, StringCache, KeyCheck, ParseNumber> {
_string_cache: PhantomData<StringCache>,
_key_check: PhantomData<KeyCheck>,
_parse_number: PhantomData<ParseNumber>,
parser: Parser<'j>,
tape: Tape,
recursion_limit: u8,
allow_inf_nan: bool,
partial_mode: PartialMode,
}
impl<'j, StringCache: StringMaybeCache, KeyCheck: MaybeKeyCheck, ParseNumber: MaybeParseNumber>
PythonParser<'j, StringCache, KeyCheck, ParseNumber>
{
fn parse<'py>(
py: Python<'py>,
json_data: &[u8],
allow_inf_nan: bool,
partial_mode: PartialMode,
) -> JsonResult<Bound<'py, PyAny>> {
let mut slf = PythonParser {
_string_cache: PhantomData::<StringCache>,
_key_check: PhantomData::<KeyCheck>,
_parse_number: PhantomData::<ParseNumber>,
parser: Parser::new(json_data),
tape: Tape::default(),
recursion_limit: DEFAULT_RECURSION_LIMIT,
allow_inf_nan,
partial_mode,
};
let peek = slf.parser.peek()?;
let v = slf.py_take_value(py, peek)?;
if !slf.partial_mode.is_active() {
slf.parser.finish()?;
}
Ok(v)
}
fn py_take_value<'py>(&mut self, py: Python<'py>, peek: Peek) -> JsonResult<Bound<'py, PyAny>> {
match peek {
Peek::Null => {
self.parser.consume_null()?;
Ok(py.None().into_bound(py))
}
Peek::True => {
self.parser.consume_true()?;
Ok(true.to_object(py).into_bound(py))
}
Peek::False => {
self.parser.consume_false()?;
Ok(false.to_object(py).into_bound(py))
}
Peek::String => {
let s = self
.parser
.consume_string::<StringDecoder>(&mut self.tape, self.partial_mode.allow_trailing_str())?;
Ok(StringCache::get_value(py, s.as_str(), s.ascii_only()).into_any())
}
Peek::Array => {
let peek_first = match self.parser.array_first() {
Ok(Some(peek)) => peek,
Err(e) if !self._allow_partial_err(&e) => return Err(e),
Ok(None) | Err(_) => return Ok(PyList::empty_bound(py).into_any()),
};
let mut vec: SmallVec<[Bound<'_, PyAny>; 8]> = SmallVec::with_capacity(8);
if let Err(e) = self._parse_array(py, peek_first, &mut vec) {
if !self._allow_partial_err(&e) {
return Err(e);
}
}
Ok(PyList::new_bound(py, vec).into_any())
}
Peek::Object => {
let dict = PyDict::new_bound(py);
if let Err(e) = self._parse_object(py, &dict) {
if !self._allow_partial_err(&e) {
return Err(e);
}
}
Ok(dict.into_any())
}
_ => ParseNumber::parse_number(py, &mut self.parser, peek, self.allow_inf_nan),
}
}
fn _parse_array<'py>(
&mut self,
py: Python<'py>,
peek_first: Peek,
vec: &mut SmallVec<[Bound<'py, PyAny>; 8]>,
) -> JsonResult<()> {
let v = self._check_take_value(py, peek_first)?;
vec.push(v);
while let Some(peek) = self.parser.array_step()? {
let v = self._check_take_value(py, peek)?;
vec.push(v);
}
Ok(())
}
fn _parse_object<'py>(&mut self, py: Python<'py>, dict: &Bound<'py, PyDict>) -> JsonResult<()> {
let set_item = |key: Bound<'py, PyString>, value: Bound<'py, PyAny>| {
let r = unsafe { ffi::PyDict_SetItem(dict.as_ptr(), key.as_ptr(), value.as_ptr()) };
assert_ne!(r, -1, "PyDict_SetItem failed");
};
let mut check_keys = KeyCheck::default();
if let Some(first_key) = self.parser.object_first::<StringDecoder>(&mut self.tape)? {
let first_key_s = first_key.as_str();
check_keys.check(first_key_s, self.parser.index)?;
let first_key = StringCache::get_key(py, first_key_s, first_key.ascii_only());
let peek = self.parser.peek()?;
let first_value = self._check_take_value(py, peek)?;
set_item(first_key, first_value);
while let Some(key) = self.parser.object_step::<StringDecoder>(&mut self.tape)? {
let key_s = key.as_str();
check_keys.check(key_s, self.parser.index)?;
let key = StringCache::get_key(py, key_s, key.ascii_only());
let peek = self.parser.peek()?;
let value = self._check_take_value(py, peek)?;
set_item(key, value);
}
}
Ok(())
}
fn _allow_partial_err(&self, e: &JsonError) -> bool {
if self.partial_mode.is_active() {
matches!(
e.error_type,
JsonErrorType::EofWhileParsingList
| JsonErrorType::EofWhileParsingObject
| JsonErrorType::EofWhileParsingString
| JsonErrorType::EofWhileParsingValue
| JsonErrorType::ExpectedListCommaOrEnd
| JsonErrorType::ExpectedObjectCommaOrEnd
)
} else {
false
}
}
fn _check_take_value<'py>(&mut self, py: Python<'py>, peek: Peek) -> JsonResult<Bound<'py, PyAny>> {
self.recursion_limit = match self.recursion_limit.checked_sub(1) {
Some(limit) => limit,
None => return json_err!(RecursionLimitExceeded, self.parser.index),
};
let r = self.py_take_value(py, peek);
self.recursion_limit += 1;
r
}
}
#[derive(Debug, Clone, Copy)]
pub enum PartialMode {
Off,
On,
TrailingStrings,
}
impl Default for PartialMode {
fn default() -> Self {
Self::Off
}
}
const PARTIAL_ERROR: &str = "Invalid partial mode, should be `'off'`, `'on'`, `'trailing-strings'` or a `bool`";
impl<'py> FromPyObject<'py> for PartialMode {
fn extract_bound(ob: &Bound<'py, PyAny>) -> PyResult<Self> {
if let Ok(bool_mode) = ob.downcast::<PyBool>() {
Ok(bool_mode.is_true().into())
} else if let Ok(str_mode) = ob.extract::<&str>() {
match str_mode {
"off" => Ok(Self::Off),
"on" => Ok(Self::On),
"trailing-strings" => Ok(Self::TrailingStrings),
_ => Err(PyValueError::new_err(PARTIAL_ERROR)),
}
} else {
Err(PyTypeError::new_err(PARTIAL_ERROR))
}
}
}
impl From<bool> for PartialMode {
fn from(mode: bool) -> Self {
if mode {
Self::On
} else {
Self::Off
}
}
}
impl PartialMode {
fn is_active(self) -> bool {
!matches!(self, Self::Off)
}
fn allow_trailing_str(self) -> bool {
matches!(self, Self::TrailingStrings)
}
}
trait MaybeKeyCheck: Default {
fn check(&mut self, key: &str, index: usize) -> JsonResult<()>;
}
#[derive(Default)]
struct NoopKeyCheck;
impl MaybeKeyCheck for NoopKeyCheck {
fn check(&mut self, _key: &str, _index: usize) -> JsonResult<()> {
Ok(())
}
}
#[derive(Default)]
struct DuplicateKeyCheck(AHashSet<String>);
impl MaybeKeyCheck for DuplicateKeyCheck {
fn check(&mut self, key: &str, index: usize) -> JsonResult<()> {
if self.0.insert(key.to_owned()) {
Ok(())
} else {
Err(JsonError::new(JsonErrorType::DuplicateKey(key.to_owned()), index))
}
}
}
trait MaybeParseNumber {
fn parse_number<'py>(
py: Python<'py>,
parser: &mut Parser,
peek: Peek,
allow_inf_nan: bool,
) -> JsonResult<Bound<'py, PyAny>>;
}
struct ParseNumberLossy;
impl MaybeParseNumber for ParseNumberLossy {
fn parse_number<'py>(
py: Python<'py>,
parser: &mut Parser,
peek: Peek,
allow_inf_nan: bool,
) -> JsonResult<Bound<'py, PyAny>> {
match parser.consume_number::<NumberAny>(peek.into_inner(), allow_inf_nan) {
Ok(number) => Ok(number.to_object(py).into_bound(py)),
Err(e) => {
if !peek.is_num() {
Err(json_error!(ExpectedSomeValue, parser.index))
} else {
Err(e)
}
}
}
}
}
struct ParseNumberLossless;
impl MaybeParseNumber for ParseNumberLossless {
fn parse_number<'py>(
py: Python<'py>,
parser: &mut Parser,
peek: Peek,
allow_inf_nan: bool,
) -> JsonResult<Bound<'py, PyAny>> {
match parser.consume_number::<NumberRange>(peek.into_inner(), allow_inf_nan) {
Ok(number_range) => {
let bytes = parser.slice(number_range.range).unwrap();
let obj = if number_range.is_int {
NumberAny::decode(bytes, 0, peek.into_inner(), allow_inf_nan)?
.0
.to_object(py)
} else {
LosslessFloat::new_unchecked(bytes.to_vec()).into_py(py)
};
Ok(obj.into_bound(py))
}
Err(e) => {
if !peek.is_num() {
Err(json_error!(ExpectedSomeValue, parser.index))
} else {
Err(e)
}
}
}
}
}
struct ParseNumberDecimal;
impl MaybeParseNumber for ParseNumberDecimal {
fn parse_number<'py>(
py: Python<'py>,
parser: &mut Parser,
peek: Peek,
allow_inf_nan: bool,
) -> JsonResult<Bound<'py, PyAny>> {
match parser.consume_number::<NumberRange>(peek.into_inner(), allow_inf_nan) {
Ok(number_range) => {
let bytes = parser.slice(number_range.range).unwrap();
if number_range.is_int {
let obj = NumberAny::decode(bytes, 0, peek.into_inner(), allow_inf_nan)?
.0
.to_object(py);
Ok(obj.into_bound(py))
} else {
let decimal_type = get_decimal_type(py)
.map_err(|e| JsonError::new(JsonErrorType::InternalError(e.to_string()), parser.index))?;
let float_str = unsafe { std::str::from_utf8_unchecked(bytes) };
decimal_type
.call1((float_str,))
.map_err(|e| JsonError::new(JsonErrorType::InternalError(e.to_string()), parser.index))
}
}
Err(e) => {
if !peek.is_num() {
Err(json_error!(ExpectedSomeValue, parser.index))
} else {
Err(e)
}
}
}
}
}