use std::{collections::HashMap, fs, path::Path, time::Duration};
use anyhow::{Context, Result, ensure};
use serde::{Deserialize, Serialize};
use crate::cfg::enums::{Digest, SessionType, YesNo};
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Config {
pub login: LoginConfig,
pub runtime: RuntimeConfig,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct LoginConfig {
pub identity: Identity,
pub auth: AuthConfig,
pub integrity: Integrity,
pub flow: Flow,
pub write_flow: WriteFlow,
pub ordering: Ordering,
pub recovery: Recovery,
pub timers: Timers,
pub limits: Limits,
pub extensions: Extensions,
pub transport: TransportHints,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Identity {
#[serde(rename = "SessionType")]
pub session_type: SessionType,
#[serde(rename = "InitiatorName")]
pub initiator_name: String,
#[serde(default, rename = "InitiatorAlias")]
pub initiator_alias: String,
#[serde(default, rename = "TargetName")]
pub target_name: String,
#[serde(rename = "IsX86")]
pub is_x86: YesNo,
}
#[derive(Deserialize, Serialize, Debug, Clone, Default)]
pub struct TransportHints {
#[serde(default, rename = "TargetAddress")]
pub target_address: String,
#[serde(default, rename = "TargetPortalGroupTag")]
pub portal_group_tag: u16,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
#[serde(tag = "AuthMethod")]
pub enum AuthConfig {
#[serde(rename = "None")]
None,
#[serde(rename = "CHAP")]
Chap(ChapConfig),
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct ChapConfig {
pub username: String,
pub secret: String,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Integrity {
#[serde(rename = "HeaderDigest")]
pub header_digest: Digest,
#[serde(rename = "DataDigest")]
pub data_digest: Digest,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Flow {
#[serde(rename = "MaxRecvDataSegmentLength")]
pub max_recv_data_segment_length: u32,
#[serde(rename = "MaxBurstLength")]
pub max_burst_length: u32,
#[serde(rename = "FirstBurstLength")]
pub first_burst_length: u32,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct WriteFlow {
#[serde(rename = "InitialR2T")]
pub initial_r2t: YesNo,
#[serde(rename = "ImmediateData")]
pub immediate_data: YesNo,
#[serde(rename = "MaxOutstandingR2T")]
pub max_outstanding_r2t: u8,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Ordering {
#[serde(rename = "DataPDUInOrder")]
pub data_pdu_in_order: YesNo,
#[serde(rename = "DataSequenceInOrder")]
pub data_sequence_in_order: YesNo,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Recovery {
#[serde(rename = "ErrorRecoveryLevel")]
pub error_recovery_level: u8,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Timers {
#[serde(rename = "DefaultTime2Wait", with = "serde_secs")]
pub default_time2wait: Duration,
#[serde(rename = "DefaultTime2Retain", with = "serde_secs")]
pub default_time2retain: Duration,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Limits {
#[serde(rename = "MaxConnections")]
pub max_connections: u16,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Extensions {
#[serde(rename = "TaskReporting", skip_serializing_if = "Option::is_none")]
pub task_reporting: Option<TaskReporting>,
#[serde(rename = "iSCSIProtocolLevel", skip_serializing_if = "Option::is_none")]
pub iscsi_protocol_level: Option<u8>,
#[serde(flatten)]
pub custom: HashMap<String, String>,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
#[serde(rename_all = "PascalCase")]
pub enum TaskReporting {
RFC3720,
ResponseFence,
FastAbort,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct RuntimeConfig {
#[serde(rename = "MaxSessions")]
pub max_sessions: u32,
#[serde(rename = "TimeoutConnection", with = "serde_secs")]
pub timeout_connection: Duration,
}
impl Config {
pub fn load_from_file<P: AsRef<Path>>(path: P) -> Result<Self> {
let s = fs::read_to_string(path)?;
let mut cfg: Config =
serde_yaml::from_str(&s).context("failed to parse config YAML")?;
cfg.validate_and_normalize()?;
Ok(cfg)
}
pub fn validate_and_normalize(&mut self) -> Result<()> {
if self.login.identity.session_type.is_discovery() {
if self.login.limits.max_connections != 1 {
self.login.limits.max_connections = 1;
}
if self.login.recovery.error_recovery_level != 0 {
self.login.recovery.error_recovery_level = 0;
}
}
if let Some(lv) = self.login.extensions.iscsi_protocol_level {
ensure!(lv >= 1, "iSCSIProtocolLevel must be >= 1");
}
ensure!(
!self.login.identity.initiator_name.is_empty(),
"InitiatorName must not be empty"
);
if self.login.identity.session_type.is_normal() {
ensure!(
!self.login.identity.target_name.is_empty(),
"TargetName is required for Normal session"
);
}
ensure!(
self.login.limits.max_connections >= 1,
"MaxConnections must be >= 1"
);
ensure!(self.runtime.max_sessions >= 1, "MaxSessions must be >= 1");
Ok(())
}
}
impl SessionType {
pub fn is_discovery(&self) -> bool {
matches!(self, SessionType::Discovery)
}
pub fn is_normal(&self) -> bool {
matches!(self, SessionType::Normal)
}
}
fn build_kv_sorted<'a, I>(items: I) -> Vec<u8>
where I: IntoIterator<Item = (&'a str, Option<String>)> {
let mut vec: Vec<(String, String)> = items
.into_iter()
.filter_map(|(k, v)| v.map(|vv| (k.to_string(), vv)))
.collect();
vec.sort_unstable_by(|a, b| a.0.cmp(&b.0));
let mut out =
Vec::with_capacity(vec.iter().map(|(k, v)| k.len() + 1 + v.len() + 1).sum());
for (k, v) in vec {
out.extend_from_slice(k.as_bytes());
out.push(b'=');
out.extend_from_slice(v.as_bytes());
out.push(0);
}
out
}
pub fn login_keys_security(cfg: &Config) -> Vec<u8> {
let id = &cfg.login.identity;
build_kv_sorted([
("SessionType", Some(id.session_type.to_string())),
("InitiatorName", Some(id.initiator_name.clone())),
(
"InitiatorAlias",
(!id.initiator_alias.is_empty()).then(|| id.initiator_alias.clone()),
),
(
"TargetName",
(id.session_type.is_normal() && !id.target_name.is_empty())
.then(|| id.target_name.clone()),
),
(
"AuthMethod",
Some(match cfg.login.auth {
AuthConfig::None => "None".to_string(),
AuthConfig::Chap(_) => "CHAP,None".to_string(),
}),
),
])
}
pub fn login_keys_chap_response(user: &str, chap_r_upper_hex_with_0x: &str) -> Vec<u8> {
build_kv_sorted([
("CHAP_N", Some(user.to_string())),
("CHAP_R", Some(chap_r_upper_hex_with_0x.to_string())),
])
}
pub fn login_keys_operational(cfg: &Config) -> Vec<u8> {
let n = &cfg.login;
let mut items: Vec<(&str, Option<String>)> = vec![
("HeaderDigest", Some(n.integrity.header_digest.to_string())),
("DataDigest", Some(n.integrity.data_digest.to_string())),
(
"DataPDUInOrder",
Some(n.ordering.data_pdu_in_order.to_string()),
),
(
"DataSequenceInOrder",
Some(n.ordering.data_sequence_in_order.to_string()),
),
(
"ErrorRecoveryLevel",
Some(n.recovery.error_recovery_level.to_string()),
),
(
"FirstBurstLength",
Some(n.flow.first_burst_length.to_string()),
),
("MaxBurstLength", Some(n.flow.max_burst_length.to_string())),
(
"MaxRecvDataSegmentLength",
Some(n.flow.max_recv_data_segment_length.to_string()),
),
(
"ImmediateData",
Some(n.write_flow.immediate_data.to_string()),
),
("InitialR2T", Some(n.write_flow.initial_r2t.to_string())),
(
"MaxOutstandingR2T",
Some(n.write_flow.max_outstanding_r2t.to_string()),
),
(
"DefaultTime2Retain",
Some(n.timers.default_time2retain.as_secs().to_string()),
),
(
"DefaultTime2Wait",
Some(n.timers.default_time2wait.as_secs().to_string()),
),
("MaxConnections", Some(n.limits.max_connections.to_string())),
];
if let Some(tr) = &n.extensions.task_reporting {
let v = match tr {
TaskReporting::RFC3720 => "RFC3720",
TaskReporting::ResponseFence => "ResponseFence",
TaskReporting::FastAbort => "FastAbort",
}
.to_string();
items.push(("TaskReporting", Some(v)));
}
if let Some(pl) = n.extensions.iscsi_protocol_level {
items.push(("iSCSIProtocolLevel", Some(pl.to_string())));
}
for (k, v) in &n.extensions.custom {
items.push((k.as_str(), Some(v.clone())));
}
build_kv_sorted(items)
}
mod serde_secs {
use std::time::Duration;
use serde::{Deserialize, Deserializer, Serializer};
pub fn serialize<S: Serializer>(d: &Duration, s: S) -> Result<S::Ok, S::Error> {
s.serialize_u64(d.as_secs())
}
pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result<Duration, D::Error> {
let secs = u64::deserialize(d)?;
Ok(Duration::from_secs(secs))
}
}