use crate::{EvaluationError, EvaluationResult};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use super::rsession::RValue;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RDataFrame {
pub columns: Vec<String>,
pub data: Vec<Vec<RValue>>,
}
impl RDataFrame {
pub fn new() -> Self {
Self {
columns: Vec::new(),
data: Vec::new(),
}
}
pub fn with_columns(columns: Vec<String>) -> Self {
Self {
columns,
data: Vec::new(),
}
}
pub fn add_column(&mut self, name: String, values: Vec<RValue>) -> Result<(), EvaluationError> {
if !self.data.is_empty() && values.len() != self.data.len() {
return Err(EvaluationError::InvalidInput {
message: format!(
"Column length {} doesn't match existing data length {}",
values.len(),
self.data.len()
),
});
}
self.columns.push(name);
if self.data.is_empty() {
for value in values {
self.data.push(vec![value]);
}
} else {
for (i, value) in values.into_iter().enumerate() {
if i < self.data.len() {
self.data[i].push(value);
}
}
}
Ok(())
}
pub fn add_row(&mut self, row: Vec<RValue>) -> Result<(), EvaluationError> {
if row.len() != self.columns.len() {
return Err(EvaluationError::InvalidInput {
message: format!(
"Row length {} doesn't match column count {}",
row.len(),
self.columns.len()
),
});
}
self.data.push(row);
Ok(())
}
pub fn get_column(&self, name: &str) -> Option<Vec<RValue>> {
if let Some(col_index) = self.columns.iter().position(|c| c == name) {
Some(self.data.iter().map(|row| row[col_index].clone()).collect())
} else {
None
}
}
pub fn get_row(&self, index: usize) -> Option<&Vec<RValue>> {
self.data.get(index)
}
pub fn filter<F>(&self, predicate: F) -> Self
where
F: Fn(&Vec<RValue>) -> bool,
{
let filtered_data: Vec<Vec<RValue>> = self
.data
.iter()
.filter(|row| predicate(row))
.cloned()
.collect();
Self {
columns: self.columns.clone(),
data: filtered_data,
}
}
pub fn nrows(&self) -> usize {
self.data.len()
}
pub fn ncols(&self) -> usize {
self.columns.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
pub fn column_names(&self) -> &Vec<String> {
&self.columns
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RSurvivalModel {
pub concordance: f64,
pub log_likelihood: f64,
pub aic: f64,
pub coefficients: Vec<String>,
pub estimates: Vec<f64>,
pub hazard_ratios: Vec<f64>,
pub p_values: Vec<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RPcaResult {
pub variance_explained: Vec<f64>,
pub cumulative_variance: Vec<f64>,
pub component_names: Vec<String>,
pub n_components: i32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RAnovaResult {
pub sources: Vec<String>,
pub df: Vec<i32>,
pub sum_squares: Vec<f64>,
pub mean_squares: Vec<f64>,
pub f_statistics: Vec<f64>,
pub p_values: Vec<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RLogisticModel {
pub coefficients: Vec<String>,
pub estimates: Vec<f64>,
pub std_errors: Vec<f64>,
pub z_values: Vec<f64>,
pub p_values: Vec<f64>,
pub aic: f64,
pub deviance: f64,
pub null_deviance: f64,
}