use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use stateset_primitives::{CurrencyCode, PriceLevelId, ProductId};
use strum::{Display, EnumString};
#[derive(
Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default, Display, EnumString,
)]
#[serde(rename_all = "snake_case")]
#[strum(serialize_all = "snake_case", ascii_case_insensitive)]
#[non_exhaustive]
pub enum PriceAdjustmentType {
#[default]
None,
PercentageDiscount,
PercentageMarkup,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PriceLevel {
pub id: PriceLevelId,
pub name: String,
pub code: String,
pub description: Option<String>,
pub adjustment_type: PriceAdjustmentType,
pub adjustment_value: Decimal,
pub currency: CurrencyCode,
pub is_active: bool,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
impl PriceLevel {
#[must_use]
pub fn adjust(&self, base: Decimal) -> Decimal {
let hundred = Decimal::from(100);
let adjusted = match self.adjustment_type {
PriceAdjustmentType::None => base,
PriceAdjustmentType::PercentageDiscount => {
base - (base * self.adjustment_value / hundred)
}
PriceAdjustmentType::PercentageMarkup => {
base + (base * self.adjustment_value / hundred)
}
};
adjusted.max(Decimal::ZERO)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PriceLevelEntry {
pub price_level_id: PriceLevelId,
pub product_id: ProductId,
pub price: Decimal,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
#[must_use]
pub fn resolve_price(
level: &PriceLevel,
entry: Option<&PriceLevelEntry>,
base: Decimal,
) -> Decimal {
match entry {
Some(e) => e.price,
None => level.adjust(base),
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CreatePriceLevel {
pub name: String,
pub code: String,
pub description: Option<String>,
#[serde(default)]
pub adjustment_type: PriceAdjustmentType,
#[serde(default)]
pub adjustment_value: Decimal,
pub currency: Option<CurrencyCode>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct UpdatePriceLevel {
pub name: Option<String>,
pub description: Option<String>,
pub adjustment_type: Option<PriceAdjustmentType>,
pub adjustment_value: Option<Decimal>,
pub is_active: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct PriceLevelFilter {
pub is_active: Option<bool>,
pub limit: Option<u32>,
pub offset: Option<u32>,
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal_macros::dec;
fn make(adjustment_type: PriceAdjustmentType, value: Decimal) -> PriceLevel {
PriceLevel {
id: PriceLevelId::new(),
name: "Wholesale".into(),
code: "WHOLESALE".into(),
description: None,
adjustment_type,
adjustment_value: value,
currency: CurrencyCode::USD,
is_active: true,
created_at: Utc::now(),
updated_at: Utc::now(),
}
}
#[test]
fn percentage_discount() {
let level = make(PriceAdjustmentType::PercentageDiscount, dec!(10));
assert_eq!(level.adjust(dec!(100)), dec!(90));
}
#[test]
fn percentage_markup() {
let level = make(PriceAdjustmentType::PercentageMarkup, dec!(20));
assert_eq!(level.adjust(dec!(100)), dec!(120));
}
#[test]
fn none_returns_base() {
let level = make(PriceAdjustmentType::None, dec!(50));
assert_eq!(level.adjust(dec!(100)), dec!(100));
}
#[test]
fn discount_clamped_at_zero() {
let level = make(PriceAdjustmentType::PercentageDiscount, dec!(150));
assert_eq!(level.adjust(dec!(100)), dec!(0));
}
#[test]
fn resolve_prefers_entry_fixed_price() {
let level = make(PriceAdjustmentType::PercentageDiscount, dec!(10));
let entry = PriceLevelEntry {
price_level_id: level.id,
product_id: ProductId::new(),
price: dec!(42),
created_at: Utc::now(),
updated_at: Utc::now(),
};
assert_eq!(resolve_price(&level, Some(&entry), dec!(100)), dec!(42));
assert_eq!(resolve_price(&level, None, dec!(100)), dec!(90));
}
#[test]
fn adjustment_type_roundtrip() {
for t in [
PriceAdjustmentType::None,
PriceAdjustmentType::PercentageDiscount,
PriceAdjustmentType::PercentageMarkup,
] {
assert_eq!(t.to_string().parse::<PriceAdjustmentType>().unwrap(), t);
}
}
}