use std::collections::{BTreeSet, HashMap, HashSet};
use futures::{
StreamExt,
stream::{self, BoxStream},
};
use openc2::{
Action, ActionTargets, Error, Feature, Message, Nsid, ProfileFeatures, TargetType, Value,
Version,
json::{Command, Headers, Response, Results, Target},
target::Features,
};
use crate::{Consume, util::stream_just};
pub struct ConsumerToken(usize);
#[derive(Default, Clone)]
pub struct Registration {
actions: HashMap<Option<Nsid>, ActionTargets>,
}
impl Registration {
pub fn new() -> Self {
Self {
actions: Default::default(),
}
}
pub fn with_actions(
mut self,
actions: impl IntoIterator<Item = (Nsid, Action, TargetType<'static>)>,
) -> Self {
for (nsid, action, target) in actions {
self.actions
.entry(Some(nsid))
.or_default()
.entry(action)
.or_default()
.insert(target);
}
self
}
pub fn with_actions_without_profile(
mut self,
actions: impl IntoIterator<Item = (Action, TargetType<'static>)>,
) -> Self {
for (action, target) in actions {
self.actions
.entry(None)
.or_default()
.entry(action)
.or_default()
.insert(target);
}
self
}
pub fn profiles(&self) -> impl Iterator<Item = &Nsid> {
self.actions.keys().flatten()
}
fn to_pairs(&self) -> impl Iterator<Item = (Action, TargetType<'static>)> {
self.actions
.values()
.flatten()
.flat_map(|(a, t)| t.iter().cloned().map(move |target| (*a, target)))
}
pub fn matches(&self, action: Action, target: &TargetType, profile: &Nsid) -> bool {
let Some(entry) = self.actions.get(&Some(profile.clone())) else {
return false;
};
entry
.get(&action)
.map(|set| set.contains(target))
.unwrap_or(false)
}
pub fn query_features(&self, features: &Features) -> Response {
if features.contains(&Feature::RateLimit) {
return Error::not_implemented("rate limit feature is not implemented")
.at("features")
.into();
}
let mut results = Results::default();
if features.contains(&Feature::Profiles) {
results.profiles = self.actions.keys().flatten().cloned().collect();
}
if features.contains(&Feature::Versions) {
results.versions = [Version::new(2, 0)].into_iter().collect();
}
if features.contains(&Feature::Pairs) {
results.pairs = Some(self.actions.values().cloned().fold(
ActionTargets::new(),
|mut acc, at| {
for (a, t) in &at {
for target in t {
acc.entry(*a).or_default().insert(target.clone());
}
}
acc
},
));
results.extensions = self
.actions
.iter()
.filter_map(|(k, v)| {
Some((
k.clone()?,
Value::from_typed(&ProfileFeatures { pairs: v.clone() }).unwrap(),
))
})
.collect();
}
results.into()
}
}
pub trait ToRegistration {
fn to_registration(&self) -> Registration;
}
impl<T: ToRegistration> ToRegistration for Box<T> {
fn to_registration(&self) -> Registration {
(**self).to_registration()
}
}
impl<T: ToRegistration> ToRegistration for std::sync::Arc<T> {
fn to_registration(&self) -> Registration {
(**self).to_registration()
}
}
struct RegEntry<T> {
registration: Registration,
value: T,
}
impl Consume for RegEntry<Box<dyn Consume + Send + Sync>> {
fn consume<'a>(&'a self, msg: Message<Headers, Command>) -> BoxStream<'a, Response> {
if let (Action::Query, Target::Features(features)) = msg.body.as_action_target() {
return stream_just(self.registration.query_features(features));
}
self.value.consume(msg)
}
}
impl<T: Consume + Send + Sync> Consume for RegEntry<T> {
fn consume<'a>(&'a self, msg: Message<Headers, Command>) -> BoxStream<'a, Response> {
if let (Action::Query, Target::Features(features)) = msg.body.as_action_target() {
return stream_just(self.registration.query_features(features));
}
self.value.consume(msg)
}
}
pub type BoxConsumer = Box<dyn Consume + Send + Sync>;
type RegistryEntry = RegEntry<BoxConsumer>;
#[derive(Default)]
pub struct Registry {
consumers: Vec<Option<RegistryEntry>>,
by_pair: HashMap<(Action, TargetType<'static>), BTreeSet<usize>>,
}
impl Registry {
pub fn add(
&mut self,
other: impl ToRegistration + Consume + Send + Sync + 'static,
) -> ConsumerToken {
self.insert(other.to_registration(), other)
}
pub fn insert(
&mut self,
registration: impl Into<Registration>,
consumer: impl Consume + Send + Sync + 'static,
) -> ConsumerToken {
self.insert_boxed(registration.into(), Box::new(consumer))
}
fn insert_boxed(&mut self, registration: Registration, consumer: BoxConsumer) -> ConsumerToken {
let idx = self.consumers.len();
for pair in registration.to_pairs() {
self.by_pair.entry(pair).or_default().insert(idx);
}
self.consumers.push(Some(RegistryEntry {
registration,
value: consumer,
}));
ConsumerToken(idx)
}
fn get_matching<'a, 'b>(&'a self, pair: &(Action, TargetType<'b>)) -> Vec<&'a RegistryEntry> {
let entry = self.by_pair.get(pair);
entry
.into_iter()
.flat_map(move |indices| {
indices
.iter()
.filter_map(|&idx| self.consumers[idx].as_ref())
})
.collect()
}
pub fn remove(&mut self, token: ConsumerToken) -> Option<(Registration, BoxConsumer)> {
let entry = self.consumers.get_mut(token.0)?.take()?;
for pair in entry.registration.to_pairs() {
if let Some(set) = self.by_pair.get_mut(&pair) {
set.remove(&token.0);
if set.is_empty() {
self.by_pair.remove(&pair);
}
}
}
Some((entry.registration, entry.value))
}
pub fn profiles(&self) -> HashSet<&Nsid> {
self.consumers
.iter()
.filter_map(|c| c.as_ref())
.flat_map(|c| c.registration.profiles())
.collect()
}
pub fn pairs(&self) -> ActionTargets {
let mut pairs = ActionTargets::new();
for (action, target) in self.by_pair.keys().cloned() {
pairs.entry(action).or_default().insert(target);
}
pairs
}
pub fn query_features(&self, features: &Features) -> Result<Response, Error> {
if features.contains(&Feature::RateLimit) {
return Err(
Error::not_implemented("rate limit feature is not implemented").at("features"),
);
}
let mut results = Results::default();
if features.contains(&Feature::Profiles) {
results.profiles = self.profiles().into_iter().cloned().collect();
}
if features.contains(&Feature::Versions) {
results.versions = [Version::new(2, 0)].into_iter().collect();
}
if features.contains(&Feature::Pairs) {
results.pairs = Some(self.pairs());
let mut profiles: HashMap<_, ActionTargets> = HashMap::new();
for consumer in self.consumers.iter().flatten() {
for (profile, actions) in &consumer.registration.actions {
let Some(profile) = profile else {
continue;
};
let profile_entry = profiles.entry(profile.clone()).or_default();
for (action, target) in actions {
profile_entry
.entry(*action)
.or_default()
.extend(target.clone());
}
}
}
results = results
.with_extensions(
profiles
.into_iter()
.map(|(ap, pairs)| (ap, ProfileFeatures { pairs })),
)
.map_err(|e| {
Error::custom(format!("unable to serialize profile-specific pairs: {e}"))
})?;
}
Ok(results.into())
}
}
impl FromIterator<(Registration, BoxConsumer)> for Registry {
fn from_iter<T: IntoIterator<Item = (Registration, BoxConsumer)>>(iter: T) -> Self {
let mut registry = Self::default();
for (registration, consumer) in iter {
registry.insert_boxed(registration, consumer);
}
registry
}
}
impl<I: Consume + Send + Sync + 'static> FromIterator<(Registration, I)> for Registry {
fn from_iter<T: IntoIterator<Item = (Registration, I)>>(iter: T) -> Self {
let mut registry = Self::default();
for (registration, item) in iter {
registry.insert_boxed(registration, Box::new(item));
}
registry
}
}
impl<T: Consume + Send + Sync + ToRegistration + 'static> FromIterator<T> for Registry {
fn from_iter<I: IntoIterator<Item = T>>(iter: I) -> Self {
let mut registry = Self::default();
for item in iter {
registry.insert(item.to_registration(), item);
}
registry
}
}
impl ToRegistration for Registry {
fn to_registration(&self) -> Registration {
let mut actions: HashMap<Option<Nsid>, ActionTargets> = HashMap::new();
for (profile, acts) in self
.consumers
.iter()
.flatten()
.flat_map(|c| &c.registration.actions)
{
let profile_entry = actions.entry(profile.clone()).or_default();
for (action, targets) in acts {
profile_entry
.entry(*action)
.or_default()
.extend(targets.iter().cloned());
}
}
Registration { actions }
}
}
impl Extend<(Registration, BoxConsumer)> for Registry {
fn extend<T: IntoIterator<Item = (Registration, BoxConsumer)>>(&mut self, iter: T) {
for (registration, consumer) in iter {
self.insert_boxed(registration, consumer);
}
}
}
impl IntoIterator for Registry {
type Item = (Registration, BoxConsumer);
type IntoIter = std::vec::IntoIter<Self::Item>;
fn into_iter(self) -> Self::IntoIter {
self.consumers
.into_iter()
.flatten()
.map(|entry| (entry.registration, entry.value))
.collect::<Vec<_>>()
.into_iter()
}
}
impl Consume for Registry {
fn consume<'a>(&'a self, msg: Message<Headers, Command>) -> BoxStream<'a, Response> {
if msg.body.action == Action::Query
&& let Target::Features(features) = &msg.body.target
{
return stream_just(match self.query_features(features) {
Ok(rsp) => rsp,
Err(e) => e.into(),
});
}
let action = msg.body.action;
let target_type = msg.body.target.kind();
let mut consumers = self.get_matching(&(action, target_type.clone()));
if consumers.is_empty() {
return stream_just(Error::not_implemented_pair(action, &target_type).into());
}
if let Some(profile) = &msg.body.profile {
consumers
.retain(|consumer| consumer.registration.matches(action, &target_type, profile));
}
if consumers.is_empty() {
return stream_just(Error::not_implemented(format!(
"No consumer for action '{action}' and target type '{target_type:?}' matches profile '{:?}'",
msg.body.profile
)).into());
}
stream::select_all(
consumers
.into_iter()
.map(|consumer| consumer.consume(msg.clone())),
)
.boxed()
}
}