use std::io;
use rmp::encode::{
write_array_len, write_bool, write_map_len, write_nil, write_sint, write_str, write_u32,
write_uint, write_uint8,
};
use rmpv::{
Value, ValueRef,
encode::{write_value, write_value_ref},
};
use tokio::io::{AsyncWrite, AsyncWriteExt};
use super::ExtType;
const MSG_TYPE_REQUEST: u8 = 0;
const MSG_TYPE_RESPONSE: u8 = 1;
const MSG_TYPE_NOTIFICATION: u8 = 2;
pub struct Encoder<W> {
writer: W,
buffer: Vec<u8>,
}
impl<W> Encoder<W>
where
W: AsyncWrite + Unpin,
{
#[inline]
#[must_use]
pub fn new(writer: W) -> Self {
Self {
writer,
buffer: Vec::new(),
}
}
pub async fn encode_request<A: EncodeArgs>(
&mut self,
msgid: u32,
method: &str,
args: A,
) -> io::Result<()> {
self.buffer.clear();
write_array_len(&mut self.buffer, 4)?;
write_uint8(&mut self.buffer, MSG_TYPE_REQUEST)?;
write_u32(&mut self.buffer, msgid)?;
write_str(&mut self.buffer, method)?;
write_array_len(&mut self.buffer, A::NUM_ARGS)?;
args.encode_args(&mut self.buffer)?;
self.writer.write_all(&self.buffer).await?;
self.writer.flush().await?;
Ok(())
}
pub async fn encode_notify<A: EncodeArgs>(&mut self, method: &str, args: A) -> io::Result<()> {
self.buffer.clear();
write_array_len(&mut self.buffer, 3)?;
write_uint8(&mut self.buffer, MSG_TYPE_NOTIFICATION)?;
write_str(&mut self.buffer, method)?;
write_array_len(&mut self.buffer, A::NUM_ARGS)?;
args.encode_args(&mut self.buffer)?;
self.writer.write_all(&self.buffer).await?;
self.writer.flush().await?;
Ok(())
}
pub async fn encode_result_response<R: Encode>(
&mut self,
msgid: u32,
result: R,
) -> io::Result<()> {
self.buffer.clear();
write_array_len(&mut self.buffer, 4)?;
write_uint8(&mut self.buffer, MSG_TYPE_RESPONSE)?;
write_u32(&mut self.buffer, msgid)?;
write_nil(&mut self.buffer)?;
result.encode(&mut self.buffer)?;
self.writer.write_all(&self.buffer).await?;
self.writer.flush().await?;
Ok(())
}
pub async fn encode_error_response<E: Encode>(
&mut self,
msgid: u32,
error: E,
) -> io::Result<()> {
self.buffer.clear();
write_array_len(&mut self.buffer, 4)?;
write_uint8(&mut self.buffer, MSG_TYPE_RESPONSE)?;
write_u32(&mut self.buffer, msgid)?;
error.encode(&mut self.buffer)?;
write_nil(&mut self.buffer)?;
self.writer.write_all(&self.buffer).await?;
self.writer.flush().await?;
Ok(())
}
}
pub trait Encode {
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()>;
}
impl<const TYPE_ID: i8> Encode for ExtType<TYPE_ID> {
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
let mut data = [0; 5];
let len = {
let mut remaining = &mut data[..];
write_sint(&mut remaining, i64::from(self.0))?;
5 - remaining.len()
};
rmp::encode::write_ext_meta(buf, len as u32, TYPE_ID)?;
buf.extend_from_slice(&data[..len]);
Ok(())
}
}
impl<T: Encode + ?Sized> Encode for &T {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
T::encode(*self, buf)
}
}
impl Encode for ValueRef<'_> {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
Ok(write_value_ref(buf, self)?)
}
}
impl Encode for Value {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
Ok(write_value(buf, self)?)
}
}
impl Encode for str {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
Ok(write_str(buf, self)?)
}
}
impl Encode for bool {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
write_bool(buf, *self)
}
}
macro_rules! impl_encode_int {
($writer:ident as $target:ty; $($ty:ty),+) => {
$(
impl Encode for $ty {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
$writer(buf, *self as $target)?;
Ok(())
}
}
)+
};
}
impl_encode_int!(write_sint as i64; i8, i32, i64);
impl_encode_int!(write_uint as u64; u8, u32, u64, usize);
impl<E: Encode> Encode for [E] {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
let len = u32::try_from(self.len())
.map_err(|err| io::Error::new(io::ErrorKind::InvalidInput, err))?;
write_array_len(buf, len)?;
for elem in self {
elem.encode(buf)?;
}
Ok(())
}
}
impl<const N: usize, E: Encode> Encode for [E; N] {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
self.as_slice().encode(buf)
}
}
impl<K: Encode, V: Encode> Encode for [(K, V)] {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
let len = u32::try_from(self.len())
.map_err(|err| io::Error::new(io::ErrorKind::InvalidInput, err))?;
write_map_len(buf, len)?;
for (key, val) in self {
key.encode(buf)?;
val.encode(buf)?;
}
Ok(())
}
}
impl<const N: usize, K: Encode, V: Encode> Encode for [(K, V); N] {
#[inline]
fn encode(&self, buf: &mut Vec<u8>) -> io::Result<()> {
self.as_slice().encode(buf)
}
}
pub trait EncodeArgs {
const NUM_ARGS: u32;
fn encode_args(self, buf: &mut Vec<u8>) -> io::Result<()>;
}
impl<T: Encode> EncodeArgs for T {
const NUM_ARGS: u32 = 1;
#[inline]
fn encode_args(self, buf: &mut Vec<u8>) -> io::Result<()> {
self.encode(buf)
}
}
macro_rules! impl_encode_tuple_args {
($len:expr, $($arg:ident),*) => {
impl<$($arg: Encode,)*> EncodeArgs for ($($arg,)*) {
const NUM_ARGS: u32 = $len;
#[inline]
fn encode_args(self, _buf: &mut Vec<u8>) -> io::Result<()> {
#[allow(non_snake_case)]
let ($($arg,)*) = self;
$(
$arg.encode(_buf)?;
)*
Ok(())
}
}
};
}
impl_encode_tuple_args!(0,);
impl_encode_tuple_args!(2, T, U);
impl_encode_tuple_args!(3, T, U, V);
impl_encode_tuple_args!(4, T, U, V, W);
impl_encode_tuple_args!(5, T, U, V, W, X);
impl_encode_tuple_args!(6, T, U, V, W, X, Y);
pub struct ArrayWriter<'buf> {
buf: &'buf mut Vec<u8>,
remaining: u32,
}
impl<'buf> ArrayWriter<'buf> {
#[inline]
pub fn new(buf: &'buf mut Vec<u8>, len: u32) -> io::Result<Self> {
write_array_len(buf, len)?;
Ok(Self {
buf,
remaining: len,
})
}
#[inline]
pub fn write(&mut self, val: impl Encode) -> io::Result<()> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
val.encode(self.buf)?;
self.remaining -= 1;
Ok(())
}
#[inline]
pub fn nest_array(&mut self, len: u32) -> io::Result<ArrayWriter<'_>> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
self.remaining -= 1;
ArrayWriter::new(self.buf, len)
}
#[inline]
pub fn nest_map(&mut self, len: u32) -> io::Result<MapWriter<'_>> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
self.remaining -= 1;
MapWriter::new(self.buf, len)
}
#[inline]
pub fn finish(self) -> io::Result<()> {
if self.remaining > 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
Ok(())
}
}
pub struct MapWriter<'buf> {
buf: &'buf mut Vec<u8>,
remaining: u32,
}
impl<'buf> MapWriter<'buf> {
#[inline]
pub fn new(buf: &'buf mut Vec<u8>, len: u32) -> io::Result<Self> {
write_map_len(buf, len)?;
Ok(Self {
buf,
remaining: len,
})
}
#[inline]
pub fn write(&mut self, key: impl Encode, val: impl Encode) -> io::Result<()> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
key.encode(self.buf)?;
val.encode(self.buf)?;
self.remaining -= 1;
Ok(())
}
#[inline]
pub fn nest_array(&mut self, key: impl Encode, len: u32) -> io::Result<ArrayWriter<'_>> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
key.encode(self.buf)?;
self.remaining -= 1;
ArrayWriter::new(self.buf, len)
}
#[inline]
pub fn nest_map(&mut self, key: impl Encode, len: u32) -> io::Result<MapWriter<'_>> {
if self.remaining == 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
key.encode(self.buf)?;
self.remaining -= 1;
MapWriter::new(self.buf, len)
}
#[inline]
pub fn finish(self) -> io::Result<()> {
if self.remaining > 0 {
return Err(io::ErrorKind::InvalidInput.into());
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use rmpv::{Value, decode::read_value};
use tokio::io::sink;
fn decoded_value(bytes: &[u8]) -> Value {
let mut input = bytes;
let value = read_value(&mut input).unwrap();
assert!(input.is_empty());
value
}
fn assert_encoded_integer(value: impl Encode, expected: Value) {
let mut bytes = Vec::new();
value.encode(&mut bytes).unwrap();
assert_eq!(decoded_value(&bytes), expected);
}
fn assert_invalid_input<T>(result: io::Result<T>) {
assert_eq!(result.err().unwrap().kind(), io::ErrorKind::InvalidInput);
}
#[test]
fn encodes_integer_types() {
assert_encoded_integer(i8::MIN, Value::from(i8::MIN));
assert_encoded_integer(i8::MAX, Value::from(i8::MAX));
assert_encoded_integer(i32::MIN, Value::from(i32::MIN));
assert_encoded_integer(i32::MAX, Value::from(i32::MAX));
assert_encoded_integer(i64::MIN, Value::from(i64::MIN));
assert_encoded_integer(i64::MAX, Value::from(i64::MAX));
assert_encoded_integer(u8::MIN, Value::from(u8::MIN));
assert_encoded_integer(u8::MAX, Value::from(u8::MAX));
assert_encoded_integer(u32::MIN, Value::from(u32::MIN));
assert_encoded_integer(u32::MAX, Value::from(u32::MAX));
assert_encoded_integer(u64::MIN, Value::from(u64::MIN));
assert_encoded_integer(u64::MAX, Value::from(u64::MAX));
assert_encoded_integer(usize::MIN, Value::from(usize::MIN as u64));
assert_encoded_integer(usize::MAX, Value::from(usize::MAX as u64));
}
#[test]
fn encodes_ext_type() {
let mut bytes = Vec::new();
ExtType::<1>(42).encode(&mut bytes).unwrap();
assert_eq!(decoded_value(&bytes), Value::Ext(1, vec![42]));
}
#[test]
fn encode_value_ref_slice_as_array() {
let values = [ValueRef::from("value"), ValueRef::from(42)];
let mut bytes = vec![];
values.as_slice().encode(&mut bytes).unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Array(vec![Value::from("value"), Value::from(42)])
);
}
#[test]
fn encode_nested_arrays_as_array() {
let chunks = [["foo"]];
let mut bytes = vec![];
chunks.encode(&mut bytes).unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Array(vec![Value::Array(vec![Value::from("foo")])])
);
}
#[test]
fn encode_value_ref_pairs_as_map() {
let entries = [
(ValueRef::from("enabled"), ValueRef::Boolean(true)),
(ValueRef::from("count"), ValueRef::from(42)),
];
let mut bytes = vec![];
entries.as_slice().encode(&mut bytes).unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Map(vec![
(Value::from("enabled"), Value::from(true)),
(Value::from("count"), Value::from(42)),
])
);
}
#[test]
fn encode_pair_array_as_map() {
let opts = [("err", true)];
let mut bytes = vec![];
opts.encode(&mut bytes).unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Map(vec![(Value::from("err"), Value::from(true))])
);
}
#[tokio::test]
async fn encode_single_arg_matches_request_encoding() {
let mut encoder = Encoder::new(sink());
encoder
.encode_request(7, "nvim_input", "<C-D>")
.await
.unwrap();
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![
Value::from(0),
Value::from(7),
Value::from("nvim_input"),
Value::from(vec![Value::from("<C-D>")]),
])
);
}
#[tokio::test]
async fn encode_single_arg_matches_notification_encoding() {
let mut encoder = Encoder::new(sink());
encoder.encode_notify("nvim_input", "<C-D>").await.unwrap();
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![
Value::from(2),
Value::from("nvim_input"),
Value::from(vec![Value::from("<C-D>")]),
])
);
}
#[tokio::test]
async fn encode_message_matches_value_ref_request_encoding() {
let cmd = ValueRef::Map(vec![
(ValueRef::from("cmd"), ValueRef::from("echo")),
(
ValueRef::from("args"),
ValueRef::Array(vec![ValueRef::from("hello")]),
),
]);
let opts = ValueRef::Map(vec![(ValueRef::from("output"), ValueRef::Boolean(true))]);
let mut encoder = Encoder::new(sink());
encoder
.encode_request(7, "nvim_cmd", (&cmd, &opts))
.await
.unwrap();
let params = Value::from(vec![cmd.to_owned(), opts.to_owned()]);
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![
Value::from(0),
Value::from(7),
Value::from("nvim_cmd"),
params,
])
);
}
#[tokio::test]
async fn encode_message_matches_integer_notification_encoding() {
let mut encoder = Encoder::new(sink());
encoder
.encode_notify("nvim_ui_try_resize", (120_i64, 40_i64))
.await
.unwrap();
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![
Value::from(2),
Value::from("nvim_ui_try_resize"),
Value::from(vec![Value::from(120), Value::from(40)]),
])
);
}
#[tokio::test]
async fn encode_error_response_matches_response_encoding() {
let mut encoder = Encoder::new(sink());
encoder
.encode_error_response(9, "Not implemented")
.await
.unwrap();
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![
Value::from(1),
Value::from(9),
Value::from("Not implemented"),
Value::Nil,
])
);
}
#[tokio::test]
async fn encode_result_response_matches_response_encoding() {
let mut encoder = Encoder::new(sink());
let result = Value::Map(vec![(Value::from("answer"), Value::from(42))]);
encoder.encode_result_response(9, &result).await.unwrap();
assert_eq!(
decoded_value(&encoder.buffer),
Value::from(vec![Value::from(1), Value::from(9), Value::Nil, result,])
);
}
#[test]
fn array_writer_encodes_nested_values() {
let mut bytes = vec![];
let mut array = ArrayWriter::new(&mut bytes, 3).unwrap();
array.write("value").unwrap();
{
let mut nested = array.nest_array(2).unwrap();
nested.write(1).unwrap();
nested.write(2).unwrap();
nested.finish().unwrap();
}
{
let mut nested = array.nest_map(1).unwrap();
nested.write("key", true).unwrap();
nested.finish().unwrap();
}
array.finish().unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Array(vec![
Value::from("value"),
Value::Array(vec![Value::from(1), Value::from(2)]),
Value::Map(vec![(Value::from("key"), Value::from(true))]),
])
);
}
#[test]
fn array_writer_reports_length_errors() {
let mut bytes = vec![];
assert_invalid_input(ArrayWriter::new(&mut bytes, 1).unwrap().finish());
let mut bytes = vec![];
let mut array = ArrayWriter::new(&mut bytes, 0).unwrap();
assert_invalid_input(array.write("value"));
assert_invalid_input(array.nest_array(0));
assert_invalid_input(array.nest_map(0));
array.finish().unwrap();
}
#[test]
fn map_writer_encodes_nested_values() {
let mut bytes = vec![];
let mut map = MapWriter::new(&mut bytes, 3).unwrap();
map.write("value", 42).unwrap();
{
let mut nested = map.nest_array("array", 2).unwrap();
nested.write(1).unwrap();
nested.write(2).unwrap();
nested.finish().unwrap();
}
{
let mut nested = map.nest_map("map", 1).unwrap();
nested.write("key", true).unwrap();
nested.finish().unwrap();
}
map.finish().unwrap();
assert_eq!(
decoded_value(&bytes),
Value::Map(vec![
(Value::from("value"), Value::from(42)),
(
Value::from("array"),
Value::Array(vec![Value::from(1), Value::from(2)]),
),
(
Value::from("map"),
Value::Map(vec![(Value::from("key"), Value::from(true))]),
),
])
);
}
#[test]
fn map_writer_reports_length_errors() {
let mut bytes = vec![];
assert_invalid_input(MapWriter::new(&mut bytes, 1).unwrap().finish());
let mut bytes = vec![];
let mut map = MapWriter::new(&mut bytes, 0).unwrap();
assert_invalid_input(map.write("key", "value"));
assert_invalid_input(map.nest_array("key", 0));
assert_invalid_input(map.nest_map("key", 0));
map.finish().unwrap();
}
}