use std::{borrow::Borrow, convert::Infallible, str::FromStr};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
fn validate_node_id(id: &str) -> Result<(), InvalidId> {
if id.is_empty() {
return Err(InvalidId("identifier must not be empty".into()));
}
if id.starts_with('.') {
return Err(InvalidId(format!(
"identifier '{id}' must not start with '.' \
(dot-segments such as '.' and '..' can traverse parent directories)"
)));
}
if let Some(ch) = id
.chars()
.find(|c| !c.is_ascii_alphanumeric() && *c != '_' && *c != '-' && *c != '.')
{
return Err(InvalidId(format!(
"identifier contains invalid character '{ch}' -- only [a-zA-Z0-9_.-] are allowed"
)));
}
Ok(())
}
fn validate_data_id(id: &str) -> Result<(), InvalidId> {
if id.is_empty() {
return Err(InvalidId("identifier must not be empty".into()));
}
if let Some(ch) = id
.chars()
.find(|c| !c.is_ascii_alphanumeric() && *c != '_' && *c != '-' && *c != '.' && *c != '/')
{
return Err(InvalidId(format!(
"identifier contains invalid character '{ch}' -- only [a-zA-Z0-9_./-] are allowed"
)));
}
if id.starts_with('/') || id.ends_with('/') || id.contains("//") {
return Err(InvalidId(format!(
"identifier '{id}' has empty path segment -- \
leading, trailing, or consecutive '/' are not allowed"
)));
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InvalidId(pub String);
impl std::fmt::Display for InvalidId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
impl std::error::Error for InvalidId {}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, JsonSchema)]
pub struct NodeId(pub(crate) String);
impl<'de> Deserialize<'de> for NodeId {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
validate_node_id(&s).map_err(serde::de::Error::custom)?;
Ok(NodeId(s))
}
}
impl FromStr for NodeId {
type Err = InvalidId;
fn from_str(s: &str) -> Result<Self, Self::Err> {
validate_node_id(s)?;
Ok(Self(s.to_owned()))
}
}
impl From<String> for NodeId {
fn from(id: String) -> Self {
if let Err(e) = validate_node_id(&id) {
panic!("invalid NodeId '{id}': {e}");
}
Self(id)
}
}
impl std::fmt::Display for NodeId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, f)
}
}
impl AsRef<str> for NodeId {
fn as_ref(&self) -> &str {
&self.0
}
}
#[derive(
Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize, JsonSchema,
)]
pub struct OperatorId(String);
impl FromStr for OperatorId {
type Err = Infallible;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(Self(s.to_owned()))
}
}
impl From<String> for OperatorId {
fn from(id: String) -> Self {
Self(id)
}
}
impl std::fmt::Display for OperatorId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, f)
}
}
impl AsRef<str> for OperatorId {
fn as_ref(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, JsonSchema)]
pub struct DataId(String);
impl<'de> Deserialize<'de> for DataId {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
validate_data_id(&s).map_err(serde::de::Error::custom)?;
Ok(DataId(s))
}
}
impl FromStr for DataId {
type Err = InvalidId;
fn from_str(s: &str) -> Result<Self, Self::Err> {
validate_data_id(s)?;
Ok(Self(s.to_owned()))
}
}
impl From<DataId> for String {
fn from(id: DataId) -> Self {
id.0
}
}
impl From<String> for DataId {
fn from(id: String) -> Self {
if let Err(e) = validate_data_id(&id) {
panic!("invalid DataId '{id}': {e}");
}
Self(id)
}
}
impl From<&str> for DataId {
fn from(id: &str) -> Self {
id.to_owned().into()
}
}
impl std::fmt::Display for DataId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, f)
}
}
impl std::ops::Deref for DataId {
type Target = String;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl AsRef<String> for DataId {
fn as_ref(&self) -> &String {
&self.0
}
}
impl AsRef<str> for DataId {
fn as_ref(&self) -> &str {
&self.0
}
}
impl Borrow<String> for DataId {
fn borrow(&self) -> &String {
&self.0
}
}
impl Borrow<str> for DataId {
fn borrow(&self) -> &str {
&self.0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn valid_node_ids() {
assert!(validate_node_id("my-node").is_ok());
assert!(validate_node_id("my_node").is_ok());
assert!(validate_node_id("MyNode123").is_ok());
assert!(validate_node_id("node.v2").is_ok());
assert!(validate_node_id("a").is_ok());
}
#[test]
fn invalid_node_ids() {
assert!(validate_node_id("").is_err());
assert!(validate_node_id("node/output").is_err());
assert!(validate_node_id("node name").is_err());
assert!(validate_node_id("node;rm").is_err());
assert!(validate_node_id("node\0").is_err());
}
#[test]
fn node_id_rejects_dot_segments() {
assert!(validate_node_id(".").is_err(), ". must be rejected");
assert!(validate_node_id("..").is_err(), ".. must be rejected");
assert!(
validate_node_id(".hidden").is_err(),
".hidden must be rejected"
);
assert!(
validate_node_id(".config").is_err(),
".config must be rejected"
);
assert!(
validate_node_id("node.v2").is_ok(),
"embedded dot is allowed"
);
assert!(
validate_node_id("a.b.c").is_ok(),
"multiple embedded dots are allowed"
);
}
#[test]
fn data_id_allows_slash_for_operators() {
assert!(validate_data_id("rust-operator/status").is_ok());
assert!(validate_data_id("op/output").is_ok());
assert!(validate_data_id("simple").is_ok());
}
#[test]
fn data_id_rejects_other_specials() {
assert!(validate_data_id("").is_err());
assert!(validate_data_id("has space").is_err());
assert!(validate_data_id("semi;colon").is_err());
}
#[test]
fn data_id_rejects_malformed_slash_patterns() {
assert!(validate_data_id("op/").is_err(), "trailing slash");
assert!(validate_data_id("/out").is_err(), "leading slash");
assert!(validate_data_id("a//b").is_err(), "double slash");
assert!(validate_data_id("/").is_err(), "bare slash");
}
#[test]
fn node_id_from_str_rejects_invalid() {
assert!(NodeId::from_str("hello").is_ok());
assert!(NodeId::from_str("hello/world").is_err());
assert!(NodeId::from_str("hello world").is_err());
assert!(NodeId::from_str("").is_err());
}
#[test]
#[should_panic(expected = "invalid NodeId")]
fn node_id_from_string_panics_on_invalid() {
let _id: NodeId = "bad/id".to_string().into();
}
#[test]
fn node_id_parse_rejects_invalid() {
assert!("hello".parse::<NodeId>().is_ok());
assert!("bad/id".parse::<NodeId>().is_err());
assert!("".parse::<NodeId>().is_err());
}
#[test]
fn data_id_parse_rejects_invalid() {
assert!("output".parse::<DataId>().is_ok());
assert!("bad;id".parse::<DataId>().is_err());
}
#[test]
fn node_id_deserialize_rejects_invalid() {
let result: Result<NodeId, _> = serde_json::from_str("\"bad/id\"");
assert!(result.is_err());
}
#[test]
fn data_id_deserialize_rejects_invalid() {
let result: Result<DataId, _> = serde_json::from_str("\"bad;id\"");
assert!(result.is_err());
}
}