use serde::ser::{self, Serialize, SerializeSeq, SerializeStruct, SerializeStructVariant};
use std::fmt::{self, Display};
pub fn to_args<T>(value: &T, cmd: &clap::Command) -> Result<Vec<String>, Error>
where
T: Serialize,
{
let mut serializer = ArgSerializer::new(cmd);
value.serialize(&mut serializer)?;
Ok(serializer.args)
}
#[derive(Debug)]
pub struct Error {
message: String,
}
impl Display for Error {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{}", self.message)
}
}
impl std::error::Error for Error {}
impl ser::Error for Error {
fn custom<T: Display>(msg: T) -> Self {
Error {
message: msg.to_string(),
}
}
}
pub struct ArgSerializer<'a> {
cmd: &'a clap::Command,
args: Vec<String>,
current_field: Option<String>,
current_arg: Option<&'a clap::Arg>,
}
impl<'a> ArgSerializer<'a> {
fn new(cmd: &'a clap::Command) -> Self {
Self {
cmd,
args: Vec::new(),
current_field: None,
current_arg: None,
}
}
fn push_flag(&mut self, short: Option<char>, long: Option<&str>) {
if let Some(long_name) = long {
self.args.push(format!("--{long_name}"));
} else if let Some(short_char) = short {
self.args.push(format!("-{short_char}"));
}
}
fn push_option(
&mut self,
short: Option<char>,
long: Option<&str>,
value: String,
is_positional: bool,
) {
if let Some(long_name) = long {
self.args.push(format!("--{long_name}"));
self.args.push(value);
} else if let Some(short_char) = short {
self.args.push(format!("-{short_char}"));
self.args.push(value);
} else if is_positional {
self.args.push(value);
}
}
}
impl<'a, 'b> ser::Serializer for &'b mut ArgSerializer<'a> {
type Ok = ();
type Error = Error;
type SerializeSeq = SeqSerializer<'a, 'b>;
type SerializeTuple = ser::Impossible<(), Error>;
type SerializeTupleStruct = ser::Impossible<(), Error>;
type SerializeTupleVariant = ser::Impossible<(), Error>;
type SerializeMap = ser::Impossible<(), Error>;
type SerializeStruct = StructSerializer<'a, 'b>;
type SerializeStructVariant = StructVariantSerializer<'a, 'b>;
fn serialize_bool(self, v: bool) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
let action = arg.get_action();
match action {
clap::ArgAction::SetTrue => {
if v {
self.push_flag(arg.get_short(), arg.get_long());
}
}
clap::ArgAction::SetFalse => {
if !v {
self.push_flag(arg.get_short(), arg.get_long());
}
}
_ => {
if v {
self.push_flag(arg.get_short(), arg.get_long());
}
}
}
}
Ok(())
}
fn serialize_i8(self, v: i8) -> Result<(), Error> {
self.serialize_i64(v as i64)
}
fn serialize_i16(self, v: i16) -> Result<(), Error> {
self.serialize_i64(v as i64)
}
fn serialize_i32(self, v: i32) -> Result<(), Error> {
self.serialize_i64(v as i64)
}
fn serialize_i64(self, v: i64) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
let value_str = v.to_string();
let is_default = arg
.get_default_values()
.first()
.map(|d| d.to_string_lossy() == value_str)
.unwrap_or(false);
if !is_default {
self.push_option(
arg.get_short(),
arg.get_long(),
value_str,
arg.is_positional(),
);
}
}
Ok(())
}
fn serialize_u8(self, v: u8) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
if matches!(arg.get_action(), clap::ArgAction::Count) {
if v > 0 {
if let Some(short_char) = arg.get_short() {
if v <= 3 {
self.args
.push(format!("-{}", short_char.to_string().repeat(v as usize)));
} else {
for _ in 0..v {
self.args.push(format!("-{short_char}"));
}
}
} else if let Some(long_name) = arg.get_long() {
for _ in 0..v {
self.args.push(format!("--{long_name}"));
}
}
}
return Ok(());
}
}
self.serialize_u64(v as u64)
}
fn serialize_u16(self, v: u16) -> Result<(), Error> {
self.serialize_u64(v as u64)
}
fn serialize_u32(self, v: u32) -> Result<(), Error> {
self.serialize_u64(v as u64)
}
fn serialize_u64(self, v: u64) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
let value_str = v.to_string();
let is_default = arg
.get_default_values()
.first()
.map(|d| d.to_string_lossy() == value_str)
.unwrap_or(false);
if !is_default {
self.push_option(
arg.get_short(),
arg.get_long(),
value_str,
arg.is_positional(),
);
}
}
Ok(())
}
fn serialize_f32(self, v: f32) -> Result<(), Error> {
self.serialize_f64(v as f64)
}
fn serialize_f64(self, v: f64) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
let value_str = v.to_string();
let is_default = arg
.get_default_values()
.first()
.map(|d| d.to_string_lossy() == value_str)
.unwrap_or(false);
if !is_default {
self.push_option(
arg.get_short(),
arg.get_long(),
value_str,
arg.is_positional(),
);
}
}
Ok(())
}
fn serialize_char(self, v: char) -> Result<(), Error> {
self.serialize_str(&v.to_string())
}
fn serialize_str(self, v: &str) -> Result<(), Error> {
if let Some(arg) = self.current_arg {
let is_default = arg
.get_default_values()
.first()
.map(|d| {
let default_str = d.to_string_lossy();
default_str == v || default_str.eq_ignore_ascii_case(v)
})
.unwrap_or(false);
if !is_default {
self.push_option(
arg.get_short(),
arg.get_long(),
v.to_string(),
arg.is_positional(),
);
}
}
Ok(())
}
fn serialize_bytes(self, _v: &[u8]) -> Result<(), Error> {
Err(ser::Error::custom("bytes not supported"))
}
fn serialize_none(self) -> Result<(), Error> {
Ok(())
}
fn serialize_some<T>(self, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_unit(self) -> Result<(), Error> {
Ok(())
}
fn serialize_unit_struct(self, _name: &'static str) -> Result<(), Error> {
Ok(())
}
fn serialize_unit_variant(
self,
_name: &'static str,
_variant_index: u32,
variant: &'static str,
) -> Result<(), Error> {
self.serialize_str(variant)
}
fn serialize_newtype_struct<T>(self, _name: &'static str, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_newtype_variant<T>(
self,
_name: &'static str,
_variant_index: u32,
variant: &'static str,
value: &T,
) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
let command_name = pascal_to_kebab(variant);
if let Some(subcmd) = self
.cmd
.get_subcommands()
.find(|c| c.get_name() == command_name)
{
self.args.push(command_name);
let mut sub_serializer = ArgSerializer {
cmd: subcmd,
args: Vec::new(),
current_field: None,
current_arg: None,
};
value.serialize(&mut sub_serializer)?;
self.args.extend(sub_serializer.args);
}
Ok(())
}
fn serialize_seq(self, _len: Option<usize>) -> Result<Self::SerializeSeq, Error> {
Ok(SeqSerializer {
parent: self,
items: Vec::new(),
})
}
fn serialize_tuple(self, _len: usize) -> Result<Self::SerializeTuple, Error> {
Err(ser::Error::custom("tuples not supported"))
}
fn serialize_tuple_struct(
self,
_name: &'static str,
_len: usize,
) -> Result<Self::SerializeTupleStruct, Error> {
Err(ser::Error::custom("tuple structs not supported"))
}
fn serialize_tuple_variant(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_len: usize,
) -> Result<Self::SerializeTupleVariant, Error> {
Err(ser::Error::custom("tuple variants not supported"))
}
fn serialize_map(self, _len: Option<usize>) -> Result<Self::SerializeMap, Error> {
Err(ser::Error::custom("maps not supported"))
}
fn serialize_struct(
self,
_name: &'static str,
_len: usize,
) -> Result<Self::SerializeStruct, Error> {
Ok(StructSerializer { parent: self })
}
fn serialize_struct_variant(
self,
_name: &'static str,
_variant_index: u32,
variant: &'static str,
_len: usize,
) -> Result<Self::SerializeStructVariant, Error> {
let command_name = pascal_to_kebab(variant);
if self
.cmd
.get_subcommands()
.any(|c| c.get_name() == command_name)
{
self.args.push(command_name.clone());
}
Ok(StructVariantSerializer {
parent: self,
command_name,
})
}
}
pub struct StructSerializer<'a, 'b> {
parent: &'b mut ArgSerializer<'a>,
}
impl<'a, 'b> SerializeStruct for StructSerializer<'a, 'b> {
type Ok = ();
type Error = Error;
fn serialize_field<T>(&mut self, key: &'static str, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
if self.parent.cmd.get_subcommands().next().is_some() {
if key == "command" || key == "cmd" || key == "subcommand" {
value.serialize(&mut *self.parent)?;
return Ok(());
}
}
if let Some(arg) = self
.parent
.cmd
.get_arguments()
.find(|a| a.get_id().as_str() == key)
{
self.parent.current_field = Some(key.to_string());
self.parent.current_arg = Some(arg);
value.serialize(&mut *self.parent)?;
self.parent.current_field = None;
self.parent.current_arg = None;
}
Ok(())
}
fn end(self) -> Result<(), Error> {
Ok(())
}
}
pub struct StructVariantSerializer<'a, 'b> {
parent: &'b mut ArgSerializer<'a>,
command_name: String,
}
impl<'a, 'b> SerializeStructVariant for StructVariantSerializer<'a, 'b> {
type Ok = ();
type Error = Error;
fn serialize_field<T>(&mut self, key: &'static str, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
if let Some(subcmd) = self
.parent
.cmd
.get_subcommands()
.find(|c| c.get_name() == self.command_name)
{
if subcmd.get_subcommands().next().is_some() {
if !subcmd.get_arguments().any(|a| a.get_id().as_str() == key) {
let mut sub_serializer = ArgSerializer {
cmd: subcmd,
args: Vec::new(),
current_field: Some(key.to_string()),
current_arg: None,
};
value.serialize(&mut sub_serializer)?;
self.parent.args.extend(sub_serializer.args);
return Ok(());
}
}
if let Some(arg) = subcmd.get_arguments().find(|a| a.get_id().as_str() == key) {
self.parent.current_field = Some(key.to_string());
self.parent.current_arg = Some(arg);
value.serialize(&mut *self.parent)?;
self.parent.current_field = None;
self.parent.current_arg = None;
}
}
Ok(())
}
fn end(self) -> Result<(), Error> {
Ok(())
}
}
pub struct SeqSerializer<'a, 'b> {
parent: &'b mut ArgSerializer<'a>,
items: Vec<String>,
}
impl<'a, 'b> SerializeSeq for SeqSerializer<'a, 'b> {
type Ok = ();
type Error = Error;
fn serialize_element<T>(&mut self, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
let mut collector = ValueCollector::new();
value.serialize(&mut collector)?;
if let Some(val) = collector.value {
self.items.push(val);
}
Ok(())
}
fn end(self) -> Result<(), Error> {
if let Some(arg) = self.parent.current_arg {
let action = arg.get_action();
match action {
clap::ArgAction::Append => {
for item in self.items {
if let Some(long) = arg.get_long() {
self.parent.args.push(format!("--{long}"));
self.parent.args.push(item);
} else if let Some(short) = arg.get_short() {
self.parent.args.push(format!("-{short}"));
self.parent.args.push(item);
} else {
self.parent.args.push(item);
}
}
}
_ => {
for item in self.items {
self.parent.args.push(item);
}
}
}
} else {
for item in self.items {
self.parent.args.push(item);
}
}
Ok(())
}
}
fn pascal_to_kebab(s: &str) -> String {
let mut result = String::new();
for (i, ch) in s.chars().enumerate() {
if i > 0 && ch.is_uppercase() {
result.push('-');
}
result.push_str(&ch.to_lowercase().to_string());
}
result
}
struct ValueCollector {
value: Option<String>,
}
impl ValueCollector {
fn new() -> Self {
Self { value: None }
}
}
impl ser::Serializer for &mut ValueCollector {
type Ok = ();
type Error = Error;
type SerializeSeq = ser::Impossible<(), Error>;
type SerializeTuple = ser::Impossible<(), Error>;
type SerializeTupleStruct = ser::Impossible<(), Error>;
type SerializeTupleVariant = ser::Impossible<(), Error>;
type SerializeMap = ser::Impossible<(), Error>;
type SerializeStruct = ser::Impossible<(), Error>;
type SerializeStructVariant = ser::Impossible<(), Error>;
fn serialize_bool(self, v: bool) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_i8(self, v: i8) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_i16(self, v: i16) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_i32(self, v: i32) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_i64(self, v: i64) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_u8(self, v: u8) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_u16(self, v: u16) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_u32(self, v: u32) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_u64(self, v: u64) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_f32(self, v: f32) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_f64(self, v: f64) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_char(self, v: char) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_str(self, v: &str) -> Result<(), Error> {
self.value = Some(v.to_string());
Ok(())
}
fn serialize_bytes(self, _v: &[u8]) -> Result<(), Error> {
Err(ser::Error::custom("bytes not supported"))
}
fn serialize_none(self) -> Result<(), Error> {
Ok(())
}
fn serialize_some<T>(self, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_unit(self) -> Result<(), Error> {
Ok(())
}
fn serialize_unit_struct(self, _name: &'static str) -> Result<(), Error> {
Ok(())
}
fn serialize_unit_variant(
self,
_name: &'static str,
_variant_index: u32,
variant: &'static str,
) -> Result<(), Error> {
self.value = Some(variant.to_string());
Ok(())
}
fn serialize_newtype_struct<T>(self, _name: &'static str, value: &T) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_newtype_variant<T>(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_value: &T,
) -> Result<(), Error>
where
T: ?Sized + Serialize,
{
Err(ser::Error::custom("newtype variants not supported"))
}
fn serialize_seq(self, _len: Option<usize>) -> Result<Self::SerializeSeq, Error> {
Err(ser::Error::custom(
"sequences not supported in value collector",
))
}
fn serialize_tuple(self, _len: usize) -> Result<Self::SerializeTuple, Error> {
Err(ser::Error::custom("tuples not supported"))
}
fn serialize_tuple_struct(
self,
_name: &'static str,
_len: usize,
) -> Result<Self::SerializeTupleStruct, Error> {
Err(ser::Error::custom("tuple structs not supported"))
}
fn serialize_tuple_variant(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_len: usize,
) -> Result<Self::SerializeTupleVariant, Error> {
Err(ser::Error::custom("tuple variants not supported"))
}
fn serialize_map(self, _len: Option<usize>) -> Result<Self::SerializeMap, Error> {
Err(ser::Error::custom("maps not supported"))
}
fn serialize_struct(
self,
_name: &'static str,
_len: usize,
) -> Result<Self::SerializeStruct, Error> {
Err(ser::Error::custom(
"structs not supported in value collector",
))
}
fn serialize_struct_variant(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_len: usize,
) -> Result<Self::SerializeStructVariant, Error> {
Err(ser::Error::custom(
"struct variants not supported in value collector",
))
}
}