use std::{
borrow::Cow,
mem,
time::{Duration, Instant},
};
use monty::{Dump, MontyRepl, ReplProgress, ReplStartError, Session, SessionRef, dump, source_within_nesting_bound};
use monty_type_checking::{SourceFile, TypeChecker};
use monty_types::{
AssertMessageAnnotations, CompileOptions, ExcType, ExtFunctionResult, MontyException, MontyObject, OsFunctionCall,
OsPolicy, PrintStream, PrintWriter, PrintWriterCallback, ResourceLimits, ResourceTracker, SOURCE_SCAN_THRESHOLD,
TypeCheckState, TypeCheckingConfig, allocate_into_baseline,
};
use super::{
BudgetVec, DEFAULT_PRINT_FLUSH_INTERVAL, FrameError, FrameReader, MAX_FRAME_LEN, ProtoConvertError,
WireFunctionCall, check_protocol_version, exceeds_max_frame_len, ext_result_from_proto, future_results_from_proto,
named_values_from_proto, os_call_from_proto, os_call_to_proto, pb, write_frame,
};
use crate::{convert::limits::micros_field, wire::uuid_to_pb};
pub trait EventSink {
fn send(&mut self, event: &pb::ChildEvent) -> Result<(), FrameError>;
}
#[derive(Default)]
pub struct VecEventSink {
frames: Vec<u8>,
}
impl VecEventSink {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn take(&mut self) -> Vec<u8> {
mem::take(&mut self.frames)
}
}
impl EventSink for VecEventSink {
fn send(&mut self, event: &pb::ChildEvent) -> Result<(), FrameError> {
write_frame(&mut self.frames, event)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HandleOutcome {
Continue,
Shutdown,
Fatal,
}
pub fn dispatch_frame(child: &mut Child, request_frame: &[u8]) -> (Vec<u8>, HandleOutcome) {
let mut sink = VecEventSink::new();
let outcome = dispatch_into(child, request_frame, &mut sink);
(sink.take(), outcome)
}
fn dispatch_into(child: &mut Child, request_frame: &[u8], sink: &mut VecEventSink) -> HandleOutcome {
let mut reader = FrameReader::new(request_frame);
match reader.read::<pb::ParentRequest>() {
Ok(Some(request)) => match child.handle(request, sink) {
Ok(outcome) => outcome,
Err(FrameError::FrameTooLarge { len, max }) => {
let _ = sink
.send(&child.fatal_event(&format!("response frame of {len} bytes exceeds maximum of {max} bytes")));
HandleOutcome::Shutdown
}
Err(_) => HandleOutcome::Shutdown,
},
Ok(None) => HandleOutcome::Continue,
Err(FrameError::Decode(err)) => {
let _ = sink.send(&protocol_violation(&format!("malformed request: {err}")));
HandleOutcome::Continue
}
Err(err) => {
let _ = sink.send(&child.fatal_event(&format!("malformed request frame: {err}")));
HandleOutcome::Shutdown
}
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct SessionBudget {
pub max_memory: Option<usize>,
pub type_check: bool,
pub max_suspensions: Option<usize>,
pub max_total_sleep: Option<Duration>,
}
enum SessionState {
Configured(Option<Box<pb::Configure>>),
Ready(Box<MontyRepl>),
Suspended(Box<ReplProgress>),
}
pub struct Child {
state: SessionState,
script_name: String,
type_checker: Option<TypeChecker>,
type_check: Option<TypeCheckState>,
print_flush_interval: Duration,
os_policy: OsPolicy,
}
impl Default for Child {
fn default() -> Self {
Self {
state: SessionState::Configured(None),
script_name: String::new(),
type_checker: None,
type_check: None,
print_flush_interval: DEFAULT_PRINT_FLUSH_INTERVAL,
os_policy: OsPolicy::default(),
}
}
}
impl Child {
pub fn handle(
&mut self,
request: pb::ParentRequest,
sink: &mut dyn EventSink,
) -> Result<HandleOutcome, FrameError> {
let Some(kind) = request.kind else {
sink.send(&protocol_violation("request has no kind"))?;
return Ok(HandleOutcome::Continue);
};
let mut event = match kind {
pb::parent_request::Kind::Configure(configure) => {
if let Err(refusal) = check_protocol_version(configure.protocol_version) {
sink.send(&self.fatal_event(&refusal))?;
return Ok(HandleOutcome::Fatal);
}
self.handle_configure(configure)
}
pb::parent_request::Kind::Feed(feed) => self.handle_repl_feed(feed, sink),
pb::parent_request::Kind::InstallDependencies(_) => error_event(
ExcType::RuntimeError,
"dependency installation is only supported by the CPython worker",
),
pb::parent_request::Kind::ResumeCall(resume) => self.handle_resume_call(resume, sink),
pb::parent_request::Kind::ResumeNameLookup(resume) => self.handle_resume_name_lookup(resume, sink),
pb::parent_request::Kind::ResumeFutures(resume) => self.handle_resume_futures(resume, sink),
pb::parent_request::Kind::AbortFeed(abort) => self.handle_abort_feed(abort, sink),
pb::parent_request::Kind::Dump(_) => self.handle_dump(),
pb::parent_request::Kind::Load(load) => self.handle_load(&load),
pb::parent_request::Kind::Reset(_) => match self.reset() {
Ok(()) => ok_event(),
Err(err) => {
sink.send(&self.fatal_event(&format!("type-check cleanup failed: {err}")))?;
return Ok(HandleOutcome::Fatal);
}
},
pb::parent_request::Kind::Shutdown(_) => {
sink.send(&ok_event())?;
return Ok(HandleOutcome::Shutdown);
}
};
self.stamp_session_budget(&mut event);
let sent = sink.send(&event);
if let Err(err) = self.reclaim_suspension_payload(&mut event) {
sink.send(&self.fatal_event(&format!("suspension payload could not be restored: {err}")))?;
return Ok(HandleOutcome::Fatal);
}
if let Err(err) = sent {
self.recover_send_error(&event, err, sink)?;
}
Ok(HandleOutcome::Continue)
}
fn reclaim_suspension_payload(&mut self, event: &mut pb::ChildEvent) -> Result<(), ProtoConvertError> {
let SessionState::Suspended(progress) = &mut self.state else {
return Ok(());
};
match (progress.as_mut(), &mut event.kind) {
(ReplProgress::FunctionCall(call), Some(pb::child_event::Kind::FunctionCall(announced))) => {
call.args = mem::take(announced).into_call_args()?;
}
(ReplProgress::OsCall(call), Some(pb::child_event::Kind::OsCall(announced)))
if announced.call.is_some() =>
{
call.function_call = os_call_from_proto(mem::take(announced))?.1;
}
_ => {}
}
Ok(())
}
#[must_use]
pub fn session_budget(&self) -> SessionBudget {
match &self.state {
SessionState::Configured(Some(config)) => SessionBudget {
max_memory: config
.limits
.as_ref()
.and_then(|limits| limits.max_memory_bytes)
.map(|v| usize::try_from(v).unwrap_or(usize::MAX)),
type_check: config.type_check,
max_suspensions: Some(ResourceLimits::from(config.limits.unwrap_or_default()).max_suspensions),
max_total_sleep: config
.limits
.as_ref()
.and_then(|limits| limits.max_total_sleep_micros)
.map(Duration::from_micros),
},
SessionState::Configured(None) => SessionBudget::default(),
SessionState::Ready(repl) => self.tracker_budget(repl.tracker()),
SessionState::Suspended(progress) => self.tracker_budget(progress.tracker()),
}
}
fn tracker_budget(&self, tracker: &ResourceTracker) -> SessionBudget {
SessionBudget {
max_memory: tracker.max_memory(),
type_check: self.type_check.is_some(),
max_suspensions: Some(tracker.max_suspensions()),
max_total_sleep: tracker.max_total_sleep(),
}
}
#[must_use]
pub fn fatal_event(&self, message: &str) -> pb::ChildEvent {
let mut event = fatal_error_event(message);
self.stamp_session_budget(&mut event);
event
}
fn recover_send_error(
&mut self,
failed: &pb::ChildEvent,
err: FrameError,
sink: &mut dyn EventSink,
) -> Result<(), FrameError> {
let announces_suspension = matches!(
failed.kind,
Some(
pb::child_event::Kind::FunctionCall(_)
| pb::child_event::Kind::OsCall(_)
| pb::child_event::Kind::NameLookup(_)
| pb::child_event::Kind::ResolveFutures(_)
)
);
match err {
FrameError::FrameTooLarge { len, max } if !announces_suspension => {
let mut event = error_event(
ExcType::RuntimeError,
&format!("result frame of {len} bytes exceeds the maximum of {max} bytes"),
);
self.stamp_session_budget(&mut event);
sink.send(&event)
}
other => Err(other),
}
}
fn stamp_session_budget(&self, event: &mut pb::ChildEvent) {
let tracker = match &self.state {
SessionState::Ready(repl) => repl.tracker(),
SessionState::Suspended(progress) => progress.tracker(),
SessionState::Configured(_) => return,
};
stamp_budget(event, tracker);
}
fn handle_configure(&mut self, configure: pb::Configure) -> pb::ChildEvent {
if matches!(self.state, SessionState::Configured(None)) {
self.print_flush_interval = configure
.print_flush_interval_ms
.map_or(DEFAULT_PRINT_FLUSH_INTERVAL, |ms| Duration::from_millis(u64::from(ms)));
self.os_policy = match configure.os_policy.clone().map(OsPolicy::try_from) {
None => OsPolicy::default(),
Some(Ok(os_policy)) => os_policy,
Some(Err(err)) => return protocol_violation(&format!("invalid os_policy: {err}")),
};
if let Some(stubs) = &configure.type_check_stubs
&& !source_within_nesting_bound(stubs, SOURCE_SCAN_THRESHOLD)
{
return protocol_violation("invalid type_check_stubs: Source is too deeply nested");
}
self.state = SessionState::Configured(Some(Box::new(configure)));
ok_event()
} else {
protocol_violation("Configure while a session already exists")
}
}
fn ensure_repl(&mut self) -> Result<(), Box<pb::ChildEvent>> {
let config = match &mut self.state {
SessionState::Configured(config) => config.take(),
SessionState::Ready(_) | SessionState::Suspended(_) => return Ok(()),
};
let Some(config) = config else {
return Err(Box::new(protocol_violation("session has not been configured")));
};
let type_check_config = TypeCheckingConfig::from(config.as_ref());
let pb::Configure {
script_name,
limits,
type_check,
type_check_stubs,
assert_message_annotations,
type_check_format: _,
type_check_color: _,
protocol_version: _,
monty_version: _,
print_flush_interval_ms: _,
os_policy: _,
persistence: _,
} = *config;
let limits = limits.unwrap_or_default().into();
self.script_name = script_name;
self.type_check = type_check.then(|| TypeCheckState {
committed_stubs: type_check_stubs.unwrap_or_default(),
pending_snippet: None,
config: type_check_config,
});
let options = CompileOptions {
assert_message_annotations: assert_message_annotations.map_or_else(
AssertMessageAnnotations::default,
AssertMessageAnnotations::from_max_bytes,
),
source_scan_threshold: SOURCE_SCAN_THRESHOLD,
};
let repl = MontyRepl::new(&self.script_name, ResourceTracker::new(limits), options)
.with_os_policy(self.os_policy.clone());
self.state = SessionState::Ready(Box::new(repl));
Ok(())
}
fn handle_repl_feed(&mut self, feed: pb::Feed, sink: &mut dyn EventSink) -> pb::ChildEvent {
if let Err(event) = self.ensure_repl() {
return *event;
}
let SessionState::Ready(repl) = &self.state else {
return protocol_violation("Feed without a session ready for input");
};
let code = match repl.check_source(&feed.code) {
Ok(code) => code,
Err(error) => {
return event(pb::child_event::Kind::Error(pb::Error {
exception: Some((&error).into()),
}));
}
};
if !feed.skip_type_check
&& let Some(event) = self.type_check_feed(&feed.code)
{
return event;
}
let inputs = match named_values_from_proto(feed.inputs, feed.values) {
Ok(inputs) => inputs,
Err(err) => return protocol_violation(&format!("invalid inputs: {err}")),
};
let SessionState::Ready(mut repl) = mem::replace(&mut self.state, SessionState::Configured(None)) else {
unreachable!("checked Ready above");
};
if !feed.cwd.is_empty() {
repl.set_cwd(&feed.cwd);
}
if !feed.skip_type_check
&& let Some(state) = &mut self.type_check
{
state.pending_snippet = Some(feed.code.clone());
}
let mut print = ProtoPrint::new(sink, self.print_flush_interval);
let result = repl.feed_start_checked(code, inputs, PrintWriter::Callback(&mut print));
let event = self.drive(result);
print.drain();
event
}
fn handle_resume_call(&mut self, resume: pb::ResumeCall, sink: &mut dyn EventSink) -> pb::ChildEvent {
let expected_call_id = match &self.state {
SessionState::Suspended(progress) => match progress.as_ref() {
ReplProgress::FunctionCall(call) => Some(call.call_id),
ReplProgress::OsCall(call) => Some(call.call_id),
_ => None,
},
_ => None,
};
let Some(call_id) = expected_call_id else {
return protocol_violation("ResumeCall without a suspended function/OS call");
};
if resume.call_id != call_id {
return protocol_violation(&format!(
"ResumeCall call_id {} does not match {call_id}",
resume.call_id
));
}
let Some(wire_result) = resume.result else {
return protocol_violation("ResumeCall has no result");
};
let result: ExtFunctionResult =
if matches!(wire_result.kind, Some(pb::ext_function_result::Kind::NotHandled(_))) {
let SessionState::Suspended(progress) = &self.state else {
unreachable!("checked above");
};
let ReplProgress::OsCall(call) = progress.as_ref() else {
return protocol_violation("NotHandled is only valid answering a suspended OS call");
};
ExtFunctionResult::Error(call.function_call.on_no_handler())
} else {
match ext_result_from_proto(wire_result, resume.values) {
Ok(result) => result,
Err(err) => return protocol_violation(&format!("invalid result: {err}")),
}
};
let SessionState::Suspended(progress) = mem::replace(&mut self.state, SessionState::Configured(None)) else {
unreachable!("checked above");
};
let mut print = ProtoPrint::new(sink, self.print_flush_interval);
let outcome = match *progress {
ReplProgress::FunctionCall(call) => call.resume(result, PrintWriter::Callback(&mut print)),
ReplProgress::OsCall(call) => call.resume(result, PrintWriter::Callback(&mut print)),
_ => unreachable!("checked above"),
};
let event = self.drive(outcome);
print.drain();
event
}
fn handle_resume_name_lookup(&mut self, resume: pb::ResumeNameLookup, sink: &mut dyn EventSink) -> pb::ChildEvent {
let SessionState::Suspended(progress) = &self.state else {
return protocol_violation("ResumeNameLookup without a suspended name lookup");
};
if !matches!(progress.as_ref(), ReplProgress::NameLookup(_)) {
return protocol_violation("ResumeNameLookup without a suspended name lookup");
}
let result = match resume.try_into() {
Ok(result) => result,
Err(err) => return protocol_violation(&format!("invalid result: {err}")),
};
let SessionState::Suspended(progress) = mem::replace(&mut self.state, SessionState::Configured(None)) else {
unreachable!("checked above");
};
let ReplProgress::NameLookup(lookup) = *progress else {
unreachable!("checked above");
};
let mut print = ProtoPrint::new(sink, self.print_flush_interval);
let outcome = lookup.resume(result, PrintWriter::Callback(&mut print));
let event = self.drive(outcome);
print.drain();
event
}
fn handle_abort_feed(&mut self, abort: pb::AbortFeed, sink: &mut dyn EventSink) -> pb::ChildEvent {
let suspended = matches!(&self.state, SessionState::Suspended(progress)
if !matches!(progress.as_ref(), ReplProgress::Complete { .. }));
if !suspended {
return protocol_violation("AbortFeed without a suspended feed");
}
let Some(exception) = abort.exception else {
return protocol_violation("AbortFeed has no exception");
};
let exc = match MontyException::try_from(exception) {
Ok(exc) => exc,
Err(err) => return protocol_violation(&format!("invalid exception: {err}")),
};
let SessionState::Suspended(progress) = mem::replace(&mut self.state, SessionState::Configured(None)) else {
unreachable!("checked above");
};
let mut print = ProtoPrint::new(sink, self.print_flush_interval);
let outcome = match *progress {
ReplProgress::FunctionCall(call) => call.abort(exc, PrintWriter::Callback(&mut print)),
ReplProgress::OsCall(call) => call.abort(exc, PrintWriter::Callback(&mut print)),
ReplProgress::NameLookup(lookup) => lookup.abort(exc, PrintWriter::Callback(&mut print)),
ReplProgress::ResolveFutures(state) => state.abort(exc, PrintWriter::Callback(&mut print)),
ReplProgress::Complete { .. } => unreachable!("checked above"),
};
let event = self.drive(outcome);
print.drain();
event
}
fn handle_resume_futures(&mut self, resume: pb::ResumeFutures, sink: &mut dyn EventSink) -> pb::ChildEvent {
let results = match future_results_from_proto(resume.results, resume.values) {
Ok(results) => results,
Err(err) => return protocol_violation(&format!("invalid results: {err}")),
};
let SessionState::Suspended(progress) = &self.state else {
return protocol_violation("ResumeFutures without suspended futures");
};
let reply = match progress.as_ref() {
ReplProgress::FunctionCall(call) if call.allow_eager_await => match eager_result(results, call.call_id) {
Ok(result) => FuturesReply::Eager(result),
Err(message) => return protocol_violation(message),
},
ReplProgress::OsCall(call) if call.allow_eager_await => match eager_result(results, call.call_id) {
Ok(result) => FuturesReply::Eager(result),
Err(message) => return protocol_violation(message),
},
ReplProgress::ResolveFutures(_) => FuturesReply::Batch(results),
_ => return protocol_violation("ResumeFutures without suspended futures"),
};
let SessionState::Suspended(progress) = mem::replace(&mut self.state, SessionState::Configured(None)) else {
unreachable!("checked above");
};
let mut print = ProtoPrint::new(sink, self.print_flush_interval);
let outcome = match (*progress, reply) {
(ReplProgress::FunctionCall(call), FuturesReply::Eager(result)) => {
call.resume_eager(result, PrintWriter::Callback(&mut print))
}
(ReplProgress::OsCall(call), FuturesReply::Eager(result)) => {
call.resume_eager(result, PrintWriter::Callback(&mut print))
}
(ReplProgress::ResolveFutures(state), FuturesReply::Batch(results)) => {
state.resume(results, PrintWriter::Callback(&mut print))
}
_ => unreachable!("reply shaped by the suspension above"),
};
let event = self.drive(outcome);
print.drain();
event
}
fn handle_dump(&mut self) -> pb::ChildEvent {
if let Err(event) = self.ensure_repl() {
return *event;
}
let session = match &self.state {
SessionState::Ready(repl) => SessionRef::Idle(repl),
SessionState::Suspended(progress) => SessionRef::Suspended(progress),
SessionState::Configured(_) => unreachable!("ensure_repl materialized the repl or errored"),
};
match dump(&self.script_name, self.type_check.as_ref(), session) {
Ok(state) => event(pb::child_event::Kind::DumpResult(pb::DumpResult {
state: state.into(),
})),
Err(err) => protocol_violation(&format!("dump failed: {err}")),
}
}
fn handle_load(&mut self, load: &pb::Load) -> pb::ChildEvent {
if !matches!(self.state, SessionState::Configured(_)) {
return protocol_violation("Load requires a session that has not started (a feed has already run)");
}
let restored = match Dump::load(&load.state) {
Ok(restored) => restored,
Err(err) => return protocol_violation(&format!("failed to load session: {err}")),
};
let Dump {
script_name,
type_check,
state,
} = restored;
let mut event = match state {
Session::Idle(repl) => {
self.state = SessionState::Ready(repl);
ok_event()
}
Session::Running(_) => protocol_violation("dump holds a one-shot run, not a repl session"),
Session::Suspended(progress) => match *progress {
ReplProgress::Complete { repl, value } => {
self.state = SessionState::Ready(Box::new(repl));
complete_event(value)
}
mut progress => {
let mut event = suspension_event(&mut progress);
stamp_budget(&mut event, progress.tracker());
if let Some(message) = oversize_suspension_error_message(&event) {
protocol_violation(&message)
} else {
self.state = SessionState::Suspended(Box::new(progress));
event
}
}
},
};
if matches!(self.state, SessionState::Ready(_) | SessionState::Suspended(_)) {
self.script_name = script_name;
self.type_check = type_check;
event.restored_script_name = Some(self.script_name.clone());
}
event
}
fn drive(&mut self, result: Result<ReplProgress, Box<ReplStartError>>) -> pb::ChildEvent {
match result {
Ok(ReplProgress::Complete { repl, value }) => {
self.state = SessionState::Ready(Box::new(repl));
if let Some(state) = &mut self.type_check
&& let Some(snippet) = state.pending_snippet.take()
{
state.committed_stubs.push('\n');
state.committed_stubs.push_str(&snippet);
}
complete_event(value)
}
Ok(ReplProgress::OsCall(mut call)) => {
let mut event = suspension_event_os_call(&mut call);
let progress = ReplProgress::OsCall(call);
stamp_budget(&mut event, progress.tracker());
if let Some(message) = oversize_suspension_error_message(&event) {
self.abort_feed_with_runtime_error(progress.into_repl(), &message)
} else {
self.state = SessionState::Suspended(Box::new(progress));
event
}
}
Ok(ReplProgress::FunctionCall(mut call)) => {
let mut event = suspension_event_function_call(&mut call);
let progress = ReplProgress::FunctionCall(call);
stamp_budget(&mut event, progress.tracker());
if let Some(message) = oversize_suspension_error_message(&event) {
self.abort_feed_with_runtime_error(progress.into_repl(), &message)
} else {
self.state = SessionState::Suspended(Box::new(progress));
event
}
}
Ok(mut progress) => {
let event = suspension_event(&mut progress);
self.state = SessionState::Suspended(Box::new(progress));
event
}
Err(err) => {
self.state = SessionState::Ready(Box::new(err.repl));
if let Some(state) = &mut self.type_check {
state.pending_snippet = None;
}
event(pb::child_event::Kind::Error(pb::Error {
exception: Some((&err.error).into()),
}))
}
}
}
fn abort_feed_with_runtime_error(&mut self, repl: MontyRepl, message: &str) -> pb::ChildEvent {
self.state = SessionState::Ready(Box::new(repl));
if let Some(state) = &mut self.type_check {
state.pending_snippet = None;
}
error_event(ExcType::RuntimeError, message)
}
fn type_check_feed(&mut self, code: &str) -> Option<pb::ChildEvent> {
let state = self.type_check.as_ref()?;
let stubs =
(!state.committed_stubs.is_empty()).then(|| SourceFile::new(&state.committed_stubs, "repl_type_stubs.pyi"));
let type_checker = self
.type_checker
.get_or_insert_with(|| allocate_into_baseline(TypeChecker::default));
match type_checker.run(&SourceFile::new(code, &self.script_name), stubs.as_ref(), state.config) {
Ok(None) => None,
Ok(Some(diagnostics)) => Some(event(pb::child_event::Kind::TypingError(pb::TypingError {
diagnostics: diagnostics.to_string(),
}))),
Err(err) => Some(protocol_violation(&format!("type checker failed: {err}"))),
}
}
fn reset(&mut self) -> Result<(), String> {
self.state = SessionState::Configured(None);
self.type_check = None;
self.script_name = String::new();
self.print_flush_interval = DEFAULT_PRINT_FLUSH_INTERVAL;
self.os_policy = OsPolicy::default();
self.type_checker.as_mut().map_or(Ok(()), TypeChecker::reset)
}
}
fn event(kind: pb::child_event::Kind) -> pb::ChildEvent {
pb::ChildEvent {
kind: Some(kind),
..Default::default()
}
}
#[must_use]
pub fn protocol_violation(message: &str) -> pb::ChildEvent {
event(pb::child_event::Kind::Error(pb::Error {
exception: Some(pb::RaisedException {
exc_type: ExcType::RuntimeError.to_string(),
message: Some(format!("protocol violation: {message}")),
traceback: BudgetVec::new(),
data: None,
}),
}))
}
#[must_use]
pub fn fatal_error_event(message: &str) -> pb::ChildEvent {
event(pb::child_event::Kind::FatalError(pb::FatalError {
message: message.to_owned(),
}))
}
fn ok_event() -> pb::ChildEvent {
event(pb::child_event::Kind::Ok(pb::Ok {}))
}
fn error_event(exc_type: ExcType, message: &str) -> pb::ChildEvent {
event(pb::child_event::Kind::Error(pb::Error {
exception: Some(pb::RaisedException {
exc_type: exc_type.to_string(),
message: Some(message.to_owned()),
traceback: BudgetVec::new(),
data: None,
}),
}))
}
fn stamp_budget(event: &mut pb::ChildEvent, tracker: &ResourceTracker) {
event.total_execution_micros = u64::try_from(tracker.elapsed().as_micros()).unwrap_or(u64::MAX);
event.feed_execution_micros = u64::try_from(tracker.feed_elapsed().as_micros()).unwrap_or(u64::MAX);
event.max_feed_duration_micros = micros_field(tracker.max_feed_duration());
event.max_turn_duration_micros = micros_field(tracker.max_turn_duration());
event.max_total_sleep_micros = micros_field(tracker.max_total_sleep());
event.max_suspensions = Some(tracker.max_suspensions() as u64);
}
fn oversize_suspension_error_message(event: &pb::ChildEvent) -> Option<String> {
exceeds_max_frame_len(event)
.map(|len| format!("argument frame of {len} bytes exceeds the maximum of {MAX_FRAME_LEN} bytes"))
}
fn suspension_event_function_call(call: &mut monty::ReplFunctionCall) -> pb::ChildEvent {
event(pb::child_event::Kind::FunctionCall(WireFunctionCall::new(
call.function_name.clone(),
mem::take(&mut call.args),
call.call_id,
call.object_id,
call.allow_eager_await,
call.position.clone(),
)))
}
fn suspension_event_os_call(call: &mut monty::ReplOsCall) -> pb::ChildEvent {
let function_call = mem::replace(&mut call.function_call, OsFunctionCall::GetEnviron);
event(pb::child_event::Kind::OsCall(os_call_to_proto(
call.call_id,
function_call,
call.allow_eager_await,
&call.position,
)))
}
fn complete_event(value: MontyObject) -> pb::ChildEvent {
event(pb::child_event::Kind::Complete(value.into()))
}
fn suspension_event(progress: &mut ReplProgress) -> pb::ChildEvent {
match progress {
ReplProgress::FunctionCall(call) => suspension_event_function_call(call),
ReplProgress::OsCall(call) => suspension_event_os_call(call),
ReplProgress::NameLookup(lookup) => event(pb::child_event::Kind::NameLookup(pb::NameLookup {
name: lookup.name.clone(),
object_id: lookup.object_id().as_ref().map(uuid_to_pb),
position: Some((&lookup.position).into()),
})),
ReplProgress::ResolveFutures(state) => event(pb::child_event::Kind::ResolveFutures(pb::ResolveFutures {
pending_call_ids: state.pending_call_ids().to_vec().into(),
position: Some(state.position().into()),
})),
ReplProgress::Complete { .. } => unreachable!("Complete is handled before suspension_event"),
}
}
enum FuturesReply {
Eager(Result<MontyObject, MontyException>),
Batch(Vec<(u32, ExtFunctionResult)>),
}
fn eager_result(
results: Vec<(u32, ExtFunctionResult)>,
call_id: u32,
) -> Result<Result<MontyObject, MontyException>, &'static str> {
match <[_; 1]>::try_from(results) {
Ok([(id, ExtFunctionResult::Return(value))]) if id == call_id => Ok(Ok(value)),
Ok([(id, ExtFunctionResult::Error(exc))]) if id == call_id => Ok(Err(exc)),
Ok([(id, _)]) if id == call_id => Err("eager coroutine must resolve to a value or exception"),
_ => Err("eager ResumeFutures must contain exactly the suspended call id"),
}
}
struct ProtoPrint<'a> {
segments: Vec<pb::PrintSegment>,
buffered_bytes: usize,
sink: &'a mut dyn EventSink,
interval: Duration,
buffered_since: Option<Instant>,
}
impl<'a> ProtoPrint<'a> {
const FLUSH_BYTES: usize = 8 * 1024;
fn new(sink: &'a mut dyn EventSink, interval: Duration) -> Self {
Self {
segments: Vec::new(),
buffered_bytes: 0,
sink,
interval,
buffered_since: None,
}
}
fn flush(&mut self) -> Result<(), MontyException> {
if self.segments.is_empty() {
return Ok(());
}
self.buffered_since = None;
self.buffered_bytes = 0;
let segments = mem::take(&mut self.segments);
self.send(segments)
}
fn flush_lines(&mut self) -> Result<(), MontyException> {
while let Some((index, end)) = self.first_line_end() {
let mut line: Vec<pb::PrintSegment> = self.segments.drain(..index).collect();
let rest = self.segments[0].text.split_off(end);
line.push(pb::PrintSegment {
stream: self.segments[0].stream,
text: mem::replace(&mut self.segments[0].text, rest),
});
if self.segments[0].text.is_empty() {
self.segments.remove(0);
}
self.buffered_bytes -= line.iter().map(|segment| segment.text.len()).sum::<usize>();
self.send(line)?;
}
if self.segments.is_empty() {
self.buffered_since = None;
}
Ok(())
}
fn first_line_end(&self) -> Option<(usize, usize)> {
self.segments
.iter()
.enumerate()
.find_map(|(index, segment)| segment.text.find('\n').map(|at| (index, at + 1)))
}
fn send(&mut self, segments: Vec<pb::PrintSegment>) -> Result<(), MontyException> {
let event = event(pb::child_event::Kind::Print(pb::Print {
segments: segments.into(),
}));
self.sink.send(&event).map_err(|err| {
MontyException::new(
ExcType::RuntimeError,
Some(format!("failed to stream print output: {err}")),
)
})
}
fn maybe_flush(&mut self) -> Result<(), MontyException> {
if self.interval.is_zero() {
self.flush_lines()?;
}
if self.buffered_bytes >= Self::FLUSH_BYTES || (!self.interval.is_zero() && self.interval_elapsed()) {
self.flush()
} else {
Ok(())
}
}
fn interval_elapsed(&self) -> bool {
self.buffered_since
.is_some_and(|since| since.elapsed() >= self.interval)
}
fn mark_buffered(&mut self) {
if self.buffered_since.is_none() {
self.buffered_since = Some(Instant::now());
}
}
fn drain(&mut self) {
let _ = self.flush();
}
fn append(&mut self, stream: PrintStream, text: &str) {
self.mark_buffered();
self.buffered_bytes += text.len();
let stream = wire_stream(stream);
match self.segments.last_mut() {
Some(segment) if segment.stream == stream => segment.text.push_str(text),
_ => self.segments.push(pb::PrintSegment {
stream,
text: text.to_owned(),
}),
}
}
fn write(&mut self, stream: PrintStream, output: &str) -> Result<(), MontyException> {
let mut rest = output;
while !rest.is_empty() {
let take = floor_char_boundary(rest, Self::FLUSH_BYTES - self.buffered_bytes);
if take == 0 {
self.flush()?;
continue;
}
self.append(stream, &rest[..take]);
rest = &rest[take..];
self.maybe_flush()?;
}
Ok(())
}
fn push(&mut self, stream: PrintStream, end: char) -> Result<(), MontyException> {
self.append(stream, end.encode_utf8(&mut [0; 4]));
self.maybe_flush()
}
}
fn wire_stream(stream: PrintStream) -> i32 {
match stream {
PrintStream::Stdout => i32::from(pb::PrintStream::Stdout),
PrintStream::Stderr => i32::from(pb::PrintStream::Stderr),
}
}
impl PrintWriterCallback for ProtoPrint<'_> {
fn stdout_write(&mut self, output: Cow<'_, str>) -> Result<(), MontyException> {
self.write(PrintStream::Stdout, &output)
}
fn stdout_push(&mut self, end: char) -> Result<(), MontyException> {
self.push(PrintStream::Stdout, end)
}
fn stderr_write(&mut self, output: Cow<'_, str>) -> Result<(), MontyException> {
self.write(PrintStream::Stderr, &output)
}
fn stderr_push(&mut self, end: char) -> Result<(), MontyException> {
self.push(PrintStream::Stderr, end)
}
fn poll_flush(&mut self) -> Result<(), MontyException> {
if !self.interval.is_zero() && self.interval_elapsed() {
self.flush()
} else {
Ok(())
}
}
}
fn floor_char_boundary(s: &str, max: usize) -> usize {
if max >= s.len() {
s.len()
} else {
let mut idx = max;
while !s.is_char_boundary(idx) {
idx -= 1;
}
idx
}
}