rust/pgsql: convert to nom 8

Ticket: #8039
pull/14207/head
Jason Ish 9 months ago
parent d25286e77a
commit dcfe247467

@ -62,6 +62,46 @@ pub mod nom7 {
}
}
pub mod nom8 {
use nom8::bytes::streaming::{tag, take_until};
use nom8::error::{Error, ParseError};
use nom8::ErrorConvert;
use nom8::{IResult, Parser};
/// Reimplementation of `take_until_and_consume` for nom 8
///
/// `take_until` does not consume the matched tag, and
/// `take_until_and_consume` was removed in nom 7. This function
/// provides an implementation (specialized for `&[u8]`).
pub fn take_until_and_consume<'a, E: ParseError<&'a [u8]>>(
t: &'a [u8],
) -> impl FnMut(&'a [u8]) -> IResult<&'a [u8], &'a [u8], E> {
move |i: &'a [u8]| {
let (i, res) = take_until(t).parse(i)?;
let (i, _) = tag(t).parse(i)?;
Ok((i, res))
}
}
/// Specialized version of the nom 8 `bits` combinator
///
/// The `bits combinator has trouble inferring the transient error type
/// used by the tuple parser, because the function is generic and any
/// error type would be valid.
/// Use an explicit error type (as described in
/// https://docs.rs/nom/7.1.0/nom/bits/fn.bits.html) to solve this problem, and
/// specialize this function for `&[u8]`.
pub fn bits<'a, O, E, P>(parser: P) -> impl FnMut(&'a [u8]) -> IResult<&'a [u8], O, E>
where
E: ParseError<&'a [u8]>,
Error<(&'a [u8], usize)>: ErrorConvert<E>,
P: FnMut((&'a [u8], usize)) -> IResult<(&'a [u8], usize), O, Error<(&'a [u8], usize)>>,
{
// use full path to disambiguate nom `bits` from this current function name
nom8::bits::bits(parser)
}
}
/// Convert a String to C-compatible string
///
/// This function will consume the provided data and use the underlying bytes to construct a new

@ -19,17 +19,17 @@
//! PostgreSQL nom parsers
use crate::common::nom7::take_until_and_consume;
use nom7::branch::alt;
use nom7::bytes::streaming::{tag, take, take_until, take_until1};
use nom7::character::streaming::{alphanumeric1, char};
use nom7::combinator::{all_consuming, cond, eof, map_parser, opt, peek, verify};
use nom7::error::{make_error, ErrorKind, ParseError};
use nom7::multi::{many1, many_m_n, many_till};
use nom7::number::streaming::{be_i16, be_i32};
use nom7::number::streaming::{be_u16, be_u32, be_u8};
use nom7::sequence::terminated;
use nom7::{Err, IResult, ToUsize};
use crate::common::nom8::take_until_and_consume;
use nom8::branch::alt;
use nom8::bytes::streaming::{tag, take, take_until, take_until1};
use nom8::character::streaming::{alphanumeric1, char};
use nom8::combinator::{all_consuming, cond, eof, map_parser, opt, peek, verify};
use nom8::error::{make_error, ErrorKind, ParseError};
use nom8::multi::{many1, many_m_n, many_till};
use nom8::number::streaming::{be_i16, be_i32};
use nom8::number::streaming::{be_u16, be_u32, be_u8};
use nom8::sequence::terminated;
use nom8::{Err, IResult, Parser, ToUsize};
const PGSQL_LENGTH_FIELD: u32 = 4;
@ -56,12 +56,12 @@ impl<I> ParseError<I> for PgsqlParseError<I> {
}
fn parse_gte_length(i: &[u8], expected_length: u32) -> IResult<&[u8], u32, PgsqlParseError<&[u8]>> {
let res = verify(be_u32::<&[u8], nom7::error::Error<_>>, |&x| {
let res = verify(be_u32::<&[u8], nom8::error::Error<_>>, |&x| {
x >= expected_length
})(i);
}).parse(i);
match res {
Ok(result) => Ok((result.0, result.1)),
Err(nom7::Err::Incomplete(needed)) => Err(Err::Incomplete(needed)),
Err(nom8::Err::Incomplete(needed)) => Err(Err::Incomplete(needed)),
Err(_) => Err(Err::Error(PgsqlParseError::InvalidLength)),
}
}
@ -69,12 +69,12 @@ fn parse_gte_length(i: &[u8], expected_length: u32) -> IResult<&[u8], u32, Pgsql
fn parse_exact_length(
i: &[u8], expected_length: u32,
) -> IResult<&[u8], u32, PgsqlParseError<&[u8]>> {
let res = verify(be_u32::<&[u8], nom7::error::Error<_>>, |&x| {
let res = verify(be_u32::<&[u8], nom8::error::Error<_>>, |&x| {
x == expected_length
})(i);
}).parse(i);
match res {
Ok(result) => Ok((result.0, result.1)),
Err(nom7::Err::Incomplete(needed)) => Err(Err::Incomplete(needed)),
Err(nom8::Err::Incomplete(needed)) => Err(Err::Incomplete(needed)),
Err(_) => Err(Err::Error(PgsqlParseError::InvalidLength)),
}
}
@ -594,10 +594,10 @@ impl From<u8> for PgsqlErrorNoticeFieldType {
fn pgsql_parse_generic_parameter(
i: &[u8],
) -> IResult<&[u8], PgsqlParameter, PgsqlParseError<&[u8]>> {
let (i, param_name) = take_until1("\x00")(i)?;
let (i, _) = tag("\x00")(i)?;
let (i, param_value) = take_until("\x00")(i)?;
let (i, _) = tag("\x00")(i)?;
let (i, param_name) = take_until1(b"\x00".as_ref()).parse(i)?;
let (i, _) = tag(b"\x00".as_ref()).parse(i)?;
let (i, param_value) = take_until(b"\x00".as_ref()).parse(i)?;
let (i, _) = tag(b"\x00".as_ref()).parse(i)?;
Ok((
i,
PgsqlParameter {
@ -612,8 +612,8 @@ fn pgsql_parse_startup_parameters(
) -> IResult<&[u8], PgsqlStartupParameters, PgsqlParseError<&[u8]>> {
let (i, mut optional) = opt(terminated(
many1(pgsql_parse_generic_parameter),
tag("\x00"),
))(i)?;
tag(b"\x00".as_ref()),
)).parse(i)?;
if let Some(ref mut params) = optional {
let mut user = PgsqlParameter {
name: PgsqlParameters::User,
@ -645,22 +645,22 @@ fn parse_sasl_initial_response_payload(
i: &[u8],
) -> IResult<&[u8], SASLInitialResponse, PgsqlParseError<&[u8]>> {
let (i, sasl_mechanism) = parse_sasl_mechanism(i)?;
let (i, param_length) = be_u32(i)?;
let (i, param_length) = be_u32.parse(i)?;
// From RFC 5802 - the client-first-message will always start w/
// 'n', 'y' or 'p', otherwise it's invalid, I think we should check that, at some point
let (i, param) = terminated(take(param_length), eof)(i)?;
let (i, param) = terminated(take(param_length), eof).parse(i)?;
Ok((i, (sasl_mechanism, param_length, param.to_vec())))
}
pub(crate) fn parse_sasl_initial_response(
i: &[u8],
) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'p')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'p').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, payload) = map_parser(
take(length - PGSQL_LENGTH_FIELD),
parse_sasl_initial_response_payload,
)(i)?;
).parse(i)?;
Ok((
i,
PgsqlFEMessage::SASLInitialResponse(SASLInitialResponsePacket {
@ -674,9 +674,9 @@ pub(crate) fn parse_sasl_initial_response(
}
pub(crate) fn parse_sasl_response(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'p')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'p').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
let resp = PgsqlFEMessage::SASLResponse(RegularPacket {
identifier,
length,
@ -689,12 +689,12 @@ fn pgsql_parse_startup_packet(
i: &[u8],
) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, length) = parse_gte_length(i, 8)?;
let (i, proto_major) = peek(be_u16)(i)?;
let (i, b) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, proto_major) = peek(be_u16).parse(i)?;
let (i, b) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
let (_, message) = match proto_major {
1..=3 => {
let (b, proto_major) = be_u16(b)?;
let (b, proto_minor) = be_u16(b)?;
let (b, proto_major) = be_u16.parse(b)?;
let (b, proto_minor) = be_u16.parse(b)?;
let (b, params) = pgsql_parse_startup_parameters(b)?;
(
b,
@ -707,8 +707,8 @@ fn pgsql_parse_startup_packet(
)
}
PGSQL_DUMMY_PROTO_MAJOR => {
let (b, proto_major) = be_u16(b)?;
let (b, proto_minor) = be_u16(b)?;
let (b, proto_major) = be_u16.parse(b)?;
let (b, proto_minor) = be_u16.parse(b)?;
let (b, message) = match proto_minor {
PGSQL_DUMMY_PROTO_CANCEL_REQUEST => parse_cancel_request(b)?,
PGSQL_DUMMY_PROTO_MINOR_SSL => (
@ -743,9 +743,9 @@ fn pgsql_parse_startup_packet(
// Password can be encrypted or in cleartext
pub(crate) fn parse_password_message(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'p')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'p').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, password) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until1("\x00"))(i)?;
let (i, password) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until1(b"\x00".as_ref())).parse(i)?;
Ok((
i,
PgsqlFEMessage::PasswordMessage(RegularPacket {
@ -757,9 +757,9 @@ pub(crate) fn parse_password_message(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage,
}
fn parse_simple_query(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'Q')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'Q').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, query) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until1("\x00"))(i)?;
let (i, query) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until1(b"\x00".as_ref())).parse(i)?;
Ok((
i,
PgsqlFEMessage::SimpleQuery(RegularPacket {
@ -771,8 +771,8 @@ fn parse_simple_query(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseErro
}
fn parse_cancel_request(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, pid) = be_u32(i)?;
let (i, backend_key) = be_u32(i)?;
let (i, pid) = be_u32.parse(i)?;
let (i, backend_key) = be_u32.parse(i)?;
Ok((
i,
PgsqlFEMessage::CancelRequest(CancelRequestMessage { pid, backend_key }),
@ -780,7 +780,7 @@ fn parse_cancel_request(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseEr
}
fn parse_terminate_message(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'X')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'X').parse(i)?;
let (i, length) = parse_exact_length(i, PGSQL_LENGTH_FIELD)?;
Ok((
i,
@ -790,7 +790,7 @@ fn parse_terminate_message(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlPars
// Messages that begin with 'p' but are not password ones are not parsed here
pub(crate) fn parse_request(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, tag) = peek(be_u8)(i)?;
let (i, tag) = peek(be_u8).parse(i)?;
let (i, message) = match tag {
b'\0' => pgsql_parse_startup_packet(i)?,
b'Q' => parse_simple_query(i)?,
@ -799,9 +799,9 @@ pub(crate) fn parse_request(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlPar
b'c' => parse_copy_in_done(i)?,
b'f' => parse_copy_fail(i)?,
_ => {
let (i, identifier) = be_u8(i)?;
let (i, identifier) = be_u8.parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
let unknown = PgsqlFEMessage::UnknownMessageType(RegularPacket {
identifier,
length,
@ -817,9 +817,9 @@ pub(crate) fn parse_request(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlPar
fn pgsql_parse_authentication_message<'a>(
i: &'a [u8],
) -> IResult<&'a [u8], PgsqlBEMessage, PgsqlParseError<&'a [u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'R')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'R').parse(i)?;
let (i, length) = parse_gte_length(i, 8)?;
let (i, auth_type) = be_u32(i)?;
let (i, auth_type) = be_u32.parse(i)?;
let (i, message) = map_parser(take(length - 8), |b: &'a [u8]| {
match auth_type {
0 => Ok((
@ -841,7 +841,7 @@ fn pgsql_parse_authentication_message<'a>(
}),
)),
5 => {
let (b, salt) = all_consuming(take(4_usize))(b)?;
let (b, salt) = all_consuming(take(4_usize)).parse(b)?;
Ok((
b,
PgsqlBEMessage::AuthenticationMD5Password(AuthenticationMessage {
@ -895,19 +895,19 @@ fn pgsql_parse_authentication_message<'a>(
// TODO add other authentication messages
_ => return Err(Err::Error(make_error(i, ErrorKind::Switch))),
}
})(i)?;
}).parse(i)?;
Ok((i, message))
}
fn parse_parameter_status_message(
i: &[u8],
) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'S')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'S').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, param) = map_parser(
take(length - PGSQL_LENGTH_FIELD),
pgsql_parse_generic_parameter,
)(i)?;
).parse(i)?;
Ok((
i,
PgsqlBEMessage::ParameterStatus(ParameterStatusMessage {
@ -919,7 +919,7 @@ fn parse_parameter_status_message(
}
pub(crate) fn parse_ssl_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, tag) = alt((char('N'), char('S')))(i)?;
let (i, tag) = alt((char('N'), char('S'))).parse(i)?;
Ok((
i,
PgsqlBEMessage::SSLResponse(SSLResponseMessage::from(tag)),
@ -929,10 +929,10 @@ pub(crate) fn parse_ssl_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, Pgs
fn parse_backend_key_data_message(
i: &[u8],
) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'K')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'K').parse(i)?;
let (i, length) = parse_exact_length(i, 12)?;
let (i, pid) = be_u32(i)?;
let (i, secret_key) = be_u32(i)?;
let (i, pid) = be_u32.parse(i)?;
let (i, secret_key) = be_u32.parse(i)?;
Ok((
i,
PgsqlBEMessage::BackendKeyData(BackendKeyDataMessage {
@ -945,9 +945,9 @@ fn parse_backend_key_data_message(
}
fn parse_command_complete(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'C')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'C').parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, payload) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until("\x00"))(i)?;
let (i, payload) = map_parser(take(length - PGSQL_LENGTH_FIELD), take_until(b"\x00".as_ref())).parse(i)?;
Ok((
i,
PgsqlBEMessage::CommandComplete(RegularPacket {
@ -959,9 +959,9 @@ fn parse_command_complete(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParse
}
fn parse_ready_for_query(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'Z')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'Z').parse(i)?;
let (i, length) = parse_exact_length(i, 5)?;
let (i, status) = verify(be_u8, |&x| x == b'I' || x == b'T' || x == b'E')(i)?;
let (i, status) = verify(be_u8, |&x| x == b'I' || x == b'T' || x == b'E').parse(i)?;
Ok((
i,
PgsqlBEMessage::ReadyForQuery(ReadyForQueryMessage {
@ -973,14 +973,14 @@ fn parse_ready_for_query(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseE
}
fn parse_row_field(i: &[u8]) -> IResult<&[u8], RowField, PgsqlParseError<&[u8]>> {
let (i, field_name) = take_until1("\x00")(i)?;
let (i, _) = tag("\x00")(i)?;
let (i, table_oid) = be_u32(i)?;
let (i, column_index) = be_u16(i)?;
let (i, data_type_oid) = be_u32(i)?;
let (i, data_type_size) = be_i16(i)?;
let (i, type_modifier) = be_i32(i)?;
let (i, format_code) = be_u16(i)?;
let (i, field_name) = take_until1(b"\x00".as_ref()).parse(i)?;
let (i, _) = tag(b"\x00".as_ref()).parse(i)?;
let (i, table_oid) = be_u32.parse(i)?;
let (i, column_index) = be_u16.parse(i)?;
let (i, data_type_oid) = be_u32.parse(i)?;
let (i, data_type_size) = be_i16.parse(i)?;
let (i, type_modifier) = be_i32.parse(i)?;
let (i, format_code) = be_u16.parse(i)?;
Ok((
i,
RowField {
@ -996,13 +996,13 @@ fn parse_row_field(i: &[u8]) -> IResult<&[u8], RowField, PgsqlParseError<&[u8]>>
}
fn parse_row_description(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'T')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'T').parse(i)?;
let (i, length) = parse_gte_length(i, 7)?;
let (i, field_count) = be_u16(i)?;
let (i, field_count) = be_u16.parse(i)?;
let (i, fields) = map_parser(
take(length - 6),
many_m_n(0, field_count.into(), parse_row_field),
)(i)?;
).parse(i)?;
Ok((
i,
PgsqlBEMessage::RowDescription(RowDescriptionMessage {
@ -1015,8 +1015,8 @@ fn parse_row_description(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseE
}
fn parse_data_row_value(i: &[u8]) -> IResult<&[u8], ColumnFieldValue, PgsqlParseError<&[u8]>> {
let (i, value_length) = be_i32(i)?;
let (i, value) = cond(value_length >= 0, take(value_length as usize))(i)?;
let (i, value_length) = be_i32.parse(i)?;
let (i, value) = cond(value_length >= 0, take(value_length as usize)).parse(i)?;
Ok((
i,
ColumnFieldValue {
@ -1044,12 +1044,12 @@ fn add_up_data_size(columns: Vec<ColumnFieldValue>) -> u64 {
}
fn parse_copy_out_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'H')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'H').parse(i)?;
// copy out message : identifier (u8), length (u32), format (u8), cols (u16), formats (u16*cols)
let (i, length) = parse_gte_length(i, 8)?;
let (i, _format) = be_u8(i)?;
let (i, columns) = be_u16(i)?;
let (i, _formats) = many_m_n(0, columns.to_usize(), be_u16)(i)?;
let (i, _format) = be_u8.parse(i)?;
let (i, columns) = be_u16.parse(i)?;
let (i, _formats) = many_m_n(0, columns.to_usize(), be_u16).parse(i)?;
Ok((
i,
PgsqlBEMessage::CopyOutResponse(CopyResponse {
@ -1061,11 +1061,11 @@ fn parse_copy_out_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlPars
}
fn parse_copy_in_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'G')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'G').parse(i)?;
let (i, length) = parse_gte_length(i, 8)?;
let (i, _format) = be_u8(i)?;
let (i, columns) = be_u16(i)?;
let (i, _formats) = many_m_n(0, columns.to_usize(), be_u16)(i)?;
let (i, _format) = be_u8.parse(i)?;
let (i, columns) = be_u16.parse(i)?;
let (i, _formats) = many_m_n(0, columns.to_usize(), be_u16).parse(i)?;
Ok((
i,
PgsqlBEMessage::CopyInResponse(CopyResponse {
@ -1077,9 +1077,9 @@ fn parse_copy_in_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParse
}
fn parse_consolidated_copy_data_out(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'd')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'd').parse(i)?;
let (i, length) = parse_gte_length(i, 5)?;
let (i, _data) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, _data) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
SCLogDebug!("data_size is {:?}", _data);
Ok((
i, PgsqlBEMessage::ConsolidatedCopyDataOut(ConsolidatedDataRowPacket {
@ -1090,9 +1090,9 @@ fn parse_consolidated_copy_data_out(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage,
}
fn parse_consolidated_copy_data_in(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'd')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'd').parse(i)?;
let (i, length) = parse_gte_length(i, 5)?;
let (i, _data) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, _data) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
SCLogDebug!("data size is {:?}", _data);
Ok((
i, PgsqlFEMessage::ConsolidatedCopyDataIn(ConsolidatedDataRowPacket {
@ -1103,7 +1103,7 @@ fn parse_consolidated_copy_data_in(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, P
}
fn parse_copy_in_done(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'c')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'c').parse(i)?;
let (i, length) = parse_exact_length(i, PGSQL_LENGTH_FIELD)?;
Ok((
i, PgsqlFEMessage::CopyDone(NoPayloadMessage {
@ -1114,7 +1114,7 @@ fn parse_copy_in_done(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseErro
}
fn parse_copy_out_done(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'c')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'c').parse(i)?;
let (i, length) = parse_exact_length(i, PGSQL_LENGTH_FIELD)?;
Ok((
i, PgsqlBEMessage::CopyDone(NoPayloadMessage {
@ -1125,9 +1125,9 @@ fn parse_copy_out_done(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseErr
}
fn parse_copy_fail(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'f')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'f').parse(i)?;
let (i, length) = parse_gte_length(i, 5)?;
let (i, data) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, data) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
Ok((
i, PgsqlFEMessage::CopyFail(RegularPacket {
identifier,
@ -1143,14 +1143,14 @@ fn parse_copy_fail(i: &[u8]) -> IResult<&[u8], PgsqlFEMessage, PgsqlParseError<&
fn parse_consolidated_data_row(
i: &[u8],
) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'D')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'D').parse(i)?;
let (i, length) = parse_gte_length(i, 7)?;
let (i, field_count) = be_u16(i)?;
let (i, field_count) = be_u16.parse(i)?;
// 6 here is for skipping length + field_count
let (i, rows) = map_parser(
take(length - 6),
many_m_n(0, field_count.into(), parse_data_row_value),
)(i)?;
).parse(i)?;
Ok((
i,
PgsqlBEMessage::ConsolidatedDataRow(ConsolidatedDataRowPacket {
@ -1164,11 +1164,11 @@ fn parse_consolidated_data_row(
fn parse_sasl_mechanism(
i: &[u8],
) -> IResult<&[u8], SASLAuthenticationMechanism, PgsqlParseError<&[u8]>> {
let res: IResult<_, _, ()> = terminated(tag("SCRAM-SHA-256-PLUS"), tag("\x00"))(i);
let res: IResult<_, _, ()> = terminated(tag("SCRAM-SHA-256-PLUS"), tag(b"\x00".as_ref())).parse(i);
if let Ok((i, _)) = res {
return Ok((i, SASLAuthenticationMechanism::ScramSha256Plus));
}
let res: IResult<_, _, ()> = terminated(tag("SCRAM-SHA-256"), tag("\x00"))(i);
let res: IResult<_, _, ()> = terminated(tag("SCRAM-SHA-256"), tag(b"\x00".as_ref())).parse(i);
if let Ok((i, _)) = res {
return Ok((i, SASLAuthenticationMechanism::ScramSha256));
}
@ -1178,14 +1178,14 @@ fn parse_sasl_mechanism(
fn parse_sasl_mechanisms(
i: &[u8],
) -> IResult<&[u8], Vec<SASLAuthenticationMechanism>, PgsqlParseError<&[u8]>> {
terminated(many1(parse_sasl_mechanism), tag("\x00"))(i)
terminated(many1(parse_sasl_mechanism), tag(b"\x00".as_ref())).parse(i)
}
fn parse_error_response_code(
i: &[u8],
) -> IResult<&[u8], PgsqlErrorNoticeMessageField, PgsqlParseError<&[u8]>> {
let (i, _field_type) = char('C')(i)?;
let (i, field_value) = map_parser(take(6_usize), alphanumeric1)(i)?;
let (i, _field_type) = char('C').parse(i)?;
let (i, field_value) = map_parser(take(6_usize), alphanumeric1).parse(i)?;
Ok((
i,
PgsqlErrorNoticeMessageField {
@ -1200,9 +1200,9 @@ fn parse_error_response_code(
fn parse_error_response_severity(
i: &[u8],
) -> IResult<&[u8], PgsqlErrorNoticeMessageField, PgsqlParseError<&[u8]>> {
let (i, field_type) = char('V')(i)?;
let (i, field_value) = alt((tag("ERROR"), tag("FATAL"), tag("PANIC")))(i)?;
let (i, _) = tag("\x00")(i)?;
let (i, field_type) = char('V').parse(i)?;
let (i, field_value) = alt((tag("ERROR"), tag("FATAL"), tag("PANIC"))).parse(i)?;
let (i, _) = tag(b"\x00".as_ref()).parse(i)?;
Ok((
i,
PgsqlErrorNoticeMessageField {
@ -1217,15 +1217,15 @@ fn parse_error_response_severity(
fn parse_notice_response_severity(
i: &[u8],
) -> IResult<&[u8], PgsqlErrorNoticeMessageField, PgsqlParseError<&[u8]>> {
let (i, field_type) = char('V')(i)?;
let (i, field_type) = char('V').parse(i)?;
let (i, field_value) = alt((
tag("WARNING"),
tag("NOTICE"),
tag("DEBUG"),
tag("INFO"),
tag("LOG"),
))(i)?;
let (i, _) = tag("\x00")(i)?;
)).parse(i)?;
let (i, _) = tag(b"\x00".as_ref()).parse(i)?;
Ok((
i,
PgsqlErrorNoticeMessageField {
@ -1238,7 +1238,7 @@ fn parse_notice_response_severity(
fn parse_error_response_field(
i: &[u8], is_err_msg: bool,
) -> IResult<&[u8], PgsqlErrorNoticeMessageField, PgsqlParseError<&[u8]>> {
let (i, field_type) = peek(be_u8)(i)?;
let (i, field_type) = peek(be_u8).parse(i)?;
let (i, data) = match field_type {
b'V' => {
if is_err_msg {
@ -1249,9 +1249,9 @@ fn parse_error_response_field(
}
b'C' => parse_error_response_code(i)?,
_ => {
let (i, field_type) = be_u8(i)?;
let (i, field_value) = take_until("\x00")(i)?;
let (i, _just_tag) = tag("\x00")(i)?;
let (i, field_type) = be_u8.parse(i)?;
let (i, field_value) = take_until(b"\x00".as_ref()).parse(i)?;
let (i, _just_tag) = tag(b"\x00".as_ref()).parse(i)?;
let message = PgsqlErrorNoticeMessageField {
field_type: PgsqlErrorNoticeFieldType::from(field_type),
field_value: field_value.to_vec(),
@ -1265,16 +1265,16 @@ fn parse_error_response_field(
fn parse_error_notice_fields(
i: &[u8], is_err_msg: bool,
) -> IResult<&[u8], Vec<PgsqlErrorNoticeMessageField>, PgsqlParseError<&[u8]>> {
let (i, data) = many_till(|b| parse_error_response_field(b, is_err_msg), tag("\x00"))(i)?;
let (i, data) = many_till(|b| parse_error_response_field(b, is_err_msg), tag(b"\x00".as_ref())).parse(i)?;
Ok((i, data.0))
}
fn pgsql_parse_error_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'E')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'E').parse(i)?;
let (i, length) = parse_gte_length(i, 11)?;
let (i, message_body) = map_parser(take(length - PGSQL_LENGTH_FIELD), |b| {
parse_error_notice_fields(b, true)
})(i)?;
}).parse(i)?;
Ok((
i,
@ -1287,11 +1287,11 @@ fn pgsql_parse_error_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlP
}
fn pgsql_parse_notice_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'N')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'N').parse(i)?;
let (i, length) = parse_gte_length(i, 11)?;
let (i, message_body) = map_parser(take(length - PGSQL_LENGTH_FIELD), |b| {
parse_error_notice_fields(b, false)
})(i)?;
}).parse(i)?;
Ok((
i,
PgsqlBEMessage::NoticeResponse(ErrorNoticeMessage {
@ -1303,15 +1303,15 @@ fn pgsql_parse_notice_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, Pgsql
}
fn parse_notification_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, identifier) = verify(be_u8, |&x| x == b'A')(i)?;
let (i, identifier) = verify(be_u8, |&x| x == b'A').parse(i)?;
// length (u32) + pid (u32) + at least one byte, for we have two str fields
let (i, length) = parse_gte_length(i, 10)?;
let (i, data) = map_parser(take(length - PGSQL_LENGTH_FIELD), |b| {
let (b, pid) = be_u32(b)?;
let (b, channel_name) = take_until_and_consume(b"\x00")(b)?;
let (b, payload) = take_until_and_consume(b"\x00")(b)?;
let (b, pid) = be_u32.parse(b)?;
let (b, channel_name) = take_until_and_consume(b"\x00").parse(b)?;
let (b, payload) = take_until_and_consume(b"\x00").parse(b)?;
Ok((b, (pid, channel_name, payload)))
})(i)?;
}).parse(i)?;
let msg = PgsqlBEMessage::NotificationResponse(NotificationResponse {
identifier,
length,
@ -1323,7 +1323,7 @@ fn parse_notification_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, Pgsql
}
pub(crate) fn pgsql_parse_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, PgsqlParseError<&[u8]>> {
let (i, tag) = peek(be_u8)(i)?;
let (i, tag) = peek(be_u8).parse(i)?;
let (i, message) = match tag {
b'E' => pgsql_parse_error_response(i)?,
b'K' => parse_backend_key_data_message(i)?,
@ -1340,9 +1340,9 @@ pub(crate) fn pgsql_parse_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, P
b'H' => parse_copy_out_response(i)?,
b'G' => parse_copy_in_response(i)?,
_ => {
let (i, identifier) = be_u8(i)?;
let (i, identifier) = be_u8.parse(i)?;
let (i, length) = parse_gte_length(i, PGSQL_LENGTH_FIELD)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD)(i)?;
let (i, payload) = take(length - PGSQL_LENGTH_FIELD).parse(i)?;
let unknown = PgsqlBEMessage::UnknownMessageType(RegularPacket {
identifier,
length,
@ -1359,7 +1359,7 @@ pub(crate) fn pgsql_parse_response(i: &[u8]) -> IResult<&[u8], PgsqlBEMessage, P
mod tests {
use super::*;
use nom7::Needed;
use nom8::Needed;
impl ErrorNoticeMessage {
fn new(identifier: u8, length: u32) -> Self {

@ -26,7 +26,7 @@ use crate::conf::*;
use crate::core::{ALPROTO_FAILED, ALPROTO_UNKNOWN, IPPROTO_TCP, *};
use crate::direction::Direction;
use crate::flow::Flow;
use nom7::{Err, IResult};
use nom8::{Err, IResult};
use std;
use std::collections::VecDeque;
use std::ffi::CString;
@ -475,7 +475,7 @@ impl PgsqlState {
return AppLayerResult::err();
}
PgsqlParseError::NomError(_i, error_kind) => {
if error_kind == nom7::error::ErrorKind::Switch {
if error_kind == nom8::error::ErrorKind::Switch {
tx.tx_data.set_event(PgsqlEvent::MalformedRequest as u8);
self.transactions.push_back(tx);
}
@ -708,7 +708,7 @@ impl PgsqlState {
return AppLayerResult::err();
}
PgsqlParseError::NomError(_i, error_kind) => {
if error_kind == nom7::error::ErrorKind::Switch {
if error_kind == nom8::error::ErrorKind::Switch {
tx.tx_data.set_event(PgsqlEvent::MalformedResponse as u8);
self.transactions.push_back(tx);
}

Loading…
Cancel
Save