use crate::{Ctx, GenerateError, ScalarBacking, numeric_tokens, same_package_scalar_backing};
use proc_macro2::TokenStream;
use quote::quote;
use ridl_ir::v2;
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum ClauseKind {
Require,
Ensure,
}
#[derive(Debug)]
pub(crate) struct ClauseBody {
pub expr: TokenStream,
pub uses_args: bool,
pub uses_reply: bool,
}
pub(crate) fn translate(
ctx: &Ctx,
contracts: &[v2::Contract],
kind: ClauseKind,
params: &[v2::Param],
reply_named: Option<&str>,
) -> Result<ClauseBody, GenerateError> {
let wanted = match kind {
ClauseKind::Require => v2::ContractKind::Require,
ClauseKind::Ensure => v2::ContractKind::Ensure,
};
let mut predicates: Vec<TokenStream> = Vec::new();
let mut uses_args = false;
let mut uses_reply = false;
for contract in contracts {
if v2::ContractKind::try_from(contract.kind).ok() != Some(wanted) {
continue;
}
let clause = translate_one(ctx, &contract.source, params, reply_named)?;
match clause.subject {
Subject::Arg => uses_args = true,
Subject::Reply => uses_reply = true,
}
predicates.push(clause.expr);
}
let expr = match predicates.into_iter().reduce(|a, b| quote! { #a && #b }) {
None => quote! { ::core::result::Result::Ok(()) },
Some(conjunction) => quote! {
if #conjunction {
::core::result::Result::Ok(())
} else {
::core::result::Result::Err(())
}
},
};
Ok(ClauseBody {
expr,
uses_args,
uses_reply,
})
}
enum Subject {
Arg,
Reply,
}
struct Clause {
expr: TokenStream,
subject: Subject,
}
#[derive(Clone, Copy)]
enum Comparison {
Lt,
Le,
Gt,
Ge,
Eq,
Ne,
}
impl Comparison {
fn tokens(self) -> TokenStream {
match self {
Comparison::Lt => quote! { < },
Comparison::Le => quote! { <= },
Comparison::Gt => quote! { > },
Comparison::Ge => quote! { >= },
Comparison::Eq => quote! { == },
Comparison::Ne => quote! { != },
}
}
}
fn translate_one(
ctx: &Ctx,
source: &str,
params: &[v2::Param],
reply_named: Option<&str>,
) -> Result<Clause, GenerateError> {
let (subject, comparison, literal) = parse(source)?;
if subject == "result" {
let named = reply_named.ok_or_else(|| {
refuse(
source,
"`result` is only a subject on a query's `ensure` clause",
)
})?;
let backing = scalar_backing(ctx, named)
.ok_or_else(|| refuse(source, "`result` is not a named integer or float scalar"))?;
let value = literal_tokens(&literal, backing).ok_or_else(|| {
refuse(
source,
"the literal does not match the reply's numeric type",
)
})?;
let op = comparison.tokens();
return Ok(Clause {
expr: quote! { reply.0 #op #value },
subject: Subject::Reply,
});
}
let [param] = params else {
return Err(refuse(
source,
"a translated clause needs exactly one declared parameter",
));
};
if param.name != subject {
return Err(refuse(
source,
"the subject is not the interaction's declared parameter",
));
}
let Some(named) = param_named_type(param) else {
return Err(refuse(source, "the parameter is not a named type"));
};
let backing = scalar_backing(ctx, named).ok_or_else(|| {
refuse(
source,
"the parameter is not a named integer or float scalar",
)
})?;
let value = literal_tokens(&literal, backing).ok_or_else(|| {
refuse(
source,
"the literal does not match the parameter's numeric type",
)
})?;
let op = comparison.tokens();
Ok(Clause {
expr: quote! { args.0 #op #value },
subject: Subject::Arg,
})
}
fn parse(source: &str) -> Result<(String, Comparison, String), GenerateError> {
let text = source.trim();
let comparisons = [
("<=", Comparison::Le),
(">=", Comparison::Ge),
("==", Comparison::Eq),
("!=", Comparison::Ne),
("<", Comparison::Lt),
(">", Comparison::Gt),
];
for (symbol, comparison) in comparisons {
let Some(position) = text.find(symbol) else {
continue;
};
let left = text[..position].trim();
let right = text[position + symbol.len()..].trim();
if is_identifier(left) && is_number(right) {
return Ok((left.to_string(), comparison, right.to_string()));
}
return Err(refuse(
source,
"not `<subject> <comparison> <numeric literal>`",
));
}
Err(refuse(source, "no accepted comparison"))
}
fn is_identifier(text: &str) -> bool {
let mut chars = text.chars();
match chars.next() {
Some(first) if first.is_ascii_alphabetic() || first == '_' => {}
_ => return false,
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn is_number(text: &str) -> bool {
!text.is_empty()
&& text
.chars()
.all(|c| c.is_ascii_digit() || matches!(c, '.' | '-' | '+' | 'e' | 'E'))
&& text.parse::<f64>().is_ok()
}
fn literal_tokens(value: &str, backing: ScalarBacking) -> Option<TokenStream> {
match backing {
ScalarBacking::Integer => {
if value.contains(['.', 'e', 'E']) {
return None;
}
Some(numeric_tokens(value, false))
}
ScalarBacking::Float => Some(numeric_tokens(value, true)),
ScalarBacking::Boolean | ScalarBacking::String | ScalarBacking::Bytes => None,
}
}
fn scalar_backing(ctx: &Ctx, reference: &str) -> Option<ScalarBacking> {
match same_package_scalar_backing(ctx, reference)? {
backing @ (ScalarBacking::Integer | ScalarBacking::Float) => Some(backing),
_ => None,
}
}
fn param_named_type(param: &v2::Param) -> Option<&str> {
match param.r#type.as_ref()?.kind.as_ref()? {
v2::field_type::Kind::Named(name) => Some(name),
_ => None,
}
}
pub(crate) const CLAUSE_REFUSAL: &str = "cannot translate contract clause";
fn refuse(source: &str, reason: &str) -> GenerateError {
GenerateError {
message: format!("{CLAUSE_REFUSAL} `{source}`: {reason}"),
}
}
#[cfg(test)]
mod tests {
use super::{ClauseKind, translate};
use crate::Ctx;
use ridl_ir::v2;
fn scalar_package() -> v2::Package {
v2::Package {
name: "p".to_string(),
decls: vec![
scalar_decl("Level", v2::PrimitiveType::Integer),
scalar_decl("Rate", v2::PrimitiveType::Float),
],
interfaces: Vec::new(),
services: Vec::new(),
retired: Vec::new(),
}
}
fn scalar_decl(name: &str, prim: v2::PrimitiveType) -> v2::Decl {
v2::Decl {
name: name.to_string(),
visibility: v2::Visibility::Public as i32,
is_error: false,
doc: String::new(),
labels: Vec::new(),
deprecated: None,
ordinal: 0,
kind: Some(v2::decl::Kind::TypeDef(v2::TypeDef {
backing: Some(v2::Backing {
kind: Some(v2::backing::Kind::Primitive(prim as i32)),
}),
constraint: None,
declared_init: None,
init: None,
width: None,
})),
}
}
fn param(name: &str, type_ref: &str) -> v2::Param {
v2::Param {
name: name.to_string(),
r#type: Some(v2::FieldType {
optional: false,
kind: Some(v2::field_type::Kind::Named(type_ref.to_string())),
}),
}
}
fn clause(kind: v2::ContractKind, source: &str) -> v2::Contract {
v2::Contract {
kind: kind as i32,
source: source.to_string(),
signal_refs: Vec::new(),
param_refs: Vec::new(),
uses_result: false,
observer_id: String::new(),
}
}
fn dense(tokens: &proc_macro2::TokenStream) -> String {
tokens
.to_string()
.chars()
.filter(|c| !c.is_whitespace())
.collect()
}
#[test]
fn an_accepted_parameter_clause_reaches_the_newtype_field() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level < 100")];
let body =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect("accepted");
assert!(body.uses_args);
assert!(!body.uses_reply);
assert_eq!(
dense(&body.expr),
"ifargs.0<100{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn a_le_comparison_emits_the_le_operator() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level <= 100")];
let body =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect("accepted");
assert_eq!(
dense(&body.expr),
"ifargs.0<=100{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn an_eq_comparison_emits_the_eq_operator() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level == 100")];
let body =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect("accepted");
assert_eq!(
dense(&body.expr),
"ifargs.0==100{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn a_ne_comparison_emits_the_ne_operator() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level != 100")];
let body =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect("accepted");
assert_eq!(
dense(&body.expr),
"ifargs.0!=100{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn a_float_backed_result_clause_emits_a_float_literal() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("window", "Level")];
let contracts = [clause(v2::ContractKind::Ensure, "result >= 0")];
let body = translate(&ctx, &contracts, ClauseKind::Ensure, ¶ms, Some("Rate"))
.expect("accepted");
assert!(body.uses_reply);
assert!(!body.uses_args);
assert_eq!(
dense(&body.expr),
"ifreply.0>=0.0{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn several_clauses_of_one_kind_are_conjoined() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [
clause(v2::ContractKind::Require, "level < 100"),
clause(v2::ContractKind::Require, "level >= 0"),
];
let body =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect("accepted");
assert_eq!(
dense(&body.expr),
"ifargs.0<100&&args.0>=0{::core::result::Result::Ok(())}else{::core::result::Result::Err(())}"
);
}
#[test]
fn no_clause_of_a_kind_emits_ok() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level < 100")];
let body =
translate(&ctx, &contracts, ClauseKind::Ensure, ¶ms, Some("Level")).expect("empty");
assert!(!body.uses_args);
assert!(!body.uses_reply);
assert_eq!(dense(&body.expr), "::core::result::Result::Ok(())");
}
#[test]
fn a_compound_clause_is_refused() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(
v2::ContractKind::Require,
"level < 100 || level == 0",
)];
let error =
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect_err("refused");
assert!(error.message.contains("cannot translate contract clause"));
}
#[test]
fn a_unit_literal_is_refused() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("window", "Level")];
let contracts = [clause(v2::ContractKind::Require, "window > 0ms")];
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect_err("refused");
}
#[test]
fn a_result_subject_outside_an_ensure_is_refused() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "result >= 0")];
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect_err("refused");
}
#[test]
fn a_fractional_literal_on_an_integer_subject_is_refused() {
let package = scalar_package();
let ctx = Ctx::new(&package);
let params = [param("level", "Level")];
let contracts = [clause(v2::ContractKind::Require, "level < 1.5")];
translate(&ctx, &contracts, ClauseKind::Require, ¶ms, None).expect_err("refused");
}
}