use std::marker::PhantomData;
use vortex_error::VortexResult;
use super::RowVisitor;
use super::check::assert_deferred_visit_contract;
use super::check::assert_owned_visit_contract;
use super::check::assert_sink_visit_contract;
use super::check::validate_owned_visit;
use super::check::validate_sink_visit;
use super::row_visitor::private;
use crate::dtype::DType;
use crate::dtype::Nullability;
use crate::scalar_fn::unstable::row::ElementTuple;
use crate::scalar_fn::unstable::row::FailureEvidence;
use crate::scalar_fn::unstable::row::IndexedElementTuple;
use crate::scalar_fn::unstable::row::OutputElement;
use crate::scalar_fn::unstable::row::OutputSink;
use crate::scalar_fn::unstable::row::RowFn;
use crate::scalar_fn::unstable::row::SinkResult;
pub(crate) struct BatchPlanner<'a, F: RowFn> {
dtypes: &'a [DType],
options: &'a F::Options,
function: PhantomData<F>,
}
impl<'a, F: RowFn> BatchPlanner<'a, F> {
pub(crate) fn new(dtypes: &'a [DType], options: &'a F::Options) -> Self {
Self {
dtypes,
options,
function: PhantomData,
}
}
}
impl<F: RowFn> private::Sealed for BatchPlanner<'_, F> {}
impl<F: RowFn> RowVisitor<F::Options> for BatchPlanner<'_, F> {
type VisitResult = BatchPlan;
fn visit_prepared<Args, Out, Prepared>(
self,
_prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
_apply: impl Fn(&Prepared, Args::Elems<'_>) -> Out,
) -> VortexResult<Self::VisitResult>
where
Args: IndexedElementTuple,
Out: OutputElement,
{
const { assert_owned_visit_contract::<F, Args, Out>() };
Ok(BatchPlan {
output_dtype: validate_owned_visit::<Args, Out>(self.dtypes)?,
policy: RowPolicy::for_owned_output::<Args>(),
})
}
fn visit_prepared_into<Args, Sink, Prepared, ApplyResult>(
self,
_prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
_apply: impl Fn(
&Prepared,
Args::Elems<'_>,
<Sink as OutputSink<F::Options>>::Row<'_>,
) -> ApplyResult,
) -> VortexResult<Self::VisitResult>
where
Args: ElementTuple,
Sink: OutputSink<F::Options>,
ApplyResult: SinkResult<WriteToken = <Sink as OutputSink<F::Options>>::WriteToken>,
{
const { assert_sink_visit_contract::<F, Args, ApplyResult>() };
Ok(BatchPlan {
output_dtype: validate_sink_visit::<Args, Sink, F::Options>(self.options, self.dtypes)?,
policy: RowPolicy::for_sink::<Args, ApplyResult>(),
})
}
fn visit_prepared_deferred<Args, Out, Prepared, Fail>(
self,
_prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
_apply: impl Fn(&Prepared, Args::Elems<'_>) -> (Out, Fail),
_finish_failure: impl FnOnce(Fail) -> VortexResult<()>,
) -> VortexResult<Self::VisitResult>
where
Args: IndexedElementTuple,
Out: OutputElement,
Fail: FailureEvidence,
{
const { assert_deferred_visit_contract::<F, Args, Out, Fail>() };
Ok(BatchPlan {
output_dtype: validate_owned_visit::<Args, Out>(self.dtypes)?,
policy: RowPolicy::for_deferred_output::<Args>(),
})
}
}
pub(crate) struct BatchPlan {
pub(crate) output_dtype: DType,
pub(crate) policy: RowPolicy,
}
impl BatchPlan {
pub(crate) fn result_dtype(&self, args: &[DType]) -> DType {
let nullability = self.output_dtype.nullability()
| Nullability::from(args.iter().any(DType::is_nullable));
self.output_dtype.with_nullability(nullability)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum RowPolicy {
Dense,
ValidOnly,
}
impl RowPolicy {
pub(crate) const fn for_owned_output<Args: ElementTuple>() -> Self {
if Args::DENSE_SAFE && Args::DECODE_INFALLIBLE {
Self::Dense
} else {
Self::ValidOnly
}
}
pub(crate) const fn for_deferred_output<Args: ElementTuple>() -> Self {
let _ = PhantomData::<Args>;
Self::ValidOnly
}
pub(crate) const fn for_sink<Args: ElementTuple, ApplyResult: SinkResult>() -> Self {
if Args::DENSE_SAFE && Args::DECODE_INFALLIBLE && ApplyResult::INFALLIBLE {
Self::Dense
} else {
Self::ValidOnly
}
}
}