use crate::{error::EdifactError, model::Segment, tokenizer::ServiceStringAdvice};
use std::borrow::Cow;
use std::io::Write;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DataElement<'a> {
Simple(&'a str),
Composite(&'a [&'a str]),
}
impl<'a> DataElement<'a> {
#[inline]
#[must_use]
pub fn components(&self) -> &[&'a str] {
match self {
Self::Simple(value) => std::slice::from_ref(value),
Self::Composite(components) => components,
}
}
}
impl<'a> From<&'a str> for DataElement<'a> {
#[inline]
fn from(value: &'a str) -> Self {
Self::Simple(value)
}
}
impl<'a> From<&'a [&'a str]> for DataElement<'a> {
#[inline]
fn from(components: &'a [&'a str]) -> Self {
Self::Composite(components)
}
}
impl<'a, const N: usize> From<&'a [&'a str; N]> for DataElement<'a> {
#[inline]
fn from(components: &'a [&'a str; N]) -> Self {
Self::Composite(components)
}
}
pub trait AsDataElement {
fn as_data_element(&self) -> DataElement<'_>;
}
impl AsDataElement for str {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Simple(self)
}
}
impl AsDataElement for &str {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Simple(self)
}
}
impl AsDataElement for String {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Simple(self.as_str())
}
}
impl AsDataElement for Cow<'_, str> {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Simple(self.as_ref())
}
}
impl<const N: usize> AsDataElement for [&str; N] {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Composite(self)
}
}
impl AsDataElement for [&str] {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Composite(self)
}
}
impl AsDataElement for Vec<&str> {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
DataElement::Composite(self)
}
}
impl AsDataElement for DataElement<'_> {
#[inline]
fn as_data_element(&self) -> DataElement<'_> {
*self
}
}
#[macro_export]
macro_rules! elements {
() => {
&[] as &[$crate::DataElement<'_>]
};
($($element:expr),+ $(,)?) => {
&[$($crate::AsDataElement::as_data_element(&$element)),+][..]
};
}
pub struct Writer<W: Write> {
inner: W,
ssa: ServiceStringAdvice,
segment_count: u64,
message_start_count: u64,
}
#[inline]
fn find_escape(ssa: &ServiceStringAdvice, hay: &[u8]) -> Option<usize> {
let first = memchr::memchr3(ssa.element_sep, ssa.component_sep, ssa.release_char, hay);
let second = if ssa.repetition_sep == b' ' {
memchr::memchr(ssa.segment_term, hay)
} else {
memchr::memchr2(ssa.segment_term, ssa.repetition_sep, hay)
};
match (first, second) {
(None, None) => None,
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(Some(a), Some(b)) => Some(a.min(b)),
}
}
impl<W: Write> Writer<W> {
pub fn new(inner: W) -> Self {
Self {
inner,
ssa: ServiceStringAdvice::default(),
segment_count: 0,
message_start_count: 0,
}
}
pub fn with_una(mut inner: W, ssa: ServiceStringAdvice) -> Result<Self, EdifactError> {
if !ssa.is_valid() {
return Err(EdifactError::InvalidUna);
}
let una = [
b'U',
b'N',
b'A',
ssa.component_sep,
ssa.element_sep,
ssa.decimal_mark,
ssa.release_char,
ssa.repetition_sep,
ssa.segment_term,
];
inner.write_all(&una)?;
Ok(Self {
inner,
ssa,
segment_count: 0,
message_start_count: 0,
})
}
pub fn write_segment(&mut self, seg: &Segment<'_>) -> Result<(), EdifactError> {
self.inner.write_all(seg.tag.as_bytes())?;
for element in &seg.elements {
self.inner.write_all(&[self.ssa.element_sep])?;
let mut first_component = true;
for (component, _) in &element.components {
if !first_component {
self.inner.write_all(&[self.ssa.component_sep])?;
}
first_component = false;
self.write_escaped(component)?;
}
}
self.inner.write_all(&[self.ssa.segment_term])?;
self.segment_count += 1;
Ok(())
}
pub fn write_raw(&mut self, tag: &str, elements: &[&str]) -> Result<(), EdifactError> {
self.inner.write_all(tag.as_bytes())?;
let comp_sep = self.ssa.component_sep;
for el in elements {
self.inner.write_all(&[self.ssa.element_sep])?;
let mut parts = el.as_bytes().split(|&b| b == comp_sep);
if let Some(first) = parts.next() {
self.write_escaped(
std::str::from_utf8(first).map_err(|_| EdifactError::InvalidUtf8)?,
)?;
}
for part in parts {
self.inner.write_all(&[comp_sep])?;
self.write_escaped(
std::str::from_utf8(part).map_err(|_| EdifactError::InvalidUtf8)?,
)?;
}
}
self.inner.write_all(&[self.ssa.segment_term])?;
if tag == "UNH" {
self.message_start_count = self.segment_count;
}
self.segment_count += 1;
Ok(())
}
pub fn write_segment_parts<E>(&mut self, tag: &str, elements: &[E]) -> Result<(), EdifactError>
where
E: AsRef<[String]>,
{
self.inner.write_all(tag.as_bytes())?;
for element in elements {
self.inner.write_all(&[self.ssa.element_sep])?;
let mut first = true;
for comp in element.as_ref() {
if !first {
self.inner.write_all(&[self.ssa.component_sep])?;
}
first = false;
self.write_escaped(comp.as_str())?;
}
}
self.inner.write_all(&[self.ssa.segment_term])?;
self.segment_count += 1;
Ok(())
}
pub fn write_composites(
&mut self,
tag: &str,
elements: &[&[&str]],
) -> Result<(), EdifactError> {
self.inner.write_all(tag.as_bytes())?;
for element in elements {
self.inner.write_all(&[self.ssa.element_sep])?;
for (i, comp) in element.iter().enumerate() {
if i > 0 {
self.inner.write_all(&[self.ssa.component_sep])?;
}
self.write_escaped(comp)?;
}
}
self.inner.write_all(&[self.ssa.segment_term])?;
if tag == "UNH" {
self.message_start_count = self.segment_count;
}
self.segment_count += 1;
Ok(())
}
pub fn write_elements(
&mut self,
tag: &str,
elements: &[DataElement<'_>],
) -> Result<(), EdifactError> {
self.inner.write_all(tag.as_bytes())?;
for element in elements {
self.inner.write_all(&[self.ssa.element_sep])?;
for (i, comp) in element.components().iter().enumerate() {
if i > 0 {
self.inner.write_all(&[self.ssa.component_sep])?;
}
self.write_escaped(comp)?;
}
}
self.inner.write_all(&[self.ssa.segment_term])?;
if tag == "UNH" {
self.message_start_count = self.segment_count;
}
self.segment_count += 1;
Ok(())
}
pub fn finish(mut self) -> Result<W, EdifactError> {
self.inner.flush()?;
Ok(self.inner)
}
pub fn finish_unt(mut self, message_ref: &str) -> Result<W, EdifactError> {
let count = self.segment_count - self.message_start_count + 1;
let count_str = count.to_string();
self.write_composites("UNT", &[&[count_str.as_str()], &[message_ref]])?;
self.finish()
}
pub fn segment_count(&self) -> u64 {
self.segment_count
}
pub fn service_string_advice(&self) -> ServiceStringAdvice {
self.ssa
}
pub fn escape_value<'v>(&self, value: &'v str) -> Cow<'v, str> {
let release = self.ssa.release_char;
let bytes = value.as_bytes();
if find_escape(&self.ssa, bytes).is_none() {
return Cow::Borrowed(value);
}
let mut out = Vec::with_capacity(value.len() + 4);
let mut last = 0;
let mut pos = 0;
while pos < bytes.len() {
let Some(hit) = find_escape(&self.ssa, &bytes[pos..]) else {
break;
};
let abs = pos + hit;
out.extend_from_slice(&bytes[last..abs]);
out.push(release);
out.push(bytes[abs]);
last = abs + 1;
pos = abs + 1;
}
out.extend_from_slice(&bytes[last..]);
Cow::Owned(
String::from_utf8(out).expect(
"escape_value: output is not valid UTF-8; this is a bug in the escape logic",
),
)
}
#[inline]
pub(crate) fn write_tag_only(&mut self, tag: &str) -> Result<(), EdifactError> {
self.inner.write_all(tag.as_bytes())?;
Ok(())
}
#[inline]
pub(crate) fn write_element_sep(&mut self) -> Result<(), EdifactError> {
self.inner.write_all(&[self.ssa.element_sep])?;
Ok(())
}
#[inline]
pub(crate) fn write_component_sep(&mut self) -> Result<(), EdifactError> {
self.inner.write_all(&[self.ssa.component_sep])?;
Ok(())
}
#[inline]
pub(crate) fn write_segment_term_and_count(&mut self) -> Result<(), EdifactError> {
self.inner.write_all(&[self.ssa.segment_term])?;
self.segment_count += 1;
Ok(())
}
pub(crate) fn write_escaped(&mut self, value: &str) -> Result<(), EdifactError> {
let release = self.ssa.release_char;
let bytes = value.as_bytes();
let mut last = 0;
let mut pos = 0;
while pos < bytes.len() {
let Some(hit) = find_escape(&self.ssa, &bytes[pos..]) else {
break;
};
let abs = pos + hit;
if abs > last {
self.inner.write_all(&bytes[last..abs])?;
}
self.inner.write_all(&[release, bytes[abs]])?;
last = abs + 1;
pos = abs + 1;
}
self.inner.write_all(&bytes[last..])?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn begin_interchange(
&mut self,
syntax_id: &str,
syntax_version: &str,
sender: &str,
recipient: &str,
date: &str,
time: &str,
control_ref: &str,
) -> Result<(), EdifactError> {
self.write_composites(
"UNB",
&[
&[syntax_id, syntax_version],
&[sender],
&[recipient],
&[date, time],
&[control_ref],
],
)
}
pub fn begin_message<'w>(
&'w mut self,
message_ref: &str,
message_type: &str,
version: &str,
release: &str,
controlling_agency: &str,
) -> Result<MessageWriter<'w, W>, EdifactError> {
self.write_composites(
"UNH",
&[
&[message_ref],
&[message_type, version, release, controlling_agency],
],
)?;
let unh_count = self.segment_count;
Ok(MessageWriter {
writer: self,
message_ref: message_ref.to_owned(),
unh_count,
finished: false,
})
}
pub fn end_interchange(
&mut self,
message_count: u32,
control_ref: &str,
) -> Result<(), EdifactError> {
let msg_count_str = message_count.to_string();
self.write_composites("UNZ", &[&[msg_count_str.as_str()], &[control_ref]])
}
}
pub struct MessageWriter<'w, W: Write> {
writer: &'w mut Writer<W>,
message_ref: String,
unh_count: u64,
finished: bool,
}
impl<W: Write> MessageWriter<'_, W> {
pub fn write_raw(&mut self, tag: &str, elements: &[&str]) -> Result<(), EdifactError> {
self.writer.write_raw(tag, elements)
}
pub fn write_elements(
&mut self,
tag: &str,
elements: &[DataElement<'_>],
) -> Result<(), EdifactError> {
self.writer.write_elements(tag, elements)
}
pub fn write_composites(
&mut self,
tag: &str,
elements: &[&[&str]],
) -> Result<(), EdifactError> {
self.writer.write_composites(tag, elements)
}
pub fn write_segment_parts<E>(&mut self, tag: &str, elements: &[E]) -> Result<(), EdifactError>
where
E: AsRef<[String]>,
{
self.writer.write_segment_parts(tag, elements)
}
pub fn write_segment(&mut self, seg: &Segment<'_>) -> Result<(), EdifactError> {
self.writer.write_segment(seg)
}
pub fn finish(mut self) -> Result<(), EdifactError> {
self.write_unt()?;
self.finished = true;
Ok(())
}
fn write_unt(&mut self) -> Result<(), EdifactError> {
let count = self.writer.segment_count - self.unh_count + 2;
let count_str = count.to_string();
self.writer.write_composites(
"UNT",
&[&[count_str.as_str()], &[self.message_ref.as_str()]],
)
}
}
impl<W: Write> Drop for MessageWriter<'_, W> {
fn drop(&mut self) {
if !self.finished {
let _ = self.write_unt();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Element;
fn exotic_ssa() -> ServiceStringAdvice {
ServiceStringAdvice {
component_sep: b'|',
element_sep: b'!',
decimal_mark: b',',
release_char: b'#',
repetition_sep: b'*',
segment_term: b'~',
}
}
#[test]
fn unh_composite_uses_the_active_component_separator() {
let mut buf = Vec::new();
{
let mut w = Writer::with_una(&mut buf, exotic_ssa()).unwrap();
let msg = w
.begin_message("1", "ORDERS", "D", "96A", "UN")
.expect("UNH");
msg.finish().expect("UNT");
}
let out = String::from_utf8(buf).unwrap();
assert!(
out.contains("UNH!1!ORDERS|D|96A|UN~"),
"S009 must use `|`, got {out}"
);
}
#[test]
fn round_trips_through_a_custom_una() {
let mut buf = Vec::new();
{
let mut w = Writer::with_una(&mut buf, exotic_ssa()).unwrap();
w.begin_interchange("UNOA", "1", "SENDER", "RECEIVER", "200101", "0900", "IC1")
.unwrap();
let mut msg = w.begin_message("1", "ORDERS", "D", "96A", "UN").unwrap();
msg.write_raw("BGM", &["220"]).unwrap();
msg.finish().unwrap();
w.end_interchange(1, "IC1").unwrap();
}
let segs: Vec<_> = crate::from_bytes(&buf)
.collect::<Result<Vec<_>, _>>()
.expect("own output must reparse");
let unh = segs.iter().find(|s| s.tag == "UNH").unwrap();
assert_eq!(unh.get_element(1).unwrap().get_component(0), Some("ORDERS"));
assert_eq!(unh.get_element(1).unwrap().get_component(2), Some("96A"));
crate::validate_envelope(&segs).expect("own output must pass envelope validation");
}
#[test]
fn finish_unt_counts_only_the_current_message() {
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
w.begin_interchange("UNOA", "1", "S", "R", "200101", "0900", "IC1")
.unwrap();
w.write_composites("UNH", &[&["1"], &["ORDERS", "D", "96A", "UN"]])
.unwrap();
w.write_raw("BGM", &["220"]).unwrap();
w.finish_unt("1").unwrap();
}
let out = String::from_utf8(buf).unwrap();
assert!(out.contains("UNT+3+1'"), "expected UNT+3, got {out}");
}
#[test]
fn repetition_separator_is_escaped_when_declared() {
let mut buf = Vec::new();
{
let mut w = Writer::with_una(
&mut buf,
ServiceStringAdvice {
repetition_sep: b'*',
..ServiceStringAdvice::default()
},
)
.unwrap();
w.write_composites("FTX", &[&["a*b"]]).unwrap();
}
let out = String::from_utf8(buf).unwrap();
assert!(out.ends_with("FTX+a?*b'"), "rep-sep unescaped in {out}");
}
#[test]
fn repetition_separator_sentinel_is_not_escaped() {
let w = Writer::new(std::io::sink());
assert_eq!(w.escape_value("a b"), "a b");
}
#[test]
fn write_composites_escapes_a_literal_component_separator() {
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
w.write_composites("NAD", &[&["MS"], &["ACME:INC"]])
.unwrap();
}
let segs: Vec<_> = crate::from_bytes(&buf)
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert_eq!(
segs[0].get_element(1).unwrap().get_component(0),
Some("ACME:INC")
);
}
#[test]
fn write_elements_mixes_simple_and_composite() {
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
w.write_elements(
"NAD",
&[
DataElement::Simple("MS"),
DataElement::Composite(&["9900112233445", "", "293"]),
],
)
.unwrap();
}
assert_eq!(buf, b"NAD+MS+9900112233445::293'");
}
#[test]
fn elements_macro_matches_the_explicit_form() {
let mut macro_buf = Vec::new();
{
let mut w = Writer::new(&mut macro_buf);
w.write_elements("NAD", elements!["MS", ["ACME", "", "9"]])
.unwrap();
w.write_elements("DTM", elements![["137", "20260101", "102"]])
.unwrap();
}
let mut explicit_buf = Vec::new();
{
let mut w = Writer::new(&mut explicit_buf);
w.write_elements(
"NAD",
&[
DataElement::Simple("MS"),
DataElement::Composite(&["ACME", "", "9"]),
],
)
.unwrap();
w.write_elements(
"DTM",
&[DataElement::Composite(&["137", "20260101", "102"])],
)
.unwrap();
}
assert_eq!(macro_buf, explicit_buf);
assert_eq!(macro_buf, b"NAD+MS+ACME::9'DTM+137:20260101:102'");
}
#[test]
fn elements_macro_accepts_arbitrary_expressions() {
let qualifier = String::from("MS");
let gln = "9900112233445";
let dtm: Vec<&str> = vec!["137", "20260101", "102"];
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
w.write_elements("NAD", elements![qualifier.as_str(), [gln, "", "293"]])
.unwrap();
w.write_elements("DTM", elements![dtm]).unwrap();
w.write_elements("FTX", elements![qualifier]).unwrap();
w.write_elements("UNS", elements![]).unwrap();
}
assert_eq!(
String::from_utf8(buf).unwrap(),
"NAD+MS+9900112233445::293'DTM+137:20260101:102'FTX+MS'UNS'"
);
}
#[test]
fn write_elements_escapes_a_literal_component_separator() {
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
w.write_elements(
"NAD",
&[DataElement::Simple("MS"), DataElement::Simple("ACME:INC")],
)
.unwrap();
}
let segs: Vec<_> = crate::from_bytes(&buf)
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert_eq!(
segs[0].get_element(1).unwrap().get_component(0),
Some("ACME:INC")
);
}
#[test]
fn write_elements_uses_the_active_component_separator() {
let mut buf = Vec::new();
{
let mut w = Writer::with_una(&mut buf, exotic_ssa()).unwrap();
w.write_elements("DTM", elements![["137", "20260101", "102"]])
.unwrap();
}
let out = String::from_utf8(buf).unwrap();
assert!(
out.ends_with("DTM!137|20260101|102~"),
"expected custom delimiters, got {out}"
);
}
#[test]
fn message_writer_counts_write_elements_segments() {
let mut buf = Vec::new();
{
let mut w = Writer::new(&mut buf);
let mut msg = w.begin_message("1", "ORDERS", "D", "96A", "UN").unwrap();
msg.write_elements("NAD", elements!["MS", ["ACME", "", "9"]])
.unwrap();
msg.write_composites("DTM", &[&["137", "20260101", "102"]])
.unwrap();
msg.finish().unwrap();
}
let out = String::from_utf8(buf).unwrap();
assert!(out.contains("UNT+4+1'"), "expected UNT+4, got {out}");
}
#[test]
fn write_and_parse_simple_segment() {
let segs: Vec<Segment<'static>> = vec![Segment::new(
"BGM",
vec![Element::of(&["220"]), Element::of(&["ORDER123"])],
)];
let bytes = crate::segments_to_bytes(&segs).unwrap();
let s = std::str::from_utf8(&bytes).unwrap();
assert!(s.starts_with("BGM+220+ORDER123'"));
}
#[test]
fn release_char_escaped() {
let segs: Vec<Segment<'static>> = vec![Segment::new(
"FTX",
vec![Element::of(&["value+with+delimiters"])],
)];
let bytes = crate::segments_to_bytes(&segs).unwrap();
let s = std::str::from_utf8(&bytes).unwrap();
assert!(s.contains("?+"), "escape missing: {s}");
}
#[test]
fn round_trip_preserves_values() {
let segs: Vec<Segment<'static>> = vec![
Segment::new(
"UNB",
vec![
Element::of(&["UNOA", "1"]),
Element::of(&["SENDER"]),
Element::of(&["RECEIVER"]),
],
),
Segment::new("UNZ", vec![Element::of(&["0"]), Element::of(&["1"])]),
];
let bytes = crate::segments_to_bytes(&segs).unwrap();
let rt: Vec<crate::OwnedSegment> = crate::parser::from_reader(std::io::Cursor::new(&bytes))
.expect("round-trip parse failed");
assert_eq!(rt[0].tag, "UNB");
assert_eq!(rt[0].as_borrowed().element_str(0), Some("UNOA"));
assert_eq!(rt[1].tag, "UNZ");
}
#[test]
fn with_una_non_default_delimiters() {
use crate::tokenizer::ServiceStringAdvice;
let ssa = ServiceStringAdvice {
component_sep: b'|',
element_sep: b'!',
release_char: b'?',
decimal_mark: b',',
repetition_sep: b'*',
segment_term: b'~',
};
let buf = Vec::new();
let mut writer = Writer::with_una(buf, ssa).expect("writer creation failed");
writer
.write_segment_parts(
"BGM",
&[
vec!["220".to_owned(), "SUB1".to_owned()],
vec!["PO1".to_owned()],
],
)
.expect("write failed");
let out = writer.finish().expect("finish failed");
let s = std::str::from_utf8(&out).unwrap();
assert!(s.contains("BGM"), "BGM segment missing: {s}");
let after_una = s.find("BGM").map(|i| &s[i..]).unwrap_or(s);
assert!(
after_una.contains('!'),
"missing element sep in segment: {after_una}"
);
assert!(
after_una.contains('|'),
"missing component sep in segment: {after_una}"
);
assert!(
after_una.ends_with('~'),
"missing segment term in segment: {after_una}"
);
assert!(s.contains(','), "missing decimal mark in UNA: {s}");
assert!(!s.contains('+'), "default element sep leaked: {s}");
assert!(!s.contains(':'), "default component sep leaked: {s}");
assert!(
!after_una.contains('\''),
"default segment term leaked after UNA: {after_una}"
);
}
}