use std::collections::HashMap;
use crate::error::{CaError, CaResult};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum AccessLevel {
NoAccess,
Read,
ReadWrite,
}
#[derive(Debug, Clone)]
pub struct AccessChecked {
pv_name: String,
level: AccessLevel,
rule_was_trap: bool,
_seal: AccessSeal,
}
#[derive(Debug, Clone)]
struct AccessSeal;
impl AccessChecked {
pub fn pv_name(&self) -> &str {
&self.pv_name
}
pub fn level(&self) -> AccessLevel {
self.level
}
pub fn allows_read(&self) -> bool {
!matches!(self.level, AccessLevel::NoAccess)
}
pub fn allows_write(&self) -> bool {
matches!(self.level, AccessLevel::ReadWrite)
}
pub fn rule_was_trap(&self) -> bool {
self.rule_was_trap
}
}
#[derive(Clone)]
pub struct AccessGate {
inner: AccessGateInner,
acl_version: AclVersionSource,
inp_resolver: Option<InpResolver>,
check_cache: std::sync::Arc<parking_lot::RwLock<HashMap<CheckKey, CachedCheck>>>,
}
#[derive(Clone, PartialEq, Eq, Hash)]
struct CheckKey {
pv_name: String,
host: String,
user: String,
method: String,
authority: String,
roles: Vec<String>,
}
#[derive(Clone, Copy)]
struct CachedCheck {
acl_version: u64,
asg_generation: u64,
cfg_ident: usize,
level: AccessLevel,
rule_was_trap: bool,
}
const CHECK_CACHE_CAP: usize = 4096;
#[derive(Clone)]
enum AclVersionSource {
Atomic(std::sync::Arc<std::sync::atomic::AtomicU64>),
Aggregator(std::sync::Arc<dyn Fn() -> u64 + Send + Sync>),
}
pub type AsgAslResolver = std::sync::Arc<
dyn Fn(String) -> std::pin::Pin<Box<dyn std::future::Future<Output = (String, u8)> + Send>>
+ Send
+ Sync,
>;
pub type InpResolver = std::sync::Arc<
dyn Fn(String) -> std::pin::Pin<Box<dyn std::future::Future<Output = Option<f64>> + Send>>
+ Send
+ Sync,
>;
#[derive(Clone)]
pub struct AcfCell(std::sync::Arc<arc_swap::ArcSwapOption<AccessSecurityConfig>>);
impl AcfCell {
pub fn load(&self) -> arc_swap::Guard<Option<std::sync::Arc<AccessSecurityConfig>>> {
self.0.load()
}
pub fn load_full(&self) -> Option<std::sync::Arc<AccessSecurityConfig>> {
self.0.load_full()
}
pub fn store(&self, value: Option<std::sync::Arc<AccessSecurityConfig>>) {
self.0.store(value);
notify_asg_field_changed();
}
}
pub fn new_acf_cell(initial: Option<AccessSecurityConfig>) -> AcfCell {
AcfCell(std::sync::Arc::new(arc_swap::ArcSwapOption::new(
initial.map(std::sync::Arc::new),
)))
}
pub fn new_acf_cell_watching(
initial: Option<AccessSecurityConfig>,
db: &std::sync::Arc<crate::server::database::PvDatabase>,
) -> AcfCell {
let cell = new_acf_cell(initial);
spawn_asg_inp_watcher(db, &cell);
cell
}
pub fn start_acf_watchers(
db: &std::sync::Arc<crate::server::database::PvDatabase>,
cell: &AcfCell,
) {
spawn_asg_inp_watcher(db, cell);
spawn_hag_refresh(cell);
}
const HAG_DNS_REFRESH: std::time::Duration = std::time::Duration::from_secs(60);
fn spawn_hag_refresh(cell: &AcfCell) {
let weak = std::sync::Arc::downgrade(&cell.0);
crate::runtime::task::spawn_background(
crate::runtime::task::CallbackPriority::Medium,
async move {
loop {
crate::runtime::task::sleep_background(HAG_DNS_REFRESH).await;
let Some(inner) = weak.upgrade() else { break };
if !as_check_client_ip() {
continue;
}
let Some(config) = inner.load_full() else {
continue;
};
if let Some(refreshed) = config.with_refreshed_hags() {
tracing::info!(
target: "epics_base_rs::access_security",
"HAG DNS refresh: re-resolved members changed; republishing policy"
);
AcfCell(inner).store(Some(std::sync::Arc::new(refreshed)));
}
}
},
);
}
const ASG_INP_RETRY: std::time::Duration = std::time::Duration::from_secs(10);
fn spawn_asg_inp_watcher(db: &std::sync::Arc<crate::server::database::PvDatabase>, cell: &AcfCell) {
let weak_cell = std::sync::Arc::downgrade(&cell.0);
let weak_db = std::sync::Arc::downgrade(db);
let mut acf_rx = subscribe_asg_changes();
crate::runtime::task::spawn_background(
crate::runtime::task::CallbackPriority::Medium,
async move {
let mut watched: Option<crate::runtime::taskwd::TaskwdEntry> = None;
enum Wake {
Values,
Rebuild,
Stop,
}
let mut readers: Vec<crate::server::event_queue::EventReader> = Vec::new();
let mut pending: Vec<(String, String)> = Vec::new();
let mut built_from: usize = 0;
loop {
let (Some(inner), Some(db)) = (weak_cell.upgrade(), weak_db.upgrade()) else {
break;
};
let config = inner.load_full();
drop(inner);
match (config.is_some(), watched.is_some()) {
(true, false) => {
watched = Some(crate::runtime::taskwd::taskwd_insert(
"asCaTask",
crate::runtime::taskwd::CheckIn::Unbounded,
None,
))
}
(false, true) => watched = None,
_ => {}
}
let id = config
.as_ref()
.map_or(0, |c| std::sync::Arc::as_ptr(c) as usize);
if id != built_from {
built_from = id;
readers.clear();
pending = config.map(|c| c.inp_link_targets()).unwrap_or_default();
}
pending.retain(|(record, field)| !attach_asg_inp(&db, record, field, &mut readers));
drop(db);
let wake = tokio::select! {
r = acf_rx.recv() => match r {
Err(tokio::sync::broadcast::error::RecvError::Closed) => Wake::Stop,
_ => Wake::Rebuild,
},
() = drain_any_asg_inp(&mut readers) => Wake::Values,
() = crate::runtime::task::sleep_background(ASG_INP_RETRY) => Wake::Rebuild,
};
match wake {
Wake::Stop => break,
Wake::Rebuild => {}
Wake::Values => notify_asg_field_changed(),
}
}
},
);
}
fn attach_asg_inp(
db: &crate::server::database::PvDatabase,
record: &str,
field: &str,
readers: &mut Vec<crate::server::event_queue::EventReader>,
) -> bool {
let Some(rec) = db.get_record(record) else {
return false;
};
let reader = rec.write().add_subscriber(
field,
0,
crate::types::DbFieldType::Double,
crate::server::recgbl::EventMask::VALUE.bits(),
);
match reader {
Some(reader) => {
readers.push(reader);
true
}
None => false,
}
}
async fn drain_any_asg_inp(readers: &mut [crate::server::event_queue::EventReader]) {
std::future::poll_fn(|cx| {
let mut fired = false;
for reader in readers.iter_mut() {
while let std::task::Poll::Ready(Some(_)) = reader.poll_recv(cx) {
fired = true;
}
}
if fired {
std::task::Poll::Ready(())
} else {
std::task::Poll::Pending
}
})
.await
}
#[derive(Clone)]
enum AccessGateInner {
Required {
acf: AcfCell,
resolver: AsgAslResolver,
},
Open,
}
impl AccessGate {
pub fn required(acf: AcfCell, resolver: AsgAslResolver) -> Self {
Self::required_with_version(
acf,
resolver,
std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0)),
)
}
pub fn required_with_version(
acf: AcfCell,
resolver: AsgAslResolver,
acl_version: std::sync::Arc<std::sync::atomic::AtomicU64>,
) -> Self {
Self {
inner: AccessGateInner::Required { acf, resolver },
acl_version: AclVersionSource::Atomic(acl_version),
inp_resolver: None,
check_cache: std::sync::Arc::new(parking_lot::RwLock::new(HashMap::new())),
}
}
pub fn with_inp_resolver(mut self, resolver: InpResolver) -> Self {
self.inp_resolver = Some(resolver);
self
}
pub fn open() -> Self {
Self {
inner: AccessGateInner::Open,
acl_version: AclVersionSource::Atomic(std::sync::Arc::new(
std::sync::atomic::AtomicU64::new(0),
)),
inp_resolver: None,
check_cache: std::sync::Arc::new(parking_lot::RwLock::new(HashMap::new())),
}
}
pub fn open_with_aggregator(f: std::sync::Arc<dyn Fn() -> u64 + Send + Sync>) -> Self {
Self {
inner: AccessGateInner::Open,
acl_version: AclVersionSource::Aggregator(f),
inp_resolver: None,
check_cache: std::sync::Arc::new(parking_lot::RwLock::new(HashMap::new())),
}
}
pub fn acl_version(&self) -> u64 {
match &self.acl_version {
AclVersionSource::Atomic(a) => a.load(std::sync::atomic::Ordering::Acquire),
AclVersionSource::Aggregator(f) => f(),
}
}
pub fn bump_acl_version(&self) {
if let AclVersionSource::Atomic(a) = &self.acl_version {
a.fetch_add(1, std::sync::atomic::Ordering::Release);
}
}
pub async fn check(
&self,
pv_name: impl Into<String>,
host: &str,
user: &str,
method: &str,
authority: &str,
) -> AccessChecked {
self.check_with_roles(pv_name, host, user, &[], method, authority)
.await
}
pub async fn check_with_roles(
&self,
pv_name: impl Into<String>,
host: &str,
user: &str,
roles: &[String],
method: &str,
authority: &str,
) -> AccessChecked {
let pv_name = pv_name.into();
let (level, rule_was_trap) = match &self.inner {
AccessGateInner::Open => (AccessLevel::ReadWrite, false),
AccessGateInner::Required { acf, resolver } => {
match acf.load_full() {
None => (AccessLevel::ReadWrite, false),
Some(cfg) => {
let acl_version = self.acl_version();
let asg_generation = asg_change_generation();
let cfg_ident = std::sync::Arc::as_ptr(&cfg) as usize;
let key = self.inp_resolver.is_none().then(|| CheckKey {
pv_name: pv_name.clone(),
host: host.to_string(),
user: user.to_string(),
method: method.to_string(),
authority: authority.to_string(),
roles: roles.to_vec(),
});
if let Some(ref key) = key
&& let Some(hit) = self.check_cache.read().get(key)
&& hit.acl_version == acl_version
&& hit.asg_generation == asg_generation
&& hit.cfg_ident == cfg_ident
{
return AccessChecked {
pv_name,
level: hit.level,
rule_was_trap: hit.rule_was_trap,
_seal: AccessSeal,
};
}
let (asg, asl) = resolver(pv_name.clone()).await;
let inp_values: Option<AsgInputs> = match self.inp_resolver {
None => None,
Some(ref res) => {
let mut inputs = AsgInputs::default();
if let Some(group) =
cfg.asg.get(&asg).or_else(|| cfg.asg.get("DEFAULT"))
{
for inp in &group.inp {
inputs.record(inp.index, res(inp.link.clone()).await);
}
}
Some(inputs)
}
};
let (level, rule_was_trap) = cfg.compute_for_name(
&asg,
host,
user,
roles,
asl,
method,
authority,
inp_values.as_ref(),
);
if let Some(key) = key {
let mut cache = self.check_cache.write();
if cache.len() >= CHECK_CACHE_CAP {
cache.clear();
}
cache.insert(
key,
CachedCheck {
acl_version,
asg_generation,
cfg_ident,
level,
rule_was_trap,
},
);
}
(level, rule_was_trap)
}
}
}
};
AccessChecked {
pv_name,
level,
rule_was_trap,
_seal: AccessSeal,
}
}
}
#[cfg(test)]
mod access_checked_tests {
use super::*;
use std::sync::Arc;
#[epics_macros_rs::epics_test]
async fn open_gate_grants_read_write() {
let gate = AccessGate::open();
let checked = gate.check("any:pv", "h", "u", "anonymous", "").await;
assert_eq!(checked.level(), AccessLevel::ReadWrite);
assert!(checked.allows_read());
assert!(checked.allows_write());
assert_eq!(checked.pv_name(), "any:pv");
}
#[epics_macros_rs::epics_test]
async fn required_gate_with_no_acf_attached_is_permissive() {
let cell = crate::server::access_security::new_acf_cell(None);
let resolver: AsgAslResolver =
Arc::new(|_pv| Box::pin(async { ("DEFAULT".to_string(), 0u8) }));
let gate = AccessGate::required(cell, resolver);
let checked = gate.check("any:pv", "h", "u", "anonymous", "").await;
assert_eq!(checked.level(), AccessLevel::ReadWrite);
}
#[epics_macros_rs::epics_test]
async fn required_gate_with_acf_denies_unprivileged_peer() {
let cfg = parse_acf(
r#"
UAG(ops) { alice }
ASG(DEFAULT) {
RULE(0, READ) { UAG(ops) }
}
"#,
)
.unwrap();
let cell = crate::server::access_security::new_acf_cell(Some(cfg));
let resolver: AsgAslResolver =
Arc::new(|_pv| Box::pin(async { ("DEFAULT".to_string(), 0u8) }));
let gate = AccessGate::required(cell, resolver);
let allowed = gate.check("x", "h", "alice", "anonymous", "").await;
assert!(allowed.allows_read());
assert!(!allowed.allows_write());
let denied = gate.check("x", "h", "intruder", "anonymous", "").await;
assert_eq!(denied.level(), AccessLevel::NoAccess);
assert!(!denied.allows_read());
}
#[epics_macros_rs::epics_test]
async fn check_cache_misses_on_acf_swap_without_version_bump() {
let cfg_deny = parse_acf(
r#"
ASG(DEFAULT) {
}
"#,
)
.unwrap();
let cfg_allow = parse_acf(
r#"
ASG(DEFAULT) {
RULE(1, WRITE)
}
"#,
)
.unwrap();
let cell = crate::server::access_security::new_acf_cell(Some(cfg_deny));
let resolver: AsgAslResolver =
Arc::new(|_pv| Box::pin(async { ("DEFAULT".to_string(), 0u8) }));
let gate = AccessGate::required(cell.clone(), resolver);
assert!(
!gate
.check("x", "h", "u", "anonymous", "")
.await
.allows_write()
);
assert!(
!gate
.check("x", "h", "u", "anonymous", "")
.await
.allows_write()
);
cell.store(Some(Arc::new(cfg_allow)));
assert!(
gate.check("x", "h", "u", "anonymous", "")
.await
.allows_write()
);
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum RuleAccess {
#[default]
None,
Read,
Write,
}
#[derive(Debug, Clone, Default)]
pub struct AccessRule {
pub level: u8,
pub access: RuleAccess,
pub uag: Vec<String>,
pub hag: Vec<String>,
pub method: Vec<String>,
pub authority: Vec<String>,
pub trap: bool,
pub calc: Option<String>,
pub calc_compiled: Option<crate::calc::CompiledExpr>,
pub inp_used: u32,
pub ignore: bool,
}
fn rule_access(rule: &AccessRule) -> AccessLevel {
if rule.ignore {
return AccessLevel::NoAccess;
}
match rule.access {
RuleAccess::None => AccessLevel::NoAccess,
RuleAccess::Read => AccessLevel::Read,
RuleAccess::Write => AccessLevel::ReadWrite,
}
}
fn calc_result_is_true(result: f64) -> bool {
result > 0.99 && result < 1.01
}
fn rule_rank(level: AccessLevel) -> u8 {
match level {
AccessLevel::NoAccess => 0,
AccessLevel::Read => 1,
AccessLevel::ReadWrite => 2,
}
}
#[derive(Debug, Clone, Default)]
pub struct AccessSecurityGroup {
pub rules: Vec<AccessRule>,
pub inp: Vec<AsgInp>,
}
#[derive(Debug, Clone, Default)]
pub struct AsgInputs {
pub values: crate::calc::NumericInputs,
pub bad: u32,
}
impl AsgInputs {
pub fn record(&mut self, index: u8, value: Option<f64>) {
let idx = index as usize;
if idx >= crate::calc::CALC_NARGS {
return;
}
match value {
Some(v) => self.values.vars[idx] = v,
None => self.bad |= 1u32 << idx,
}
}
}
#[derive(Debug, Clone)]
pub struct AsgInp {
pub index: u8,
pub link: String,
}
pub fn inp_link_target(link: &str) -> (&str, &str) {
let (record, field) = crate::server::database::parse_pv_name(link);
(record, if field.is_empty() { "VAL" } else { field })
}
#[derive(Debug, Clone)]
pub struct AccessSecurityConfig {
pub uag: HashMap<String, Vec<String>>,
pub hag: HashMap<String, Vec<String>>,
pub hag_raw: HashMap<String, Vec<String>>,
pub asg: HashMap<String, AccessSecurityGroup>,
pub unknown_access: AccessLevel,
}
impl AccessSecurityConfig {
pub fn with_refreshed_hags(&self) -> Option<Self> {
let refreshed: HashMap<String, Vec<String>> = self
.hag_raw
.iter()
.map(|(name, raw)| (name.clone(), hag_members(raw)))
.collect();
if refreshed == self.hag {
return None;
}
let mut new = self.clone();
new.hag = refreshed;
Some(new)
}
pub fn dump_report(&self) -> String {
let mut out = String::new();
let mut uags: Vec<_> = self.uag.keys().collect();
uags.sort();
for name in uags {
out.push_str(&format!("UAG({name})\n"));
for m in &self.uag[name] {
out.push('\t');
dump_quoted(&mut out, m);
out.push('\n');
}
}
let mut hags: Vec<_> = self.hag.keys().collect();
hags.sort();
for name in hags {
out.push_str(&format!("HAG({name})\n"));
for h in &self.hag[name] {
out.push('\t');
dump_quoted(&mut out, h);
out.push('\n');
}
}
let mut asgs: Vec<_> = self.asg.keys().collect();
asgs.sort();
for name in asgs {
out.push_str(&format!("ASG({name})\n"));
self.fmt_asg(name, &mut out);
}
out
}
pub fn fmt_asg(&self, name: &str, out: &mut String) {
let Some(asg) = self.asg.get(name) else {
return;
};
for inp in &asg.inp {
let letter = (b'A' + inp.index) as char;
out.push_str(&format!("\tINP{letter}(\"{}\")\n", inp.link));
}
for rule in &asg.rules {
let access = match rule.access {
RuleAccess::None => "NONE",
RuleAccess::Read => "READ",
RuleAccess::Write => "WRITE",
};
let disabled = if rule.ignore { " [DISABLED]" } else { "" };
out.push_str(&format!("\tRULE({},{access}){disabled}\n", rule.level));
for u in &rule.uag {
out.push_str(&format!("\t\tUAG({u})\n"));
}
for h in &rule.hag {
out.push_str(&format!("\t\tHAG({h})\n"));
}
for m in &rule.method {
out.push_str(&format!("\t\tMETHOD(\"{m}\")\n"));
}
for a in &rule.authority {
out.push_str(&format!("\t\tAUTHORITY(\"{a}\")\n"));
}
if let Some(calc) = &rule.calc {
out.push_str(&format!("\t\tCALC(\"{calc}\")\n"));
}
}
}
pub fn check_access(&self, asg_name: &str, host: &str, user: &str) -> AccessLevel {
self.check_access_asl(asg_name, host, user, 0)
}
pub fn check_access_method(
&self,
asg_name: &str,
host: &str,
user: &str,
record_asl: u8,
method: &str,
authority: &str,
) -> AccessLevel {
self.check_access_method_trap(asg_name, host, user, record_asl, method, authority)
.0
}
#[allow(clippy::too_many_arguments)]
pub fn compute_for_name(
&self,
asg_name: &str,
host: &str,
user: &str,
roles: &[String],
record_asl: u8,
method: &str,
authority: &str,
inputs: Option<&AsgInputs>,
) -> (AccessLevel, bool) {
let asg = match self.asg.get(asg_name) {
Some(a) => a,
None => match self.asg.get("DEFAULT") {
Some(a) => a,
None => return (AccessLevel::NoAccess, false),
},
};
self.compute_rules(
asg, host, user, roles, record_asl, method, authority, inputs,
)
}
pub fn check_access_method_trap(
&self,
asg_name: &str,
host: &str,
user: &str,
record_asl: u8,
method: &str,
authority: &str,
) -> (AccessLevel, bool) {
let asg = match self.asg.get(asg_name) {
Some(a) => a,
None => match self.asg.get("DEFAULT") {
Some(a) => a,
None => return (AccessLevel::NoAccess, false),
},
};
self.compute_rules(asg, host, user, &[], record_asl, method, authority, None)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn compute_rules(
&self,
asg: &AccessSecurityGroup,
host: &str,
user: &str,
roles: &[String],
record_asl: u8,
method: &str,
authority: &str,
inputs: Option<&AsgInputs>,
) -> (AccessLevel, bool) {
let mut access = AccessLevel::NoAccess;
let mut trap = false;
for rule in &asg.rules {
if rule.ignore {
continue;
}
if access == AccessLevel::ReadWrite {
break;
}
if rule_rank(rule_access(rule)) <= rule_rank(access) {
continue;
}
if record_asl > rule.level {
continue;
}
let user_match = rule.uag.is_empty()
|| rule.uag.iter().any(|g| {
self.uag
.get(g)
.map(|members| {
members.iter().any(|m| {
m == user
|| matches!(
m.strip_prefix("role/"),
Some(role) if roles.iter().any(|r| r == role)
)
})
})
.unwrap_or(false)
});
if !user_match {
continue;
}
let host_lc = host.to_ascii_lowercase();
let host_match = rule.hag.is_empty()
|| rule.hag.iter().any(|g| {
self.hag
.get(g)
.map(|members| members.iter().any(|m| m.eq_ignore_ascii_case(&host_lc)))
.unwrap_or(false)
});
if !host_match {
continue;
}
let method_match = rule.method.is_empty()
|| rule.method.iter().any(|m| m.eq_ignore_ascii_case(method));
if !method_match {
continue;
}
let authority_match = rule.authority.is_empty()
|| rule
.authority
.iter()
.any(|a| a.eq_ignore_ascii_case(authority));
if !authority_match {
continue;
}
if rule.calc.is_some() {
let Some(compiled) = rule.calc_compiled.as_ref() else {
continue;
};
let Some(inputs) = inputs else {
continue;
};
if inputs.bad & rule.inp_used != 0 {
continue;
}
match crate::calc::eval(compiled, &mut inputs.values.clone()) {
Ok(result) if calc_result_is_true(result) => {}
_ => continue,
}
}
access = rule_access(rule);
trap = rule.trap;
}
(access, trap)
}
pub fn resolve_asg_inputs(
&self,
asg_name: &str,
resolve: &dyn Fn(&str) -> Option<f64>,
) -> AsgInputs {
let mut inputs = AsgInputs::default();
let Some(group) = self.asg.get(asg_name).or_else(|| self.asg.get("DEFAULT")) else {
return inputs;
};
for inp in &group.inp {
inputs.record(inp.index, resolve(&inp.link));
}
inputs
}
pub fn inp_link_targets(&self) -> Vec<(String, String)> {
let mut targets = std::collections::BTreeSet::new();
for group in self.asg.values() {
for inp in &group.inp {
let (record, field) = inp_link_target(&inp.link);
targets.insert((record.to_string(), field.to_string()));
}
}
targets.into_iter().collect()
}
pub fn check_access_asl(
&self,
asg_name: &str,
host: &str,
user: &str,
record_asl: u8,
) -> AccessLevel {
self.check_access_method(asg_name, host, user, record_asl, "", "")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TrapWriteOp {
BeforeWrite,
AfterWrite,
}
#[derive(Debug, Clone, Copy)]
pub struct TrapWriteMessage<'a> {
pub op: TrapWriteOp,
pub pv_name: &'a str,
pub user: &'a str,
pub host: &'a str,
pub peer: &'a str,
pub value_str: &'a str,
pub dbr_type: u16,
pub no_elements: u32,
pub event_id: u64,
pub status: Option<&'a str>,
pub rule_was_trap: bool,
}
static TRAP_WRITE_EVENT_ID: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
pub fn next_trap_write_event_id() -> u64 {
TRAP_WRITE_EVENT_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
}
pub type TrapWriteListener = std::sync::Arc<dyn Fn(&TrapWriteMessage<'_>) + Send + Sync>;
pub struct TrapWriteListenerHandle {
id: u64,
}
impl Drop for TrapWriteListenerHandle {
fn drop(&mut self) {
if let Some(reg) = TRAP_WRITE_REGISTRY.get() {
let mut guard = reg.write().expect("trap-write registry poisoned");
guard.retain(|(id, _)| *id != self.id);
}
}
}
static TRAP_WRITE_REGISTRY: std::sync::OnceLock<std::sync::RwLock<Vec<(u64, TrapWriteListener)>>> =
std::sync::OnceLock::new();
static TRAP_WRITE_NEXT_ID: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
fn trap_write_registry() -> &'static std::sync::RwLock<Vec<(u64, TrapWriteListener)>> {
TRAP_WRITE_REGISTRY.get_or_init(|| std::sync::RwLock::new(Vec::new()))
}
pub fn register_trap_write_listener(listener: TrapWriteListener) -> TrapWriteListenerHandle {
let id = TRAP_WRITE_NEXT_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut guard = trap_write_registry()
.write()
.expect("trap-write registry poisoned");
guard.push((id, listener));
TrapWriteListenerHandle { id }
}
pub fn has_trap_write_listeners() -> bool {
let Some(reg) = TRAP_WRITE_REGISTRY.get() else {
return false;
};
let guard = reg.read().expect("trap-write registry poisoned");
!guard.is_empty()
}
pub fn dispatch_trap_write(msg: &TrapWriteMessage<'_>) {
let Some(reg) = TRAP_WRITE_REGISTRY.get() else {
return;
};
let snapshot: Vec<TrapWriteListener> = {
let guard = reg.read().expect("trap-write registry poisoned");
if guard.is_empty() {
return;
}
guard.iter().map(|(_, l)| l.clone()).collect()
};
for listener in snapshot {
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
listener(msg);
}));
if let Err(payload) = result {
let descr = if let Some(s) = payload.downcast_ref::<&'static str>() {
(*s).to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"(non-string panic payload)".to_string()
};
tracing::error!(
target: "epics_base_rs::server::access_security",
pv = msg.pv_name,
event_id = msg.event_id,
op = ?msg.op,
panic = %descr,
"TRAPWRITE listener panicked — isolating; remaining listeners will still run. \
C asTrapWriteWithData has no unwind concept; this is a Rust-only safety net \
to keep the per-circuit task alive."
);
}
}
}
pub struct TrapWriteFields {
pub pv_name: String,
pub user: String,
pub host: String,
pub peer: String,
pub value_str: String,
pub dbr_type: u16,
pub no_elements: u32,
pub event_id: u64,
pub rule_was_trap: bool,
pub cancel_status: String,
}
impl TrapWriteFields {
fn message<'a>(&'a self, op: TrapWriteOp, status: Option<&'a str>) -> TrapWriteMessage<'a> {
TrapWriteMessage {
op,
pv_name: &self.pv_name,
user: &self.user,
host: &self.host,
peer: &self.peer,
value_str: &self.value_str,
dbr_type: self.dbr_type,
no_elements: self.no_elements,
event_id: self.event_id,
status,
rule_was_trap: self.rule_was_trap,
}
}
}
pub struct TrapWriteGuard {
armed: Option<Box<TrapWriteFields>>,
}
impl TrapWriteGuard {
pub fn begin(fields: TrapWriteFields) -> Self {
if !has_trap_write_listeners() {
return Self { armed: None };
}
dispatch_trap_write(&fields.message(TrapWriteOp::BeforeWrite, None));
Self {
armed: Some(Box::new(fields)),
}
}
pub fn complete(&mut self, status: &str) {
if let Some(fields) = self.armed.take() {
dispatch_trap_write(&fields.message(TrapWriteOp::AfterWrite, Some(status)));
}
}
}
impl Drop for TrapWriteGuard {
fn drop(&mut self) {
if let Some(fields) = self.armed.take() {
dispatch_trap_write(
&fields.message(TrapWriteOp::AfterWrite, Some(&fields.cancel_status)),
);
}
}
}
pub struct TrapWriteMeta<'a> {
pub pv_name: &'a str,
pub user: &'a str,
pub host: &'a str,
pub peer: &'a str,
pub dbr_type: u16,
}
pub fn trap_write_armed(rule_was_trap: bool) -> bool {
rule_was_trap && has_trap_write_listeners()
}
pub async fn put_with_trap<T, E, F, Fut>(
rule_was_trap: bool,
meta: TrapWriteMeta<'_>,
value: crate::types::EpicsValue,
write: F,
) -> Result<T, E>
where
F: FnOnce(crate::types::EpicsValue) -> Fut,
Fut: std::future::Future<Output = Result<T, E>>,
{
if !trap_write_armed(rule_was_trap) {
return write(value).await;
}
let mut guard = TrapWriteGuard::begin(trap_fields(&meta, &value));
let result = write(value).await;
guard.complete(if result.is_ok() { "ok" } else { "fail" });
result
}
pub fn put_with_trap_blocking<T, E, F>(
rule_was_trap: bool,
meta: TrapWriteMeta<'_>,
value: crate::types::EpicsValue,
write: F,
) -> Result<T, E>
where
F: FnOnce(crate::types::EpicsValue) -> Result<T, E>,
{
if !trap_write_armed(rule_was_trap) {
return write(value);
}
let mut guard = TrapWriteGuard::begin(trap_fields(&meta, &value));
let result = write(value);
guard.complete(if result.is_ok() { "ok" } else { "fail" });
result
}
fn trap_fields(meta: &TrapWriteMeta<'_>, value: &crate::types::EpicsValue) -> TrapWriteFields {
TrapWriteFields {
pv_name: meta.pv_name.to_string(),
user: meta.user.to_string(),
host: meta.host.to_string(),
peer: meta.peer.to_string(),
value_str: value.display_truncated(64),
dbr_type: meta.dbr_type,
no_elements: value.count(),
event_id: next_trap_write_event_id(),
rule_was_trap: true,
cancel_status: "cancel".to_string(),
}
}
static ASG_CHANGE_BROADCAST: std::sync::OnceLock<tokio::sync::broadcast::Sender<()>> =
std::sync::OnceLock::new();
fn asg_change_broadcast() -> &'static tokio::sync::broadcast::Sender<()> {
ASG_CHANGE_BROADCAST.get_or_init(|| {
let (tx, _rx) = tokio::sync::broadcast::channel(16);
tx
})
}
static ASG_CHANGE_GENERATION: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub fn notify_asg_field_changed() {
ASG_CHANGE_GENERATION.fetch_add(1, std::sync::atomic::Ordering::Release);
let _ = asg_change_broadcast().send(());
}
pub fn asg_change_generation() -> u64 {
ASG_CHANGE_GENERATION.load(std::sync::atomic::Ordering::Acquire)
}
pub fn subscribe_asg_changes() -> tokio::sync::broadcast::Receiver<()> {
asg_change_broadcast().subscribe()
}
fn dump_quoted(out: &mut String, s: &str) {
let bytes = s.as_bytes();
let len = bytes.iter().position(|&b| b == 0).unwrap_or(bytes.len());
out.push('"');
out.push_str(&crate::runtime::epics_string::print_escaped(&bytes[..len]));
out.push('"');
}
struct AcfScanner<'a> {
src: &'a str,
pos: usize,
line: u32,
}
impl<'a> AcfScanner<'a> {
fn new(content: &'a str) -> Self {
Self {
src: content,
pos: 0,
line: 1,
}
}
fn peek(&mut self) -> Option<char> {
self.src[self.pos..].chars().next()
}
fn next(&mut self) -> Option<char> {
let c = self.src[self.pos..].chars().next()?;
self.pos += c.len_utf8();
if c == '\n' {
self.line += 1;
}
Some(c)
}
fn offending(&self, accepted: &str) -> String {
const CAP: usize = 32;
let rest: String = self.src[self.pos..]
.chars()
.take_while(|c| !c.is_whitespace() && !matches!(c, '(' | ')' | '{' | '}' | ',' | '#'))
.collect();
let token: String = accepted.chars().chain(rest.chars()).collect();
if token.chars().count() > CAP {
let head: String = token.chars().take(CAP).collect();
format!("{head}…")
} else {
token
}
}
fn skip_ws_comments(&mut self) {
while let Some(c) = self.peek() {
if c.is_whitespace() {
self.next();
} else if c == '#' {
while let Some(c) = self.peek() {
self.next();
if c == '\n' {
break;
}
}
} else {
break;
}
}
}
fn reject(&self, what: impl std::fmt::Display) -> CaError {
self.reject_at(self.line, what)
}
fn reject_at(&self, line: u32, what: impl std::fmt::Display) -> CaError {
CaError::Protocol(format!("ACF line {line}: {what}"))
}
}
pub fn parse_acf(content: &str) -> CaResult<AccessSecurityConfig> {
let mut config = AccessSecurityConfig {
uag: HashMap::new(),
hag: HashMap::new(),
hag_raw: HashMap::new(),
asg: HashMap::new(),
unknown_access: AccessLevel::Read,
};
config
.asg
.insert("DEFAULT".to_string(), AccessSecurityGroup::default());
let mut sc = AcfScanner::new(content);
let mut buf = String::new();
while sc.peek().is_some() {
sc.skip_ws_comments();
buf.clear();
read_word(&mut sc, &mut buf);
match buf.as_str() {
"UAG" => {
let name = read_paren_name(&mut sc)?;
let members = read_brace_list(&mut sc)?;
config.uag.insert(name, members);
}
"HAG" => {
let name = read_paren_name(&mut sc)?;
let members = read_brace_list(&mut sc)?;
config.hag.insert(name.clone(), hag_members(&members));
config.hag_raw.insert(name, members);
}
"ASG" => {
let name = read_paren_name(&mut sc)?;
let asg = parse_asg_body(&mut sc)?;
config.asg.insert(name, asg);
}
"" => {
match sc.peek() {
Some(c) if matches!(c, '(' | ')' | '{' | '}' | ',') => {
return Err(sc.reject(format!(
"unexpected '{c}' where a top-level block keyword is expected"
)));
}
_ => break,
}
}
other => {
skip_unknown_top_level_block(other, &mut sc)?;
}
}
}
Ok(config)
}
fn skip_unknown_top_level_block(keyword: &str, sc: &mut AcfScanner) -> CaResult<()> {
let at = sc.line;
sc.skip_ws_comments();
if sc.peek() != Some('(') {
return Err(sc.reject_at(
at,
format!(
"unexpected token '{keyword}' — expected a top-level \
UAG/HAG/ASG block or an unknown keyword followed by '('"
),
));
}
let mut depth = 0;
let mut closed = false;
while let Some(c) = sc.peek() {
sc.next();
match c {
'(' => depth += 1,
')' => {
depth -= 1;
if depth == 0 {
closed = true;
break;
}
}
_ => {}
}
}
if !closed {
return Err(sc.reject_at(
at,
format!("unbalanced '(' in unsupported top-level block '{keyword}'"),
));
}
sc.skip_ws_comments();
if sc.peek() == Some('{') {
let mut depth = 0;
let mut closed = false;
while let Some(c) = sc.peek() {
sc.next();
match c {
'{' => depth += 1,
'}' => {
depth -= 1;
if depth == 0 {
closed = true;
break;
}
}
_ => {}
}
}
if !closed {
return Err(sc.reject_at(
at,
format!("unbalanced '{{' in unsupported top-level block '{keyword}'"),
));
}
}
tracing::warn!(
target: "epics_base_rs::access_security",
line = at,
keyword = %keyword,
"ACF: ignoring unsupported top-level block"
);
Ok(())
}
static AS_CHECK_CLIENT_IP: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
pub fn as_check_client_ip() -> bool {
AS_CHECK_CLIENT_IP.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_as_check_client_ip(on: bool) {
AS_CHECK_CLIENT_IP.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn hag_members(members: &[String]) -> Vec<String> {
if !as_check_client_ip() {
return members.iter().map(|m| m.to_ascii_lowercase()).collect();
}
use std::net::ToSocketAddrs;
members
.iter()
.map(|m| match format!("{m}:0").to_socket_addrs() {
Ok(iter) => match iter.filter(|sa| sa.is_ipv4()).map(|sa| sa.ip()).next() {
Some(ip) => ip.to_string(),
None => format!("unresolved:{m}"),
},
Err(e) => {
tracing::warn!(
target: "epics_base_rs::access_security",
host = %m,
error = %e,
"ACF: Unable to resolve host (asCheckClientIP=1)"
);
format!("unresolved:{m}")
}
})
.collect()
}
fn read_word(sc: &mut AcfScanner, buf: &mut String) {
while let Some(c) = sc.peek() {
if c.is_alphanumeric() || c == '_' {
buf.push(c);
sc.next();
} else {
break;
}
}
}
fn read_paren_name(sc: &mut AcfScanner) -> CaResult<String> {
sc.skip_ws_comments();
if sc.next() != Some('(') {
return Err(sc.reject("expected '('"));
}
let opened = sc.line;
sc.skip_ws_comments();
let mut name = String::new();
if sc.peek() == Some('"') {
sc.next();
let mut closed = false;
while let Some(c) = sc.peek() {
sc.next();
if c == '"' {
closed = true;
break;
}
name.push(c);
}
if !closed {
return Err(sc.reject_at(opened, "unterminated quoted name"));
}
sc.skip_ws_comments();
if sc.next() != Some(')') {
return Err(sc.reject("expected ')' after quoted name"));
}
return Ok(name);
}
loop {
match sc.peek() {
Some(')') => {
sc.next();
break;
}
Some(c) if c.is_whitespace() => {
sc.skip_ws_comments();
match sc.peek() {
Some(')') => {
sc.next();
break;
}
Some(_) => {
return Err(sc.reject("whitespace inside parenthesised name"));
}
None => {
return Err(sc.reject_at(opened, "unterminated '(' — missing ')'"));
}
}
}
Some(c) => {
name.push(c);
sc.next();
}
None => {
return Err(sc.reject_at(opened, "unterminated '(' — missing ')'"));
}
}
}
Ok(name)
}
fn read_brace_list(sc: &mut AcfScanner) -> CaResult<Vec<String>> {
sc.skip_ws_comments();
if sc.next() != Some('{') {
return Err(sc.reject("expected '{'"));
}
let opened = sc.line;
let mut items = Vec::new();
let mut current = String::new();
loop {
sc.skip_ws_comments();
match sc.peek() {
Some('}') => {
sc.next();
break;
}
Some(',') => {
sc.next();
if !current.is_empty() {
items.push(current.clone());
current.clear();
}
}
Some('"') => {
sc.next(); let quote_opened = sc.line;
if !current.is_empty() {
items.push(current.clone());
current.clear();
}
let mut quoted = String::new();
loop {
match sc.next() {
Some('"') => break,
Some('\\') => {
if let Some(esc) = sc.next() {
quoted.push(esc);
}
}
Some('\n') | None => {
return Err(sc.reject_at(quote_opened, "unterminated quoted string"));
}
Some(c) => quoted.push(c),
}
}
if !quoted.is_empty() {
items.push(quoted);
}
}
Some(c)
if c.is_alphanumeric()
|| matches!(c, '_' | '.' | '-' | '+' | ':' | '[' | ']' | '<' | '>' | ';') =>
{
current.push(c);
sc.next();
}
Some(_) => {
sc.next();
}
None => return Err(sc.reject_at(opened, "unterminated '{'")),
}
}
if !current.is_empty() {
items.push(current);
}
Ok(items)
}
fn parse_asg_body(sc: &mut AcfScanner) -> CaResult<AccessSecurityGroup> {
sc.skip_ws_comments();
if sc.next() != Some('{') {
return Err(sc.reject("expected '{' after ASG name"));
}
let opened = sc.line;
let mut asg = AccessSecurityGroup::default();
loop {
sc.skip_ws_comments();
match sc.peek() {
Some('}') => {
sc.next();
break;
}
Some(_) => {
let mut kw = String::new();
read_word(sc, &mut kw);
if kw == "RULE" {
let rule = parse_rule(sc)?;
asg.rules.push(rule);
} else if let Some(stripped) = kw.strip_prefix("INP") {
let index = match parse_inp_index(stripped) {
Some(i) => i,
None => {
return Err(sc.reject(format!(
"invalid INP link selector 'INP{stripped}' \
(expected INPA..INPU)"
)));
}
};
let link = read_paren_name(sc)?;
asg.inp.push(AsgInp { index, link });
} else if kw.is_empty() {
sc.next(); }
}
None => return Err(sc.reject_at(opened, "unterminated ASG")),
}
}
Ok(asg)
}
fn parse_inp_index(suffix: &str) -> Option<u8> {
let mut it = suffix.chars();
let c = it.next()?;
if it.next().is_some() {
return None; }
if ('A'..='U').contains(&c) {
Some((c as u8) - b'A')
} else {
None
}
}
fn parse_rule(sc: &mut AcfScanner) -> CaResult<AccessRule> {
sc.skip_ws_comments();
if sc.next() != Some('(') {
return Err(sc.reject("expected '(' after RULE"));
}
sc.skip_ws_comments();
let mut level_str = String::new();
if matches!(sc.peek(), Some('+') | Some('-')) {
level_str.push(sc.next().unwrap());
}
while let Some(c) = sc.peek() {
if c.is_ascii_digit() {
level_str.push(c);
sc.next();
} else {
break;
}
}
let level_num: i64 = match level_str.parse() {
Ok(n) => n,
Err(_) => {
let got = sc.offending(&level_str);
return Err(sc.reject(format!("RULE level must be an integer, got '{got}'")));
}
};
if level_num < 0 {
return Err(sc.reject(format!("RULE LEVEL must be positive: {level_num}")));
}
let level: u8 = u8::try_from(level_num)
.map_err(|_| sc.reject(format!("RULE level out of range: {level_num}")))?;
sc.skip_ws_comments();
if sc.peek() == Some(',') {
sc.next();
}
sc.skip_ws_comments();
let mut access_str = String::new();
read_word(sc, &mut access_str);
let (access, mut ignore) = if access_str == "WRITE" {
(RuleAccess::Write, false)
} else if access_str == "READ" {
(RuleAccess::Read, false)
} else if access_str == "NONE" {
(RuleAccess::None, false)
} else {
let keyword = sc.offending(&access_str);
tracing::warn!(
target: "epics_base_rs::access_security",
line = sc.line,
keyword = %keyword,
"ACF: ignoring RULE with unsupported access keyword"
);
(RuleAccess::None, true)
};
let mut trap = false;
sc.skip_ws_comments();
if sc.peek() == Some(',') {
sc.next();
sc.skip_ws_comments();
let mut log_opt = String::new();
read_word(sc, &mut log_opt);
if log_opt == "TRAPWRITE" {
trap = true;
} else if log_opt != "NOTRAPWRITE" {
let got = sc.offending(&log_opt);
return Err(sc.reject(format!(
"RULE log option must be TRAPWRITE or NOTRAPWRITE, got '{got}'"
)));
}
}
sc.skip_ws_comments();
if sc.peek() == Some(')') {
sc.next();
}
let mut uag = Vec::new();
let mut hag = Vec::new();
let mut method = Vec::new();
let mut authority = Vec::new();
let mut calc: Option<(u32, String)> = None;
sc.skip_ws_comments();
if sc.peek() == Some('{') {
sc.next();
loop {
sc.skip_ws_comments();
match sc.peek() {
Some('}') => {
sc.next();
break;
}
Some(_) => {
let mut kw = String::new();
read_word(sc, &mut kw);
if kw == "UAG" {
let name = read_paren_name(sc)?;
uag.push(name);
} else if kw == "HAG" {
let name = read_paren_name(sc)?;
hag.push(name);
} else if kw == "METHOD" {
method.extend(read_paren_string_list(sc)?);
} else if kw == "AUTHORITY" {
authority.extend(read_paren_string_list(sc)?);
} else if kw == "CALC" {
let at = sc.line;
let expr = read_paren_name_raw(sc)?;
calc = Some((at, expr));
} else if kw.is_empty() {
sc.next();
} else {
tracing::warn!(
target: "epics_base_rs::access_security",
line = sc.line,
keyword = %kw,
"ACF: ignoring RULE with unsupported keyword — rule disabled"
);
ignore = true;
sc.skip_ws_comments();
if sc.peek() == Some('(') {
let _ = read_paren_name(sc)?;
}
}
}
None => break,
}
}
}
let mut inp_used: u32 = 0;
let calc_compiled = match calc {
Some((at, ref expr)) => {
let compiled = crate::calc::compile(expr)
.map_err(|e| sc.reject_at(at, format!("bad CALC expression '{expr}': {e}")))?;
let (used, stores) = compiled.arg_usage();
if stores != 0 {
return Err(sc.reject_at(
at,
format!("assignment operator used in CALC expression '{expr}'"),
));
}
inp_used = used;
Some(compiled)
}
None => None,
};
Ok(AccessRule {
level,
access,
uag,
hag,
method,
authority,
trap,
calc: calc.map(|(_, expr)| expr),
calc_compiled,
inp_used,
ignore,
})
}
fn read_paren_name_raw(sc: &mut AcfScanner) -> CaResult<String> {
sc.skip_ws_comments();
if sc.next() != Some('(') {
return Err(sc.reject("expected '(' after CALC"));
}
let opened = sc.line;
sc.skip_ws_comments();
let mut body = String::new();
if sc.peek() == Some('"') {
sc.next();
while let Some(c) = sc.peek() {
sc.next();
if c == '"' {
break;
}
body.push(c);
}
sc.skip_ws_comments();
if sc.next() != Some(')') {
return Err(sc.reject_at(opened, "expected ')' after CALC expression"));
}
} else {
while let Some(c) = sc.peek() {
if c == ')' {
sc.next();
break;
}
body.push(c);
sc.next();
}
}
Ok(body.trim().to_string())
}
fn read_paren_string_list(sc: &mut AcfScanner) -> CaResult<Vec<String>> {
sc.skip_ws_comments();
if sc.next() != Some('(') {
return Err(sc.reject("expected '(' after METHOD/AUTHORITY"));
}
let opened = sc.line;
let mut items = Vec::new();
let mut current = String::new();
let mut in_quotes = false;
loop {
match sc.peek() {
Some('"') => {
sc.next();
in_quotes = !in_quotes;
}
Some(')') if !in_quotes => {
sc.next();
break;
}
Some(',') if !in_quotes => {
sc.next();
let trimmed = current.trim().to_string();
if !trimmed.is_empty() {
items.push(trimmed);
}
current.clear();
}
Some(c) => {
current.push(c);
sc.next();
}
None => {
return Err(sc.reject_at(opened, "unterminated METHOD/AUTHORITY list"));
}
}
}
let trimmed = current.trim().to_string();
if !trimmed.is_empty() {
items.push(trimmed);
}
Ok(items)
}
#[cfg(test)]
mod tests {
use super::*;
static AS_CHECK_CLIENT_IP_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn lock_as_check_client_ip() -> std::sync::MutexGuard<'static, ()> {
AS_CHECK_CLIENT_IP_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner())
}
#[test]
fn test_parse_acf_basic() {
let _guard = lock_as_check_client_ip();
let acf = r#"
UAG(admins) { user1, user2 }
HAG(operators) { host1, host2 }
ASG(DEFAULT) {
RULE(1, WRITE) { UAG(admins) HAG(operators) }
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(config.uag.get("admins").unwrap(), &["user1", "user2"]);
assert_eq!(config.hag.get("operators").unwrap(), &["host1", "host2"]);
assert!(config.asg.contains_key("DEFAULT"));
assert_eq!(config.asg["DEFAULT"].rules.len(), 2);
}
#[test]
fn test_parse_acf_hag_uag() {
let _guard = lock_as_check_client_ip();
let acf = r#"
UAG(ops) { alice, bob }
HAG(lab) { lab-pc1.invalid }
ASG(SECURE) {
RULE(1, WRITE) { UAG(ops) HAG(lab) }
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(config.uag["ops"], vec!["alice", "bob"]);
assert_eq!(config.hag["lab"], vec!["lab-pc1.invalid"]);
}
#[test]
fn hag_stores_names_by_default() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(false);
let acf = r#"
HAG(local) { LocalHost }
ASG(DEFAULT) {
RULE(1, WRITE) { HAG(local) }
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.hag["local"],
vec!["localhost"],
"C stores the lowercased literal name and resolves nothing"
);
assert_eq!(
config.check_access("DEFAULT", "localhost", "alice"),
AccessLevel::ReadWrite,
"the claimed host name matches the HAG"
);
assert_eq!(
config.check_access("DEFAULT", "127.0.0.1", "alice"),
AccessLevel::NoAccess,
"a peer IP does not match a name HAG — that is what asCheckClientIP=1 is for"
);
}
#[test]
fn hag_stores_resolved_ips_under_as_check_client_ip() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(true);
let acf = r#"
HAG(local) { localhost }
ASG(DEFAULT) {
RULE(1, WRITE) { HAG(local) }
}
"#;
let config = parse_acf(acf).unwrap();
set_as_check_client_ip(false); assert_eq!(
config.hag["local"],
vec!["127.0.0.1"],
"C resolves the host to its IP under asCheckClientIP"
);
assert_eq!(
config.check_access("DEFAULT", "127.0.0.1", "alice"),
AccessLevel::ReadWrite,
"the peer IP matches the resolved HAG"
);
}
#[test]
fn hag_unresolvable_under_as_check_client_ip_becomes_sentinel() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(true);
let config = parse_acf("HAG(lab) { lab-pc1.invalid }\n").unwrap();
set_as_check_client_ip(false);
assert_eq!(config.hag["lab"], vec!["unresolved:lab-pc1.invalid"]);
}
#[test]
fn hag_unresolvable_name_does_not_abort_parser() {
let _guard = lock_as_check_client_ip();
let acf = r#"
HAG(quarantine) { gone.invalid, alive.invalid }
ASG(DEFAULT) {
RULE(1, WRITE) { HAG(quarantine) }
}
"#;
let config = parse_acf(acf).expect("parser must not abort on bad DNS");
let entries = &config.hag["quarantine"];
assert_eq!(
entries.len(),
2,
"literal entries preserved verbatim; no resolved IPs appended"
);
assert_eq!(entries[0], "gone.invalid");
assert_eq!(entries[1], "alive.invalid");
}
#[test]
fn with_refreshed_hags_recovers_a_stale_resolution() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(true);
let mut config = parse_acf("HAG(local) { localhost }\n").unwrap();
assert_eq!(config.hag_raw["local"], vec!["localhost"]);
config.hag.insert("local".into(), vec!["192.0.2.1".into()]);
let refreshed = config
.with_refreshed_hags()
.expect("a moved resolution must produce a refreshed config");
set_as_check_client_ip(false);
assert_eq!(refreshed.hag["local"], vec!["127.0.0.1"]);
assert_eq!(
refreshed.hag_raw["local"],
vec!["localhost"],
"raw spellings survive the refresh for the next round"
);
}
#[test]
fn with_refreshed_hags_is_none_when_resolution_is_unchanged() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(true);
let config = parse_acf("HAG(local) { localhost }\n").unwrap();
let idempotent = config.with_refreshed_hags();
set_as_check_client_ip(false);
assert!(
idempotent.is_none(),
"a freshly parsed config re-resolves to itself"
);
}
#[test]
fn with_refreshed_hags_is_none_in_name_mode() {
let _guard = lock_as_check_client_ip();
set_as_check_client_ip(false);
let config = parse_acf("HAG(local) { LocalHost }\n").unwrap();
assert_eq!(config.hag_raw["local"], vec!["LocalHost"]);
assert_eq!(config.hag["local"], vec!["localhost"]);
assert!(config.with_refreshed_hags().is_none());
}
#[test]
fn test_check_access_default_rw() {
let acf = "ASG(DEFAULT) { RULE(1, WRITE) RULE(1, READ) }";
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("DEFAULT", "host1", "user1"),
AccessLevel::ReadWrite
);
}
#[test]
fn test_check_access_read_only() {
let acf = r#"
UAG(admins) { admin1 }
ASG(READONLY) {
RULE(1, READ)
RULE(1, WRITE) { UAG(admins) }
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("READONLY", "host1", "admin1"),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access("READONLY", "host1", "regular"),
AccessLevel::Read
);
}
#[test]
fn test_check_access_hag_uag_match() {
let _guard = lock_as_check_client_ip();
let acf = r#"
UAG(ops) { alice }
HAG(lab) { lab-pc1 }
ASG(CONTROLLED) {
RULE(1, WRITE) { UAG(ops) HAG(lab) }
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("CONTROLLED", "lab-pc1", "alice"),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access("CONTROLLED", "other-host", "alice"),
AccessLevel::Read
);
assert_eq!(
config.check_access("CONTROLLED", "lab-pc1", "bob"),
AccessLevel::Read
);
}
#[test]
fn test_check_access_unknown_user() {
let acf = r#"
ASG(DEFAULT) {
RULE(1, WRITE)
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("DEFAULT", "", ""),
AccessLevel::ReadWrite
);
}
#[test]
fn dump_report_quotes_uag_and_hag_members() {
let cfg =
parse_acf("UAG(special) { someone, \"role/op\", \"a\\\"b\" }\nHAG(hosts) { HostA }\n")
.unwrap();
let dump = cfg.dump_report();
assert!(dump.contains("\t\"someone\"\n"), "{dump}");
assert!(dump.contains("\t\"role/op\"\n"), "{dump}");
assert!(dump.contains("\t\"a\\\"b\"\n"), "{dump}");
assert!(dump.contains("\t\"hosta\"\n"), "{dump}");
}
#[test]
fn parse_acf_captures_method_and_authority() {
let acf = r#"
ASG(SECURE) {
RULE(1, WRITE) {
METHOD("ca", "x509")
AUTHORITY("ANL CA")
}
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
let asg = &config.asg["SECURE"];
assert_eq!(asg.rules.len(), 2);
assert_eq!(asg.rules[0].method, vec!["ca", "x509"]);
assert_eq!(asg.rules[0].authority, vec!["ANL CA"]);
assert!(
asg.rules[1].method.is_empty(),
"READ rule must not inherit METHOD list",
);
assert!(asg.rules[1].authority.is_empty());
}
#[test]
fn tls_x509_acf_rule_grants_write_on_issuer_match() {
let cfg = parse_acf(
r#"
ASG(TLS_ONLY) {
RULE(1, WRITE) { METHOD("x509") AUTHORITY("CN=ops-ca, O=Lab") }
RULE(1, READ)
}
"#,
)
.unwrap();
assert_eq!(
cfg.check_access_method("TLS_ONLY", "h", "u", 0, "", ""),
AccessLevel::Read
);
assert_eq!(
cfg.check_access_method("TLS_ONLY", "h", "u", 0, "x509", "CN=other-ca"),
AccessLevel::Read
);
assert_eq!(
cfg.check_access_method("TLS_ONLY", "h", "u", 0, "x509", "CN=ops-ca, O=Lab"),
AccessLevel::ReadWrite
);
}
#[test]
fn check_access_method_gates_on_method() {
let acf = r#"
ASG(METHOD_GATED) {
RULE(1, WRITE) {
METHOD("x509")
}
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access_method("METHOD_GATED", "h", "u", 0, "x509", ""),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access_method("METHOD_GATED", "h", "u", 0, "ca", ""),
AccessLevel::Read
);
}
#[test]
fn check_access_method_gates_on_authority() {
let acf = r#"
ASG(AUTH_GATED) {
RULE(1, WRITE) {
AUTHORITY("Trusted Root")
}
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access_method("AUTH_GATED", "h", "u", 0, "x509", "Trusted Root"),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access_method("AUTH_GATED", "h", "u", 0, "x509", "Other CA"),
AccessLevel::Read
);
}
#[test]
fn check_access_asl_legacy_path_matches_when_method_empty() {
let acf = r#"
ASG(LEGACY) {
RULE(1, WRITE)
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access_asl("LEGACY", "h", "u", 0),
AccessLevel::ReadWrite
);
}
#[test]
fn check_access_method_match_is_case_insensitive() {
let acf = r#"
ASG(MIXED_CASE) {
RULE(1, WRITE) {
METHOD("X509")
}
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access_method("MIXED_CASE", "h", "u", 0, "x509", ""),
AccessLevel::ReadWrite
);
}
#[test]
fn empty_rule_asg_denies_access() {
let config = parse_acf("ASG(LOCKED) { }").unwrap();
assert_eq!(
config.check_access("LOCKED", "host", "user"),
AccessLevel::NoAccess,
"ASG with no RULE must deny — C asComputePvt fails closed"
);
}
#[test]
fn unknown_asg_falls_back_to_empty_default_and_denies() {
let config = parse_acf("UAG(ops) { alice }").unwrap();
assert!(config.asg.contains_key("DEFAULT"));
assert_eq!(
config.check_access("TYPO", "host", "alice"),
AccessLevel::NoAccess,
"unknown ASG must resolve to empty DEFAULT ⇒ NoAccess"
);
}
#[test]
fn default_asg_without_rules_denies() {
let config = parse_acf("UAG(ops) { alice }").unwrap();
assert_eq!(
config.check_access("DEFAULT", "host", "alice"),
AccessLevel::NoAccess
);
}
#[test]
fn a_rejected_token_is_quoted_as_the_operator_wrote_it() {
for (acf, got) in [
("ASG(G) { RULE(abc, READ) }", "abc"),
("ASG(G) { RULE(-abc, READ) }", "-abc"),
("ASG(G) { RULE(1, READ, %%%) }", "%%%"),
("ASG(G) { RULE(1, READ, trapwrite) }", "trapwrite"),
("ASG(G) { RULE(1, READ, TRAP%WRITE) }", "TRAP%WRITE"),
] {
let e = parse_acf(acf).unwrap_err().to_string();
assert!(
e.contains(&format!("got '{got}'")),
"expected got '{got}', of {acf:?}, got {e:?}"
);
}
let long = "z".repeat(4096);
let e = parse_acf(&format!("ASG(G) {{ RULE({long}, READ) }}"))
.unwrap_err()
.to_string();
assert!(e.contains('…') && e.len() < 200, "uncapped quote: {e:?}");
}
#[test]
fn every_parse_rejection_names_its_line() {
for (acf, line, needle) in [
(
"UAG(a) { x }\n\nUAG(my group) { b }\n",
3,
"whitespace inside",
),
("ASG(G) {\n RULE(-1, READ)\n}\n", 2, "must be positive"),
(
"ASG(G) {\n RULE(1, READ, BOGUS)\n}\n",
2,
"TRAPWRITE or NOTRAPWRITE",
),
("UAG(a) { x }\n}\n", 2, "top-level block keyword"),
] {
let e = parse_acf(acf).unwrap_err().to_string();
assert!(
e.contains(&format!("ACF line {line}:")) && e.contains(needle),
"expected line {line} and {needle:?}, got {e:?}"
);
}
let tail = "\n\n\n\n\n";
for (acf, line, needle) in [
(format!("UAG(a{tail}"), 1, "missing ')'"),
(format!("UAG(\"abc{tail}"), 1, "unterminated quoted name"),
(format!("UAG(a) {{ x{tail}"), 1, "unterminated '{'"),
(
format!("UAG(a) {{ \"unterminated{tail}"),
1,
"unterminated quoted string",
),
(
format!("ASG(G) {{ INPA(\"g\"){tail}"),
1,
"unterminated ASG",
),
(format!("BOGUS{tail}"), 1, "unexpected token 'BOGUS'"),
(
"ASG(G) {\n RULE(1, WRITE) {\n CALC(\"A:=1;1\")\n }\n}\n".to_string(),
3,
"assignment operator",
),
(
"ASG(G) {\n RULE(1, WRITE) {\n CALC(\"A+\")\n }\n}\n".to_string(),
3,
"bad CALC expression",
),
] {
let e = parse_acf(&acf).unwrap_err().to_string();
assert!(
e.contains(&format!("ACF line {line}:")) && e.contains(needle),
"expected line {line} and {needle:?}, got {e:?}"
);
}
}
#[test]
fn empty_acf_denies_all_access() {
let _guard = lock_as_check_client_ip();
for acf in ["", "# just a comment\n", "UAG(ops){alice}\nHAG(h){pc1}\n"] {
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("DEFAULT", "host", "alice"),
AccessLevel::NoAccess,
"empty/rule-less ACF must deny (input was {acf:?})"
);
assert_eq!(
config.check_access("ANY_GROUP", "host", "alice"),
AccessLevel::NoAccess,
"unknown ASG against empty ACF must deny (input was {acf:?})"
);
}
}
#[test]
fn handbuilt_config_missing_default_denies() {
let config = AccessSecurityConfig {
uag: HashMap::new(),
hag: HashMap::new(),
hag_raw: HashMap::new(),
asg: HashMap::new(),
unknown_access: AccessLevel::Read,
};
assert_eq!(
config.check_access("WHATEVER", "host", "user"),
AccessLevel::NoAccess
);
}
#[test]
fn rule_none_grants_no_access() {
let config = parse_acf("ASG(N) { RULE(0, NONE) }").unwrap();
assert_eq!(
config.check_access("N", "host", "user"),
AccessLevel::NoAccess
);
}
#[test]
fn rule_unsupported_access_keyword_is_inert() {
let config = parse_acf("ASG(B) { RULE(0, WRIET) }").unwrap();
assert_eq!(config.asg["B"].rules.len(), 1);
assert!(config.asg["B"].rules[0].ignore, "bad keyword ⇒ inert rule");
assert_eq!(
config.check_access("B", "host", "user"),
AccessLevel::NoAccess
);
}
#[test]
fn rule_negative_level_is_rejected() {
let err = parse_acf("ASG(X) { RULE(-1, READ) }");
assert!(err.is_err(), "negative RULE level must fail the parse");
}
#[test]
fn rule_non_numeric_level_is_rejected() {
let err = parse_acf("ASG(X) { RULE(abc, READ) }");
assert!(err.is_err(), "non-numeric RULE level must fail the parse");
}
#[test]
fn unknown_top_level_block_is_skipped_not_fatal() {
let acf = r#"
VENDOR(extension) { whatever }
ASG(DEFAULT) { RULE(1, READ) }
"#;
let config = parse_acf(acf).expect("unknown top-level block must not abort the parse");
assert_eq!(
config.check_access("DEFAULT", "host", "user"),
AccessLevel::Read,
"the ASG after the unknown block must still parse"
);
}
#[test]
fn unknown_well_formed_block_parses_ok_with_warning() {
let acf = r#"
VENDOR(x) { FOO(1) }
ASG(DEFAULT) { RULE(1, READ) }
"#;
let config = parse_acf(acf)
.expect("a well-formed unknown top-level block must warn-and-continue, not fail");
assert_eq!(
config.check_access("DEFAULT", "host", "user"),
AccessLevel::Read
);
}
#[test]
fn unknown_block_bare_head_parses_ok() {
let acf = "VENDOR(x) ASG(DEFAULT) { RULE(1, READ) }";
let config = parse_acf(acf).expect("bare unknown-block head must warn-and-continue");
assert!(config.asg.contains_key("DEFAULT"));
}
#[test]
fn genuine_garbage_acf_is_rejected() {
assert!(
parse_acf("this is not valid ACF (((").is_err(),
"unparseable ACF must fail, not silently skip to EOF"
);
}
#[test]
fn stray_top_level_punctuation_is_rejected() {
assert!(
parse_acf("(((").is_err(),
"a file of only '(((' must fail, not silently skip to EOF"
);
assert!(
parse_acf("}").is_err(),
"a file of only '}}' must fail, not silently skip to EOF"
);
}
#[test]
fn empty_and_comment_only_acf_still_parses_ok() {
assert!(parse_acf("").is_ok(), "empty file must parse Ok");
assert!(
parse_acf(" \n\t \n").is_ok(),
"whitespace-only file must parse Ok"
);
assert!(
parse_acf("# just a comment\n# another\n").is_ok(),
"comment-only file must parse Ok"
);
}
#[test]
fn unknown_keyword_without_paren_head_is_rejected() {
assert!(parse_acf("VENDOR something").is_err());
}
#[test]
fn unknown_keyword_at_eof_is_rejected() {
assert!(parse_acf("VENDOR").is_err());
}
#[test]
fn unknown_block_unbalanced_paren_is_rejected() {
assert!(parse_acf("VENDOR(((").is_err());
}
#[test]
fn unknown_block_unbalanced_brace_is_rejected() {
assert!(parse_acf("VENDOR(x) { unterminated").is_err());
}
#[test]
fn hag_host_match_is_case_insensitive() {
let _guard = lock_as_check_client_ip();
let acf = r#"
HAG(lab) { LabPC1.invalid }
ASG(C) {
RULE(1, WRITE) { HAG(lab) }
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
assert_eq!(
config.check_access("C", "labpc1.invalid", "user"),
AccessLevel::ReadWrite,
"lowercased HAG entry must match a mixed-case client host"
);
assert_eq!(
config.check_access("C", "LABPC1.INVALID", "user"),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access("C", "other.invalid", "user"),
AccessLevel::Read
);
}
#[test]
fn rule_trapwrite_log_option_parses() {
let config =
parse_acf("ASG(T) { RULE(1, WRITE, TRAPWRITE) RULE(1, READ, NOTRAPWRITE) }").unwrap();
assert_eq!(config.asg["T"].rules.len(), 2);
assert_eq!(config.asg["T"].rules[0].access, RuleAccess::Write);
assert!(
config.asg["T"].rules[0].trap,
"TRAPWRITE must set the trap mask"
);
assert_eq!(config.asg["T"].rules[1].access, RuleAccess::Read);
assert!(
!config.asg["T"].rules[1].trap,
"NOTRAPWRITE must clear the trap mask"
);
}
#[test]
fn rule_bad_log_option_is_rejected() {
assert!(parse_acf("ASG(T) { RULE(1, WRITE, BOGUS) }").is_err());
}
#[test]
fn rule_log_option_is_case_sensitive() {
assert!(
parse_acf("ASG(T) { RULE(1, WRITE, trapwrite) }").is_err(),
"lowercase `trapwrite` is not a valid log option (C strcmp)"
);
assert!(
parse_acf("ASG(T) { RULE(1, WRITE, notrapwrite) }").is_err(),
"lowercase `notrapwrite` is not a valid log option (C strcmp)"
);
}
#[test]
fn rule_access_keyword_is_case_sensitive() {
let cfg = parse_acf("ASG(L) { RULE(1, write) }").unwrap();
assert_eq!(
cfg.check_access_method("L", "h", "u", 0, "", ""),
AccessLevel::NoAccess,
"lowercase `write` is an unsupported keyword (C strcmp); grants nothing"
);
let cfg = parse_acf("ASG(U) { RULE(1, WRITE) }").unwrap();
assert_eq!(
cfg.check_access_method("U", "h", "u", 0, "", ""),
AccessLevel::ReadWrite
);
}
#[test]
fn mr_r20_trap_mask_reflects_matched_rule() {
let cfg = parse_acf(
r#"
ASG(TRAPPED) { RULE(0, WRITE, TRAPWRITE) }
ASG(UNTRAPPED) { RULE(0, WRITE, NOTRAPWRITE) }
ASG(PLAIN) { RULE(0, WRITE) }
ASG(LOCKED) { }
"#,
)
.unwrap();
let (lvl, trap) = cfg.check_access_method_trap("TRAPPED", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::ReadWrite);
assert!(trap, "a TRAPWRITE rule must resolve rule_was_trap = true");
let (lvl, trap) = cfg.check_access_method_trap("UNTRAPPED", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::ReadWrite);
assert!(
!trap,
"a NOTRAPWRITE rule must resolve rule_was_trap = false"
);
let (lvl, trap) = cfg.check_access_method_trap("PLAIN", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::ReadWrite);
assert!(
!trap,
"a rule with no trap option must resolve rule_was_trap = false"
);
let (lvl, trap) = cfg.check_access_method_trap("LOCKED", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::NoAccess);
assert!(
!trap,
"a denied resolution must carry rule_was_trap = false"
);
}
#[test]
fn mr_r20_trap_mask_follows_last_access_raising_rule() {
let cfg = parse_acf("ASG(M) { RULE(0, READ) RULE(0, WRITE, TRAPWRITE) }").unwrap();
let (lvl, trap) = cfg.check_access_method_trap("M", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::ReadWrite);
assert!(
trap,
"trap mask must follow the WRITE rule that raised access"
);
let cfg = parse_acf("ASG(N) { RULE(0, READ) RULE(0, WRITE, NOTRAPWRITE) }").unwrap();
let (lvl, trap) = cfg.check_access_method_trap("N", "h", "u", 0, "", "");
assert_eq!(lvl, AccessLevel::ReadWrite);
assert!(!trap, "NOTRAPWRITE on the access-raising rule must win");
}
#[test]
fn calc_rule_is_conditionally_active_and_fails_closed_without_resolver() {
let config = parse_acf(r#"ASG(G) { INPA("ref") RULE(1, WRITE) { CALC("A=1") } }"#).unwrap();
let rule = &config.asg["G"].rules[0];
assert!(rule.calc.is_some(), "CALC clause must be parsed and stored");
assert!(
!rule.ignore,
"a CALC rule is conditionally active, not unconditionally ignored"
);
assert_eq!(
config.check_access("G", "host", "user"),
AccessLevel::NoAccess,
"CALC rule with no resolver must not grant WRITE"
);
}
#[test]
fn inp_link_targets_are_deduplicated_across_groups() {
let cfg = parse_acf(
r#"
ASG(A) { INPA("gate") INPB("gate.RVAL") RULE(1, WRITE) { CALC("A") } }
ASG(B) { INPA("gate") INPB("other.SEVR") RULE(1, WRITE) { CALC("A") } }
"#,
)
.expect("parse");
assert_eq!(
cfg.inp_link_targets(),
vec![
("gate".to_string(), "RVAL".to_string()),
("gate".to_string(), "VAL".to_string()),
("other".to_string(), "SEVR".to_string()),
],
"`gate` is read by both groups but is one subscription"
);
}
#[epics_macros_rs::epics_test]
async fn calc_gated_rule_evaluates_against_inp_resolver() {
use std::sync::Arc;
let cfg =
parse_acf(r#"ASG(OPS) { INPA("permit.VAL") RULE(1, WRITE) { CALC("A=1") } }"#).unwrap();
let cell = crate::server::access_security::new_acf_cell(Some(cfg));
let asg_resolver: AsgAslResolver =
Arc::new(|_name| Box::pin(async { ("OPS".to_string(), 0u8) }));
let grant = AccessGate::required(cell.clone(), asg_resolver.clone()).with_inp_resolver(
Arc::new(|link: String| Box::pin(async move { (link == "permit.VAL").then_some(1.0) })),
);
assert!(
grant.check("x", "h", "u", "ca", "").await.allows_write(),
"CALC A=1 with permit=1 grants WRITE"
);
let deny = AccessGate::required(cell.clone(), asg_resolver.clone()).with_inp_resolver(
Arc::new(|link: String| Box::pin(async move { (link == "permit.VAL").then_some(0.0) })),
);
assert!(
!deny.check("x", "h", "u", "ca", "").await.allows_write(),
"CALC A=1 with permit=0 denies WRITE"
);
let bad = AccessGate::required(cell.clone(), asg_resolver.clone())
.with_inp_resolver(Arc::new(|_link: String| Box::pin(async move { None })));
assert!(
!bad.check("x", "h", "u", "ca", "").await.allows_write(),
"a bad/disconnected INP denies the CALC-gated rule"
);
let none = AccessGate::required(cell, asg_resolver);
assert!(
!none.check("x", "h", "u", "ca", "").await.allows_write(),
"no INP resolver installed → CALC rule fails closed"
);
}
#[test]
fn uag_role_member_matches_client_role() {
let cfg =
parse_acf(r#"UAG(special) { "role/op" } ASG(G) { RULE(1, WRITE) { UAG(special) } }"#)
.unwrap();
let (lvl, _) =
cfg.compute_for_name("G", "h", "acct", &["op".to_string()], 0, "ca", "", None);
assert_eq!(
lvl,
AccessLevel::ReadWrite,
"role/op member matches a client holding role 'op'"
);
let (lvl_none, _) = cfg.compute_for_name("G", "h", "acct", &[], 0, "ca", "", None);
assert_eq!(
lvl_none,
AccessLevel::NoAccess,
"a client without role 'op' must not match role/op"
);
}
#[test]
fn calc_rule_with_bad_expression_is_rejected() {
assert!(
parse_acf(r#"ASG(G) { RULE(1, WRITE) { CALC("A=") } }"#).is_err(),
"syntactically broken CALC must fail the parse"
);
}
#[test]
fn asg_inp_links_are_parsed() {
let acf = r#"
ASG(G) {
INPA("rec1.VAL")
INPC("rec3.VAL")
RULE(1, READ)
}
"#;
let config = parse_acf(acf).unwrap();
let inp = &config.asg["G"].inp;
assert_eq!(inp.len(), 2);
assert_eq!(inp[0].index, 0);
assert_eq!(inp[0].link, "rec1.VAL");
assert_eq!(inp[1].index, 2);
assert_eq!(inp[1].link, "rec3.VAL");
}
#[test]
fn asg_inp_bad_selector_is_rejected() {
assert!(parse_acf(r#"ASG(G) { INPZ("x") }"#).is_err());
}
#[test]
fn asg_inp_selector_is_case_sensitive() {
assert!(
parse_acf(r#"ASG(G) { INPa("x") }"#).is_err(),
"lowercase INP selector must be rejected (C flex [A-U])"
);
}
#[test]
fn paren_name_rejects_embedded_whitespace() {
assert!(parse_acf("UAG(my group) { x }").is_err());
}
#[test]
fn paren_name_rejects_unterminated() {
assert!(parse_acf("UAG(unterminated").is_err());
}
#[test]
fn paren_name_accepts_quoted_form() {
let config = parse_acf(r#"UAG("my group") { x }"#).unwrap();
assert!(config.uag.contains_key("my group"));
}
#[test]
fn asl_gate_still_honoured_after_fail_closed_rewrite() {
let config = parse_acf("ASG(A) { RULE(0, READ) RULE(1, WRITE) }").unwrap();
assert_eq!(
config.check_access_method("A", "h", "u", 0, "", ""),
AccessLevel::ReadWrite
);
assert_eq!(
config.check_access_method("A", "h", "u", 2, "", ""),
AccessLevel::NoAccess
);
}
fn trap_capture(
pv: &'static str,
) -> (
std::sync::Arc<std::sync::Mutex<Vec<(TrapWriteOp, Option<String>)>>>,
TrapWriteListenerHandle,
) {
let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let sink = events.clone();
let handle = register_trap_write_listener(std::sync::Arc::new(move |msg| {
if msg.pv_name == pv {
sink.lock()
.unwrap()
.push((msg.op, msg.status.map(str::to_owned)));
}
}));
(events, handle)
}
fn trap_fields(pv: &'static str) -> TrapWriteFields {
TrapWriteFields {
pv_name: pv.to_string(),
user: "u".to_string(),
host: "h".to_string(),
peer: "h:5064".to_string(),
value_str: "42".to_string(),
dbr_type: 5,
no_elements: 1,
event_id: next_trap_write_event_id(),
rule_was_trap: true,
cancel_status: "superseded".to_string(),
}
}
#[test]
fn trap_write_guard_complete_fires_one_after_and_disarms_drop() {
let (events, _handle) = trap_capture("guard:complete");
{
let mut guard = TrapWriteGuard::begin(trap_fields("guard:complete"));
guard.complete("ok");
} let got = events.lock().unwrap().clone();
assert_eq!(
got,
vec![
(TrapWriteOp::BeforeWrite, None),
(TrapWriteOp::AfterWrite, Some("ok".to_string())),
]
);
}
#[test]
fn trap_write_guard_drop_without_complete_fires_cancel_after() {
let (events, _handle) = trap_capture("guard:cancel");
{
let _guard = TrapWriteGuard::begin(trap_fields("guard:cancel"));
}
let got = events.lock().unwrap().clone();
assert_eq!(
got,
vec![
(TrapWriteOp::BeforeWrite, None),
(TrapWriteOp::AfterWrite, Some("superseded".to_string())),
]
);
}
}
#[cfg(test)]
mod as_ca_task_tests {
use super::*;
#[epics_macros_rs::epics_test]
async fn the_as_ca_task_row_tracks_the_loaded_policy() {
fn table() -> String {
let out = std::cell::RefCell::new(String::new());
crate::runtime::taskwd::taskwd_show(1, &|line| {
out.borrow_mut().push_str(line);
out.borrow_mut().push('\n');
});
out.into_inner()
}
async fn wait_until(want: bool) {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10);
while table().contains("asCaTask") != want {
assert!(
std::time::Instant::now() < deadline,
"`asCaTask` never became {}:\n{}",
if want { "present" } else { "absent" },
table()
);
crate::runtime::task::sleep_background(std::time::Duration::from_millis(10)).await;
}
}
let db = std::sync::Arc::new(crate::server::database::PvDatabase::new());
let cell = new_acf_cell_watching(None, &db);
assert!(
!table().contains("asCaTask"),
"a cell built with no policy registered the access-security task"
);
cell.store(Some(std::sync::Arc::new(
parse_acf("ASG(DEFAULT) { RULE(1, READ) }").expect("minimal ACF"),
)));
wait_until(true).await;
cell.store(None);
wait_until(false).await;
}
}