use crate::journal::frame::SaturatingFrom;
use std::fmt;
use super::{Coverage, MAX_FIELD_NAME_BYTES, Outcome, Provenance, Refusal, Source, Surface};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct PlanLimits {
pub max_bytes: usize,
pub max_field_bytes: usize,
pub max_fields: u32,
}
impl Default for PlanLimits {
fn default() -> Self {
Self {
max_bytes: Self::default_bytes(),
max_field_bytes: Self::default_field_bytes(),
max_fields: 8,
}
}
}
impl PlanLimits {
#[must_use]
pub const fn new(max_bytes: usize, max_field_bytes: usize, max_fields: u32) -> Self {
Self {
max_bytes,
max_field_bytes,
max_fields,
}
}
#[must_use]
pub const fn with_max_bytes(mut self, max_bytes: usize) -> Self {
self.max_bytes = max_bytes;
self
}
#[must_use]
pub const fn with_max_field_bytes(mut self, max_field_bytes: usize) -> Self {
self.max_field_bytes = max_field_bytes;
self
}
#[must_use]
pub const fn with_max_fields(mut self, max_fields: u32) -> Self {
self.max_fields = max_fields;
self
}
#[must_use]
pub const fn default_bytes() -> usize {
super::MAX_PLAN_BYTES
}
#[must_use]
pub const fn default_field_bytes() -> usize {
super::MAX_FIELD_BYTES
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Declared {
ops: Vec<String>,
notes: Vec<String>,
paths: Vec<String>,
coverage: Option<String>,
}
impl Declared {
#[must_use]
pub fn ops(&self) -> &[String] {
&self.ops
}
#[must_use]
pub fn notes(&self) -> &[String] {
&self.notes
}
#[must_use]
pub fn paths(&self) -> &[String] {
&self.paths
}
#[must_use]
pub fn coverage(&self) -> Option<&str> {
self.coverage.as_deref()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.ops.is_empty() && self.paths.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Wanted {
operation: String,
path: Option<String>,
}
impl Wanted {
#[must_use]
pub fn operation(&self) -> &str {
&self.operation
}
#[must_use]
pub fn path(&self) -> Option<&str> {
self.path.as_deref()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Plan {
wanted: Vec<Wanted>,
notes: Vec<String>,
coverage: Coverage,
provenance: Provenance,
}
impl Plan {
#[must_use]
pub fn wanted(&self) -> &[Wanted] {
&self.wanted
}
#[must_use]
pub fn notes(&self) -> &[String] {
&self.notes
}
#[must_use]
pub const fn coverage(&self) -> Coverage {
self.coverage
}
#[must_use]
pub const fn provenance(&self) -> &Provenance {
&self.provenance
}
pub(crate) fn for_repair(fingerprint: &str, provenance: Provenance) -> Self {
Self {
wanted: Vec::new(),
notes: vec![format!("retry: {fingerprint}")],
coverage: Coverage::Partial {
covered: 0,
asked: 1,
},
provenance,
}
}
}
impl fmt::Display for Plan {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
for wanted in &self.wanted {
match wanted.path() {
Some(path) => write!(formatter, "{} {} ", wanted.operation(), path)?,
None => write!(formatter, "{} ", wanted.operation())?,
}
}
write!(formatter, "[{}]", self.coverage.label())
}
}
#[must_use]
pub fn payload_digest(payload: &[u8]) -> lgwks_std::hash::Digest {
super::framed_digest("lgwks-bot/proposal/plan", payload)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Decoder {
limits: PlanLimits,
}
impl Decoder {
#[must_use]
pub const fn new(limits: PlanLimits) -> Self {
Self { limits }
}
#[must_use]
pub const fn limits(&self) -> PlanLimits {
self.limits
}
#[must_use]
pub fn decode(&self, surface: &Surface, payload: &[u8], source: Source) -> Outcome {
let provenance = Provenance::of(source, surface.tenant(), payload);
if let Err(refusal) = self.check_size(payload) {
return Outcome::refused(provenance, refusal);
}
match self.read(surface, payload) {
Err(refusal) => Outcome::refused(provenance, refusal),
Ok(declared) => self.authorize(surface, &declared, provenance),
}
}
fn check_size(&self, payload: &[u8]) -> Result<(), Refusal> {
if payload.len() > self.limits.max_bytes {
let refusal = Err(Refusal::Oversized {
got: payload.len(),
limit: self.limits.max_bytes,
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "check_size: returning an error to the caller");
return refusal;
}
Ok(())
}
fn authorize(&self, surface: &Surface, declared: &Declared, provenance: Provenance) -> Outcome {
match self.plan(surface, declared, provenance.clone()) {
Ok(plan) => Outcome::Admitted { plan, provenance },
Err(refusal) => Outcome::refused(provenance, refusal),
}
}
fn plan(
&self,
surface: &Surface,
declared: &Declared,
provenance: Provenance,
) -> Result<Plan, Refusal> {
if declared.is_empty() {
let refusal = Err(Refusal::Empty);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "plan: returning an error to the caller");
return refusal;
}
for path in declared.paths() {
refuse_traversal(path)?;
}
let mut wanted = Vec::with_capacity(declared.ops().len());
for op in declared.ops() {
surface.authorize(op)?;
wanted.push(Wanted {
operation: op.clone(),
path: declared.paths().first().cloned(),
});
}
Ok(Plan {
wanted,
notes: declared.notes().to_vec(),
coverage: Coverage::from_claim(declared.coverage()),
provenance,
})
}
fn read(&self, surface: &Surface, payload: &[u8]) -> Result<Declared, Refusal> {
Reader {
bytes: payload,
cursor: 0,
limits: self.limits,
surface: Some(surface),
}
.document()
}
}
fn refuse_traversal(path: &str) -> Result<(), Refusal> {
if path.starts_with('/') || path.starts_with('\\') || path.contains(':') {
let refusal = Err(Refusal::escape("path", path));
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "refuse_traversal: returning an error to the caller");
return refusal;
}
if path.split('/').any(|segment| segment == "..") {
let refusal = Err(Refusal::escape("path", path));
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "refuse_traversal: returning an error to the caller");
return refusal;
}
Ok(())
}
struct Reader<'a> {
bytes: &'a [u8],
cursor: usize,
limits: PlanLimits,
surface: Option<&'a Surface>,
}
impl<'a> Reader<'a> {
fn document(&mut self) -> Result<Declared, Refusal> {
let mut declared = Declared::default();
let mut fields = 0_u32;
loop {
self.skip_newlines();
if self.cursor >= self.bytes.len() {
return Ok(declared);
}
let start = self.cursor;
let line = self.line()?;
if line.is_empty() {
continue;
}
fields = fields.saturating_add(1);
if fields > self.limits.max_fields {
let refusal = Err(Refusal::Limit {
what: "the number of fields in a plan",
got: u64::from(fields),
limit: u64::from(self.limits.max_fields),
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "document: returning an error to the caller");
return refusal;
}
self.field(&mut declared, line, start)?;
}
}
fn line(&mut self) -> Result<&'a [u8], Refusal> {
let start = self.cursor;
while self.cursor < self.bytes.len() {
if self.bytes[self.cursor] == b'\n' {
let line = &self.bytes[start..self.cursor];
self.cursor = self.cursor.saturating_add(1);
return Ok(trim_cr(line));
}
self.cursor = self.cursor.saturating_add(1);
let taken = self.cursor.saturating_sub(start);
if taken > self.line_ceiling() {
let refusal = Err(Refusal::Limit {
what: "the length of a line in a plan",
got: u64::saturating_from(taken),
limit: u64::saturating_from(self.line_ceiling()),
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "line: returning an error to the caller");
return refusal;
}
}
Ok(trim_cr(&self.bytes[start..self.cursor]))
}
fn line_ceiling(&self) -> usize {
self.limits
.max_field_bytes
.saturating_add(MAX_FIELD_NAME_BYTES)
.saturating_add(1)
}
fn field(&mut self, declared: &mut Declared, line: &[u8], start: usize) -> Result<(), Refusal> {
let Some(split) = line.iter().position(|byte| *byte == b'=') else {
let refusal = Err(Refusal::Malformed {
cause: "a line with no `=` separator",
at: start,
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "field: returning an error to the caller");
return refusal;
};
let raw_name = &line[..split];
let raw_value = &line[split.saturating_add(1)..];
self.check_len(
raw_name.len(),
MAX_FIELD_NAME_BYTES,
"a field name in a plan",
start,
)?;
self.check_len(
raw_value.len(),
self.limits.max_field_bytes,
"a field value in a plan",
start,
)?;
let name = decode_text(raw_name);
let value = decode_text(raw_value);
self.store(declared, &name, value, start)
}
fn check_len(
&self,
len: usize,
limit: usize,
what: &'static str,
at: usize,
) -> Result<(), Refusal> {
if len > limit {
let refusal = Err(Refusal::Malformed { cause: what, at });
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "check_len: returning an error to the caller");
return refusal;
}
Ok(())
}
fn store(
&mut self,
declared: &mut Declared,
name: &str,
value: String,
start: usize,
) -> Result<(), Refusal> {
match name {
"op" => declared.ops.push(value),
"note" => declared.notes.push(value),
"path" => declared.paths.push(value),
"coverage" => declared.coverage = Some(value),
"install" => {
let refusal = Err(Refusal::install_tool(&value));
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "store: returning an error to the caller");
return refusal;
}
"credential" => {
let refusal = Err(Refusal::credential_read(&value));
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "store: returning an error to the caller");
return refusal;
}
"host" => return self.refuse_foreign_host(&value, start),
_ => {
let refusal = Err(Refusal::Malformed {
cause: "a field this decoder does not recognise",
at: start,
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "store: returning an error to the caller");
return refusal;
}
}
Ok(())
}
fn refuse_foreign_host(&self, value: &str, at: usize) -> Result<(), Refusal> {
let Some(surface) = self.surface else {
let refusal = Err(Refusal::Malformed {
cause: "a `host` field outside a run surface",
at,
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "refuse_foreign_host: returning an error to the caller");
return refusal;
};
if value == surface.tenant() {
Ok(())
} else {
Err(Refusal::escape("host", value))
}
}
fn skip_newlines(&mut self) {
while self.cursor < self.bytes.len() && self.bytes[self.cursor] == b'\n' {
self.cursor = self.cursor.saturating_add(1);
}
}
}
fn trim_cr(line: &[u8]) -> &[u8] {
match line.split_last() {
Some((&b'\r', rest)) => rest,
Some(_) | None => line,
}
}
fn trim(value: &[u8]) -> &[u8] {
let start = value.iter().take_while(|byte| **byte == b' ').count();
let end = value.iter().rev().take_while(|byte| **byte == b' ').count();
&value[start..value.len().saturating_sub(end)]
}
fn decode_text(raw: &[u8]) -> String {
String::from_utf8_lossy(trim(raw)).into_owned()
}