use ndarray::{Array1, Array2, Array3, Axis};
use num_complex::Complex64;
use refeff_core::{
RixsTransitionMatrixInput, RixsTransitionMatrixSetup, RixsTransitionPhaseShiftInput,
rixs_transition_matrix_setup, rixs_transition_phase_shifts,
};
use crate::error::Result;
use crate::global_input::GlobalInput;
use super::common::invalid_phase_bin;
use super::types::{
PHASE_BIN_DEFAULT_TRANSITION_COUNT, PhaseBinData, PhaseBinPotential, PhaseBinScalars,
};
use super::validate::validate_phase_bin;
const RIXS_PHASE_MIN: f64 = 1.0e-7;
#[derive(Debug, Clone, PartialEq)]
pub struct PhaseBinRixsHandoff {
pub spin_count: usize,
pub energy_count: usize,
pub main_energy_count: usize,
pub auxiliary_energy_count: usize,
pub ihole: i32,
pub fermi_index: i32,
pub scalars: PhaseBinScalars,
pub energy_grid: Array1<Complex64>,
pub reference_energy: Array2<Complex64>,
pub potentials: Vec<PhaseBinPotential>,
pub transition_moments: Array3<Complex64>,
pub angular_limits: Array2<usize>,
pub max_angular_limit_plus_one: usize,
}
impl PhaseBinRixsHandoff {
#[must_use]
pub fn potential_count(&self) -> usize {
self.potentials.len()
}
}
pub fn phase_bin_rixs_transition_setup_from_handoffs(
global: &GlobalInput,
phase: &PhaseBinRixsHandoff,
) -> Result<RixsTransitionMatrixSetup> {
rixs_transition_matrix_setup(RixsTransitionMatrixInput {
lmax: phase.max_angular_limit_plus_one.saturating_sub(1),
hole: phase.ihole,
polarization: global.control.ipol,
polarization_tensor: rixs_polarization_tensor(global),
multipole: global.control.le2,
trace_orbital: false,
spin: global.control.ispin,
spin_channel_count: phase.spin_count,
spin_vector_angle: global.control.angks,
})
.map_err(|source| invalid_phase_bin("rixs_transition_setup", source.to_string()))
}
pub fn phase_bin_rixs_transition_phase_shifts_from_handoff(
phase: &PhaseBinRixsHandoff,
setup: &RixsTransitionMatrixSetup,
) -> Result<Array2<Complex64>> {
let potential = phase
.potentials
.first()
.ok_or_else(|| invalid_phase_bin("potentials", "phase.bin contains no RIXS potentials"))?;
let signed_l_min = isize::try_from(potential.lmax)
.map_err(|_| invalid_phase_bin("lmax", "phase lmax exceeds isize"))?
.checked_neg()
.ok_or_else(|| invalid_phase_bin("lmax", "phase lmax exceeds signed range"))?;
let phase_shifts = potential.phase_shifts.index_axis(Axis(2), 0);
rixs_transition_phase_shifts(RixsTransitionPhaseShiftInput {
phase_shifts,
signed_l_min,
transition_angular_momenta: &setup.transition_angular_momenta,
})
.map_err(|source| invalid_phase_bin("rixs_transition_phase_shifts", source.to_string()))
}
impl PhaseBinData {
pub fn to_rixs_handoff(&self) -> Result<PhaseBinRixsHandoff> {
phase_bin_rixs_handoff_from_phase_bin(self)
}
}
pub fn phase_bin_rixs_handoff_from_phase_bin(phase: &PhaseBinData) -> Result<PhaseBinRixsHandoff> {
validate_rixs_phase_handoff(phase)?;
let angular_limits = rixs_angular_limits_from_phase_bin(phase)?;
let max_angular_limit_plus_one = match angular_limits.iter().copied().max() {
Some(limit) => limit
.checked_add(1)
.ok_or_else(|| invalid_phase_bin("lmaxp1", "angular limit overflowed"))?,
None => 0,
};
let transition_moments = rixs_transition_moments_from_phase_bin(phase)?;
Ok(PhaseBinRixsHandoff {
spin_count: phase.spin_count,
energy_count: phase.energy_count,
main_energy_count: phase.main_energy_count,
auxiliary_energy_count: phase.auxiliary_energy_count,
ihole: phase.ihole,
fermi_index: phase.fermi_index,
scalars: phase.scalars,
energy_grid: phase.energy_grid.clone(),
reference_energy: phase.reference_energy.clone(),
potentials: phase.potentials.clone(),
transition_moments,
angular_limits,
max_angular_limit_plus_one,
})
}
pub fn rixs_angular_limits_from_phase_bin(phase: &PhaseBinData) -> Result<Array2<usize>> {
validate_phase_bin(phase)?;
let mut angular_limits = Array2::zeros((phase.energy_count, phase.potential_count()));
for (potential_index, potential) in phase.potentials.iter().enumerate() {
for energy in 0..phase.energy_count {
angular_limits[(energy, potential_index)] =
rixs_active_angular_limit(phase, potential, energy)?;
}
}
Ok(angular_limits)
}
pub fn rixs_transition_moments_from_phase_bin(phase: &PhaseBinData) -> Result<Array3<Complex64>> {
validate_rixs_phase_handoff(phase)?;
Ok(Array3::from_shape_fn(
(
phase.energy_count,
PHASE_BIN_DEFAULT_TRANSITION_COUNT,
phase.spin_count,
),
|(energy, transition, spin)| phase.transition_moments[(energy, 0, transition, spin)],
))
}
fn rixs_active_angular_limit(
phase: &PhaseBinData,
potential: &PhaseBinPotential,
energy: usize,
) -> Result<usize> {
for angular in (0..=potential.lmax).rev() {
let signed_l_slot = potential
.lmax
.checked_add(angular)
.ok_or_else(|| invalid_phase_bin("lmax", "positive signed-l slot overflowed"))?;
let first_spin = potential.phase_shifts[(energy, signed_l_slot, 0)];
let last_spin = potential.phase_shifts[(energy, signed_l_slot, phase.spin_count - 1)];
if first_spin.sin().norm() > RIXS_PHASE_MIN || last_spin.sin().norm() > RIXS_PHASE_MIN {
return Ok(angular);
}
}
Ok(0)
}
fn validate_rixs_phase_handoff(phase: &PhaseBinData) -> Result<()> {
validate_phase_bin(phase)?;
if phase.q_count == 0 {
return Err(invalid_phase_bin(
"nq",
"RIXS phase handoff requires at least one q block",
));
}
if phase.transition_count < PHASE_BIN_DEFAULT_TRANSITION_COUNT {
return Err(invalid_phase_bin(
"rkk",
format!(
"RIXS phase handoff requires at least {} transition channels, got {}",
PHASE_BIN_DEFAULT_TRANSITION_COUNT, phase.transition_count
),
));
}
Ok(())
}
fn rixs_polarization_tensor(global: &GlobalInput) -> [[Complex64; 3]; 3] {
let mut tensor = [[Complex64::new(0.0, 0.0); 3]; 3];
for (row_index, row) in global.polarization_tensor.iter().enumerate() {
tensor[row_index] = [
Complex64::new(row[0], row[1]),
Complex64::new(row[2], row[3]),
Complex64::new(row[4], row[5]),
];
}
tensor
}