use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; use litellm_host::call::HostedCompletion; use litellm_host::machine::{HostFailure, Machine, MachineStep}; use litellm_host::protocol::Protocol; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; use tokio::sync::Mutex; use crate::missing_state; use crate::runtime::{poll_async_value, run_async_value, run_sync_value}; type NativeResult = Result< MachineStep< ::Protocol, HostedCompletion<<::Protocol as Protocol>::Response>, >, <::Protocol as Protocol>::Error, >; type MachineResult = Result< MachineStep<::Protocol, ::Complete>, <::Protocol as Protocol>::Error, >; struct MachineState { machine: M, result: Option>, } pub(super) enum NativePoll { Ready(T), Suspend(Py), } pub(super) struct NativeMachine { state: Option>>>, abort: Option, asynchronous: bool, } impl NativeMachine where M::Complete: Into::Response>>, { pub(super) fn new(asynchronous: bool) -> Self { Self { state: None, abort: None, asynchronous, } } pub(super) fn start(&mut self, machine: M) { self.state = Some(Arc::new(Mutex::new(MachineState { machine, result: None, }))); } pub(super) fn resume( &mut self, py: Python<'_>, interruption: Option::Error>>, ) -> PyResult>> { let state = Arc::clone(self.state.as_ref().ok_or_else(missing_state)?); let future = async move { let mut state = state.lock().await; let result = match interruption { Some(failure) => state .machine .interrupt(failure) .await .map(MachineStep::Complete), None => state.machine.resume().await, }; state.result = Some(result); Ok(()) }; if self.asynchronous { let mut future = Box::pin(future); if let Poll::Ready(()) = poll_async_value(py, future.as_mut())? { return Ok(NativePoll::Ready(self.take_result()?)); } let (abort, registration) = AbortHandle::new_pair(); self.abort = Some(abort); Ok(NativePoll::Suspend( run_async_value(py, async move { Abortable::new(future, registration) .await .map_err(|_| PyRuntimeError::new_err("native execution closed"))? })? .unbind(), )) } else { run_sync_value(py, future)?; Ok(NativePoll::Ready(self.take_result()?)) } } pub(super) fn take_result(&self) -> PyResult> { self.state .as_ref() .ok_or_else(missing_state)? .try_lock() .map_err(|_| missing_state())? .result .take() .ok_or_else(missing_state) .map(|result| { result.map(|step| match step { MachineStep::Suspended(op) => MachineStep::Suspended(op), MachineStep::Complete(response) => MachineStep::Complete(response.into()), }) }) } } impl NativeMachine { pub(super) fn close(&mut self) { if let Some(abort) = self.abort.take() { abort.abort(); } self.state = None; } } impl Drop for NativeMachine { fn drop(&mut self) { self.close(); } }