Merge pull request #2799 from bakaq/callback_streams

Callback streams for use as library
This commit is contained in:
Mark Thom
2025-02-16 22:47:34 -08:00
committed by GitHub
5 changed files with 550 additions and 41 deletions

View File

@@ -181,6 +181,8 @@ pub enum ArenaHeaderTag {
ReadlineStream = 0b110000,
StaticStringStream = 0b110100,
ByteStream = 0b111000,
CallbackStream = 0b111001,
InputChannelStream = 0b111010,
StandardOutputStream = 0b1100,
StandardErrorStream = 0b11000,
NullStream = 0b111100,
@@ -841,6 +843,12 @@ unsafe fn drop_slab_in_place(value: NonNull<AllocSlab>, tag: ArenaHeaderTag) {
ArenaHeaderTag::ByteStream => {
drop_typed_slab_in_place!(ByteStream, value);
}
ArenaHeaderTag::CallbackStream => {
drop_typed_slab_in_place!(CallbackStream, value);
}
ArenaHeaderTag::InputChannelStream => {
drop_typed_slab_in_place!(InputChannelStream, value);
}
ArenaHeaderTag::LiveLoadState | ArenaHeaderTag::InactiveLoadState => {
drop_typed_slab_in_place!(LiveLoadState, value);
}

View File

@@ -1,43 +1,228 @@
use std::borrow::Cow;
use std::io::Write;
use std::sync::mpsc::{channel, Receiver, Sender};
use rand::{rngs::StdRng, SeedableRng};
use crate::Machine;
use super::{
bootstrapping_compile, current_dir, import_builtin_impls, libraries, load_module, Atom,
CompilationTarget, IndexStore, ListingSource, MachineArgs, MachineState, Stream, StreamOptions,
bootstrapping_compile, current_dir, import_builtin_impls, libraries, load_module, Arena, Atom,
Callback, CompilationTarget, IndexStore, ListingSource, MachineArgs, MachineState, Stream,
};
/// Describes how the streams of a [`Machine`](crate::Machine) will be handled.
#[derive(Default)]
enum OutputStreamConfigInner {
#[default]
Memory,
Stdout,
Stderr,
Callback(Callback),
}
impl std::fmt::Debug for OutputStreamConfigInner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Memory => write!(f, "Memory"),
Self::Stdout => write!(f, "Stdout"),
Self::Stderr => write!(f, "Stderr"),
Self::Callback(_) => f.debug_tuple("Callback").field(&"<callback>").finish(),
}
}
}
/// Configuration for an output stream.
#[derive(Debug, Default)]
pub struct OutputStreamConfig {
inner: OutputStreamConfigInner,
}
impl OutputStreamConfig {
/// Sends output to stdout.
pub fn stdout() -> Self {
Self {
inner: OutputStreamConfigInner::Stdout,
}
}
/// Sends output to stderr.
pub fn stderr() -> Self {
Self {
inner: OutputStreamConfigInner::Stderr,
}
}
/// Keeps output in a memory buffer.
pub fn memory() -> Self {
Self {
inner: OutputStreamConfigInner::Memory,
}
}
/// Calls a callback with the output whenever the stream is written to.
pub fn callback(callback: Callback) -> Self {
Self {
inner: OutputStreamConfigInner::Callback(callback),
}
}
fn into_stream(self, arena: &mut Arena) -> Stream {
match self.inner {
OutputStreamConfigInner::Memory => Stream::from_owned_string("".to_owned(), arena),
OutputStreamConfigInner::Stdout => Stream::stdout(arena),
OutputStreamConfigInner::Stderr => Stream::stderr(arena),
OutputStreamConfigInner::Callback(callback) => Stream::from_callback(callback, arena),
}
}
}
#[derive(Debug)]
enum InputStreamConfigInner {
String(String),
Stdin,
Channel(Receiver<Vec<u8>>),
}
impl Default for InputStreamConfigInner {
fn default() -> Self {
Self::String("".into())
}
}
/// Configuration for an input stream;
#[derive(Debug, Default)]
pub struct InputStreamConfig {
inner: InputStreamConfigInner,
}
impl InputStreamConfig {
/// Gets input from string.
pub fn string(s: impl Into<String>) -> Self {
Self {
inner: InputStreamConfigInner::String(s.into()),
}
}
/// Gets input from stdin.
pub fn stdin() -> Self {
Self {
inner: InputStreamConfigInner::Stdin,
}
}
/// Connects the input to the receiving end of a channel.
pub fn channel() -> (UserInput, Self) {
let (sender, receiver) = channel();
(
UserInput { inner: sender },
Self {
inner: InputStreamConfigInner::Channel(receiver),
},
)
}
fn into_stream(self, arena: &mut Arena, add_history: bool) -> Stream {
match self.inner {
InputStreamConfigInner::String(s) => Stream::from_owned_string(s, arena),
InputStreamConfigInner::Stdin => Stream::stdin(arena, add_history),
InputStreamConfigInner::Channel(channel) => Stream::input_channel(channel, arena),
}
}
}
/// Describes how the streams of a [`Machine`](crate::Machine) will be handled.
pub struct StreamConfig {
inner: StreamConfigInner,
user_input: InputStreamConfig,
user_output: OutputStreamConfig,
user_error: OutputStreamConfig,
}
impl Default for StreamConfig {
fn default() -> Self {
Self::in_memory()
}
}
impl StreamConfig {
/// Binds the input, output and error streams to stdin, stdout and stderr.
pub fn stdio() -> Self {
StreamConfig {
inner: StreamConfigInner::Stdio,
user_input: InputStreamConfig::stdin(),
user_output: OutputStreamConfig::stdout(),
user_error: OutputStreamConfig::stderr(),
}
}
/// Binds the output stream to a memory buffer, and the error stream to stderr.
///
/// The input stream is ignored.
/// Binds the output and error streams to memory buffers and has an empty input.
pub fn in_memory() -> Self {
StreamConfig {
inner: StreamConfigInner::Memory,
user_input: InputStreamConfig::string(""),
user_output: OutputStreamConfig::memory(),
user_error: OutputStreamConfig::memory(),
}
}
/// Calls the given callbacks when the respective streams are written to.
///
/// This also returns a handler to the stdin of the [`Machine`](crate::Machine).
pub fn from_callbacks(stdout: Option<Callback>, stderr: Option<Callback>) -> (UserInput, Self) {
let (user_input, channel_stream) = InputStreamConfig::channel();
(
user_input,
StreamConfig {
user_input: channel_stream,
user_output: stdout
.map_or_else(OutputStreamConfig::memory, OutputStreamConfig::callback),
user_error: stderr
.map_or_else(OutputStreamConfig::memory, OutputStreamConfig::callback),
},
)
}
/// Configures the `user_input` stream.
pub fn with_user_input(self, user_input: InputStreamConfig) -> Self {
Self { user_input, ..self }
}
/// Configures the `user_output` stream.
pub fn with_user_output(self, user_output: OutputStreamConfig) -> Self {
Self {
user_output,
..self
}
}
/// Configures the `user_error` stream.
pub fn with_user_error(self, user_error: OutputStreamConfig) -> Self {
Self { user_error, ..self }
}
fn into_streams(self, arena: &mut Arena, add_history: bool) -> (Stream, Stream, Stream) {
(
self.user_input.into_stream(arena, add_history),
self.user_output.into_stream(arena),
self.user_error.into_stream(arena),
)
}
}
#[derive(Default)]
enum StreamConfigInner {
Stdio,
#[default]
Memory,
/// A handler for the stdin of the [`Machine`](crate::Machine).
#[derive(Debug)]
pub struct UserInput {
inner: Sender<Vec<u8>>,
}
impl Write for UserInput {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.inner
.send(buf.into())
.map(|_| buf.len())
.map_err(|_| std::io::ErrorKind::BrokenPipe.into())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
/// Describes how a [`Machine`](crate::Machine) will be configured.
@@ -79,18 +264,9 @@ impl MachineBuilder {
let args = MachineArgs::new();
let mut machine_st = MachineState::new();
let (user_input, user_output, user_error) = match self.streams.inner {
StreamConfigInner::Stdio => (
Stream::stdin(&mut machine_st.arena, args.add_history),
Stream::stdout(&mut machine_st.arena),
Stream::stderr(&mut machine_st.arena),
),
StreamConfigInner::Memory => (
Stream::Null(StreamOptions::default()),
Stream::from_owned_string("".to_owned(), &mut machine_st.arena),
Stream::stderr(&mut machine_st.arena),
),
};
let (user_input, user_output, user_error) = self
.streams
.into_streams(&mut machine_st.arena, args.add_history);
let mut wam = Machine {
machine_st,

View File

@@ -6,11 +6,13 @@ use crate::heap_iter::{stackful_post_order_iter, NonListElider};
use crate::machine::machine_indices::VarKey;
use crate::machine::mock_wam::CompositeOpDir;
use crate::machine::{
F64Offset, F64Ptr, Fixnum, Number, BREAK_FROM_DISPATCH_LOOP_LOC, LIB_QUERY_SUCCESS,
ArenaHeaderTag, F64Offset, F64Ptr, Fixnum, Number, BREAK_FROM_DISPATCH_LOOP_LOC,
LIB_QUERY_SUCCESS,
};
use crate::parser::ast::{Var, VarPtr};
use crate::parser::parser::{Parser, Tokens};
use crate::read::{write_term_to_heap, TermWriteResult};
use crate::types::UntypedArenaPtr;
use dashu::{Integer, Rational};
use indexmap::IndexMap;
@@ -280,11 +282,32 @@ impl Term {
(HeapCellValueTag::Fixnum, n) => {
term_stack.push(Term::Integer(n.into()));
}
(HeapCellValueTag::Cons) => {
match Number::try_from(addr) {
Ok(Number::Integer(i)) => term_stack.push(Term::Integer((*i).clone())),
Ok(Number::Rational(r)) => term_stack.push(Term::Rational((*r).clone())),
_ => {}
(HeapCellValueTag::Cons, ptr) => {
if let Ok(n) = Number::try_from(addr) {
match n {
Number::Integer(i) => term_stack.push(Term::Integer((*i).clone())),
Number::Rational(r) => term_stack.push(Term::Rational((*r).clone())),
_ => { unreachable!() },
}
} else {
match_untyped_arena_ptr!(ptr,
(ArenaHeaderTag::Stream, stream) => {
let stream_term = if let Some(alias) = stream.options().get_alias() {
Term::atom(alias.as_str().to_string())
} else {
Term::compound("$stream", [
Term::integer(stream.as_ptr() as usize)
])
};
term_stack.push(stream_term);
}
(ArenaHeaderTag::Dropped, _stream) => {
term_stack.push(Term::atom("$dropped_value"));
}
_ => {
unreachable!();
}
);
}
}
(HeapCellValueTag::CStr, s) => {
@@ -394,6 +417,7 @@ impl Term {
}
*/
_ => {
unreachable!();
}
);
}

View File

@@ -24,10 +24,13 @@ use std::fs::{File, OpenOptions};
use std::hash::Hash;
use std::io;
use std::io::{Cursor, ErrorKind, Read, Seek, SeekFrom, Write};
use std::mem::ManuallyDrop;
use std::net::{Shutdown, TcpStream};
use std::ops::{Deref, DerefMut};
use std::path::PathBuf;
use std::ptr;
use std::sync::mpsc::Receiver;
use std::sync::mpsc::TryRecvError;
#[cfg(feature = "tls")]
use native_tls::TlsStream;
@@ -375,6 +378,89 @@ impl Write for StandardErrorStream {
}
}
pub type Callback = Box<dyn FnMut(&mut Cursor<Vec<u8>>)>;
pub struct CallbackStream {
pub(crate) inner: Cursor<Vec<u8>>,
callback: Callback,
}
impl Debug for CallbackStream {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CallbackStream")
.field("inner", &self.inner)
.finish()
}
}
impl Write for CallbackStream {
#[inline]
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let pos = self.inner.position();
self.inner.seek(SeekFrom::End(0))?;
let result = self.inner.write(buf);
self.inner.seek(SeekFrom::Start(pos))?;
result
}
#[inline]
fn flush(&mut self) -> std::io::Result<()> {
(self.callback)(&mut self.inner);
self.inner.flush()
}
}
#[derive(Debug)]
pub struct InputChannelStream {
pub(crate) inner: Cursor<Vec<u8>>,
pub eof: bool,
channel: Receiver<Vec<u8>>,
}
impl Read for InputChannelStream {
#[inline]
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.eof {
return Ok(0);
}
let to_read = buf.len();
let mut total_read = 0;
loop {
total_read += self.inner.read(&mut buf[total_read..])?;
if total_read < to_read {
// We need to get more data to read
match self.channel.try_recv() {
Ok(data) => {
// Append into self.inner
let pos = self.inner.position();
assert_eq!(pos as usize, self.inner.get_ref().len());
self.inner.write_all(&data)?;
self.inner.seek(SeekFrom::Start(pos))?;
}
Err(TryRecvError::Empty) => {
// Data is pending
break;
}
Err(TryRecvError::Disconnected) => {
// The other end of the channel was closed so we are EOF
self.eof = true;
break;
}
}
} else {
assert_eq!(total_read, to_read);
break;
}
}
Ok(total_read)
}
}
#[bitfield]
#[repr(u64)]
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
@@ -500,6 +586,8 @@ arena_allocated_impl_for_stream!(ReadlineStream, ReadlineStream);
arena_allocated_impl_for_stream!(StaticStringStream, StaticStringStream);
arena_allocated_impl_for_stream!(StandardOutputStream, StandardOutputStream);
arena_allocated_impl_for_stream!(StandardErrorStream, StandardErrorStream);
arena_allocated_impl_for_stream!(CharReader<CallbackStream>, CallbackStream);
arena_allocated_impl_for_stream!(CharReader<InputChannelStream>, InputChannelStream);
#[derive(Debug, Copy, Clone)]
pub enum Stream {
@@ -518,6 +606,8 @@ pub enum Stream {
Readline(TypedArenaPtr<ReadlineStream>),
StandardOutput(TypedArenaPtr<StandardOutputStream>),
StandardError(TypedArenaPtr<StandardErrorStream>),
Callback(TypedArenaPtr<CallbackStream>),
InputChannel(TypedArenaPtr<InputChannelStream>),
}
impl From<TypedArenaPtr<ReadlineStream>> for Stream {
@@ -548,6 +638,19 @@ impl Stream {
))
}
#[inline]
pub fn input_channel(channel: Receiver<Vec<u8>>, arena: &mut Arena) -> Stream {
let inner = Cursor::new(Vec::new());
Stream::InputChannel(arena_alloc!(
StreamLayout::new(CharReader::new(InputChannelStream {
inner,
eof: false,
channel
})),
arena
))
}
#[inline]
pub fn stdin(arena: &mut Arena, add_history: bool) -> Stream {
Stream::Readline(arena_alloc!(
@@ -581,6 +684,10 @@ impl Stream {
ArenaHeaderTag::Dropped | ArenaHeaderTag::NullStream => {
Stream::Null(StreamOptions::default())
}
ArenaHeaderTag::CallbackStream => Stream::Callback(unsafe { ptr.as_typed_ptr() }),
ArenaHeaderTag::InputChannelStream => {
Stream::InputChannel(unsafe { ptr.as_typed_ptr() })
}
_ => unreachable!(),
}
}
@@ -617,6 +724,8 @@ impl Stream {
Stream::Readline(ptr) => ptr.header_ptr(),
Stream::StandardOutput(ptr) => ptr.header_ptr(),
Stream::StandardError(ptr) => ptr.header_ptr(),
Stream::Callback(ptr) => ptr.header_ptr(),
Stream::InputChannel(ptr) => ptr.header_ptr(),
}
}
@@ -637,6 +746,8 @@ impl Stream {
Stream::Readline(ref ptr) => &ptr.options,
Stream::StandardOutput(ref ptr) => &ptr.options,
Stream::StandardError(ref ptr) => &ptr.options,
Stream::Callback(ref ptr) => &ptr.options,
Stream::InputChannel(ref ptr) => &ptr.options,
}
}
@@ -657,6 +768,8 @@ impl Stream {
Stream::Readline(ref mut ptr) => &mut ptr.options,
Stream::StandardOutput(ref mut ptr) => &mut ptr.options,
Stream::StandardError(ref mut ptr) => &mut ptr.options,
Stream::Callback(ref mut ptr) => &mut ptr.options,
Stream::InputChannel(ref mut ptr) => &mut ptr.options,
}
}
@@ -678,6 +791,8 @@ impl Stream {
Stream::Readline(ptr) => ptr.lines_read += incr_num_lines_read,
Stream::StandardOutput(ptr) => ptr.lines_read += incr_num_lines_read,
Stream::StandardError(ptr) => ptr.lines_read += incr_num_lines_read,
Stream::Callback(ptr) => ptr.lines_read += incr_num_lines_read,
Stream::InputChannel(ptr) => ptr.lines_read += incr_num_lines_read,
}
}
@@ -699,6 +814,8 @@ impl Stream {
Stream::Readline(ptr) => ptr.lines_read = value,
Stream::StandardOutput(ptr) => ptr.lines_read = value,
Stream::StandardError(ptr) => ptr.lines_read = value,
Stream::Callback(ptr) => ptr.lines_read = value,
Stream::InputChannel(ptr) => ptr.lines_read = value,
}
}
@@ -720,6 +837,8 @@ impl Stream {
Stream::Readline(ptr) => ptr.lines_read,
Stream::StandardOutput(ptr) => ptr.lines_read,
Stream::StandardError(ptr) => ptr.lines_read,
Stream::Callback(ptr) => ptr.lines_read,
Stream::InputChannel(ptr) => ptr.lines_read,
}
}
}
@@ -736,6 +855,7 @@ impl CharRead for Stream {
Stream::Readline(rl_stream) => (*rl_stream).peek_char(),
Stream::StaticString(src) => (*src).peek_char(),
Stream::Byte(cursor) => (*cursor).peek_char(),
Stream::InputChannel(cursor) => (*cursor).peek_char(),
#[cfg(feature = "http")]
Stream::HttpWrite(_) => Some(Err(std::io::Error::new(
ErrorKind::PermissionDenied,
@@ -744,7 +864,8 @@ impl CharRead for Stream {
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Null(_) => Some(Err(std::io::Error::new(
| Stream::Null(_)
| Stream::Callback(_) => Some(Err(std::io::Error::new(
ErrorKind::PermissionDenied,
StreamError::ReadFromOutputStream,
))),
@@ -762,6 +883,7 @@ impl CharRead for Stream {
Stream::Readline(rl_stream) => (*rl_stream).read_char(),
Stream::StaticString(src) => (*src).read_char(),
Stream::Byte(cursor) => (*cursor).read_char(),
Stream::InputChannel(cursor) => (*cursor).read_char(),
#[cfg(feature = "http")]
Stream::HttpWrite(_) => Some(Err(std::io::Error::new(
ErrorKind::PermissionDenied,
@@ -770,7 +892,8 @@ impl CharRead for Stream {
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Null(_) => Some(Err(std::io::Error::new(
| Stream::Null(_)
| Stream::Callback(_) => Some(Err(std::io::Error::new(
ErrorKind::PermissionDenied,
StreamError::ReadFromOutputStream,
))),
@@ -793,7 +916,9 @@ impl CharRead for Stream {
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Null(_) => {}
| Stream::Null(_)
| Stream::Callback(_) => {}
Stream::InputChannel(_) => {}
}
}
@@ -808,12 +933,14 @@ impl CharRead for Stream {
Stream::Readline(ref mut rl_stream) => rl_stream.consume(nread),
Stream::StaticString(ref mut src) => src.consume(nread),
Stream::Byte(ref mut cursor) => cursor.consume(nread),
Stream::InputChannel(ref mut cursor) => cursor.consume(nread),
#[cfg(feature = "http")]
Stream::HttpWrite(_) => {}
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Null(_) => {}
| Stream::Null(_)
| Stream::Callback(_) => {}
}
}
}
@@ -831,6 +958,7 @@ impl Read for Stream {
Stream::Readline(rl_stream) => (*rl_stream).read(buf),
Stream::StaticString(src) => (*src).read(buf),
Stream::Byte(cursor) => (*cursor).read(buf),
Stream::InputChannel(cursor) => (*cursor).read(buf),
#[cfg(feature = "http")]
Stream::HttpWrite(_) => Err(std::io::Error::new(
ErrorKind::PermissionDenied,
@@ -839,7 +967,8 @@ impl Read for Stream {
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Null(_) => Err(std::io::Error::new(
| Stream::Null(_)
| Stream::Callback(_) => Err(std::io::Error::new(
ErrorKind::PermissionDenied,
StreamError::ReadFromOutputStream,
)),
@@ -855,6 +984,7 @@ impl Write for Stream {
#[cfg(feature = "tls")]
Stream::NamedTls(ref mut tls_stream) => tls_stream.get_mut().write(buf),
Stream::Byte(ref mut cursor) => cursor.get_mut().write(buf),
Stream::Callback(ref mut callback_stream) => callback_stream.get_mut().write(buf),
Stream::StandardOutput(stream) => stream.write(buf),
Stream::StandardError(stream) => stream.write(buf),
#[cfg(feature = "http")]
@@ -865,6 +995,7 @@ impl Write for Stream {
StreamError::WriteToInputStream,
)),
Stream::StaticString(_)
| Stream::InputChannel(_)
| Stream::Readline(_)
| Stream::InputFile(..)
| Stream::Null(_) => Err(std::io::Error::new(
@@ -881,6 +1012,7 @@ impl Write for Stream {
#[cfg(feature = "tls")]
Stream::NamedTls(ref mut tls_stream) => tls_stream.stream.get_mut().flush(),
Stream::Byte(ref mut cursor) => cursor.stream.get_mut().flush(),
Stream::Callback(ref mut callback_stream) => callback_stream.stream.get_mut().flush(),
Stream::StandardError(stream) => stream.stream.flush(),
Stream::StandardOutput(stream) => stream.stream.flush(),
#[cfg(feature = "http")]
@@ -891,6 +1023,7 @@ impl Write for Stream {
StreamError::FlushToInputStream,
)),
Stream::StaticString(_)
| Stream::InputChannel(_)
| Stream::Readline(_)
| Stream::InputFile(_)
| Stream::Null(_) => Err(std::io::Error::new(
@@ -1043,6 +1176,8 @@ impl Stream {
Stream::Readline(stream) => stream.past_end_of_stream,
Stream::StandardOutput(stream) => stream.past_end_of_stream,
Stream::StandardError(stream) => stream.past_end_of_stream,
Stream::Callback(stream) => stream.past_end_of_stream,
Stream::InputChannel(stream) => stream.past_end_of_stream,
}
}
@@ -1069,6 +1204,8 @@ impl Stream {
Stream::Readline(stream) => stream.past_end_of_stream = value,
Stream::StandardOutput(stream) => stream.past_end_of_stream = value,
Stream::StandardError(stream) => stream.past_end_of_stream = value,
Stream::Callback(stream) => stream.past_end_of_stream = value,
Stream::InputChannel(stream) => stream.past_end_of_stream = value,
}
}
@@ -1144,6 +1281,13 @@ impl Stream {
AtEndOfStream::Past
}
}
Stream::InputChannel(stream_layout) => {
if stream_layout.stream.get_ref().eof {
AtEndOfStream::At
} else {
AtEndOfStream::Not
}
}
_ => AtEndOfStream::Not,
}
}
@@ -1168,6 +1312,7 @@ impl Stream {
#[cfg(feature = "tls")]
Stream::NamedTls(..) => atom!("read_append"),
Stream::Byte(_)
| Stream::InputChannel(_)
| Stream::Readline(_)
| Stream::StaticString(_)
| Stream::InputFile(..) => atom!("read"),
@@ -1175,7 +1320,10 @@ impl Stream {
Stream::OutputFile(file) if file.is_append => atom!("append"),
#[cfg(feature = "http")]
Stream::HttpWrite(_) => atom!("write"),
Stream::OutputFile(_) | Stream::StandardError(_) | Stream::StandardOutput(_) => {
Stream::OutputFile(_)
| Stream::StandardError(_)
| Stream::StandardOutput(_)
| Stream::Callback(_) => {
atom!("write")
}
Stream::Null(_) => atom!(""),
@@ -1198,6 +1346,17 @@ impl Stream {
))
}
#[inline]
pub fn from_callback(callback: Callback, arena: &mut Arena) -> Self {
Stream::Callback(arena_alloc!(
ManuallyDrop::new(StreamLayout::new(CharReader::new(CallbackStream {
inner: Cursor::new(Vec::new()),
callback,
}))),
arena
))
}
#[inline]
pub(crate) fn from_tcp_stream(address: Atom, tcp_stream: TcpStream, arena: &mut Arena) -> Self {
tcp_stream.set_read_timeout(None).unwrap();
@@ -1325,6 +1484,14 @@ impl Stream {
stream.drop_payload();
Ok(())
}
Stream::Callback(mut stream) => {
stream.drop_payload();
Ok(())
}
Stream::InputChannel(mut stream) => {
stream.drop_payload();
Ok(())
}
Stream::StaticString(mut stream) => {
stream.drop_payload();
Ok(())
@@ -1352,6 +1519,7 @@ impl Stream {
Stream::HttpRead(..) => true,
Stream::NamedTcp(..)
| Stream::Byte(_)
| Stream::InputChannel(_)
| Stream::Readline(_)
| Stream::StaticString(_)
| Stream::InputFile(..) => true,
@@ -1370,6 +1538,7 @@ impl Stream {
| Stream::StandardOutput(_)
| Stream::NamedTcp(..)
| Stream::Byte(_)
| Stream::Callback(_)
| Stream::OutputFile(..) => true,
_ => false,
}
@@ -1399,6 +1568,10 @@ impl Stream {
readline_stream.reset();
true
}
Stream::InputChannel(ref mut input_channel_stream) => {
input_channel_stream.stream.get_mut().inner.set_position(0);
true
}
_ => false,
}
}
@@ -1924,9 +2097,135 @@ impl MachineState {
}
#[cfg(test)]
mod test {
use super::*;
use crate::machine::config::*;
mod tests {
use crate::*;
use std::{cell::RefCell, io::Read, io::Write, rc::Rc};
fn succeeded(answer: Vec<Result<LeafAnswer, Term>>) -> bool {
// Ideally this should be a method in QueryState or LeafAnswer.
matches!(
answer[0].as_ref(),
Ok(LeafAnswer::True) | Ok(LeafAnswer::LeafAnswer { .. })
)
}
#[test]
#[cfg_attr(miri, ignore)]
fn user_input_string_stream() {
let streams =
StreamConfig::default().with_user_input(InputStreamConfig::string("a(1,2,3)."));
let mut machine = MachineBuilder::default().with_streams(streams).build();
let complete_answer: Vec<_> = machine
.run_query(r#"current_input(_), \+ at_end_of_stream."#)
.collect();
assert!(succeeded(complete_answer));
let complete_answer: Vec<_> = machine.run_query("read(A).").collect();
assert_eq!(
complete_answer,
[Ok(LeafAnswer::from_bindings([(
"A",
Term::compound("a", [Term::integer(1), Term::integer(2), Term::integer(3),])
)]))]
);
let complete_answer: Vec<_> = machine.run_query(r#"at_end_of_stream."#).collect();
assert!(succeeded(complete_answer));
}
#[test]
#[cfg_attr(miri, ignore)]
fn user_input_channel_stream() {
let (mut user_input, channel_stream) = InputStreamConfig::channel();
let streams = StreamConfig::default().with_user_input(channel_stream);
let mut machine = MachineBuilder::default().with_streams(streams).build();
let complete_answer: Vec<_> = machine
.run_query(r#"current_input(_), \+ at_end_of_stream."#)
.collect();
assert!(succeeded(complete_answer));
write!(user_input, "a(1,2,3).").unwrap();
let complete_answer: Vec<_> = machine
.run_query(r#"\+ at_end_of_stream, read(A)."#)
.collect();
assert_eq!(
complete_answer,
[Ok(LeafAnswer::from_bindings([(
"A",
Term::compound("a", [Term::integer(1), Term::integer(2), Term::integer(3),])
)]))]
);
// End-of-data but not end-of-stream;
let complete_answer: Vec<_> = machine
.run_query(
r#"
use_module(library(charsio)),
current_input(In), get_n_chars(In, N, C),
N == 0, \+ at_end_of_stream.
"#,
)
.collect();
assert!(succeeded(complete_answer));
// Dropping the sender closes the input
drop(user_input);
let complete_answer: Vec<_> = machine
.run_query(
r#"
current_input(In), get_n_chars(In, N, _),
N == 0, at_end_of_stream.
"#,
)
.collect();
assert!(succeeded(complete_answer));
}
#[test]
#[cfg_attr(miri, ignore)]
fn user_output_callback_stream() {
let test_string = Rc::new(RefCell::new(String::new()));
let streams =
StreamConfig::default().with_user_output(OutputStreamConfig::callback(Box::new({
let test_string = test_string.clone();
move |x| {
x.read_to_string(&mut test_string.borrow_mut()).unwrap();
}
})));
let mut machine = MachineBuilder::default().with_streams(streams).build();
let complete_answer: Vec<_> = machine
.run_query(r#"current_output(Out), \+ at_end_of_stream(Out)."#)
.collect();
assert!(succeeded(complete_answer));
let complete_answer: Vec<_> = machine
.run_query(r#"write(asdf), nl, flush_output."#)
.collect();
assert!(succeeded(complete_answer));
assert_eq!(test_string.borrow().as_str(), "asdf\n");
let complete_answer: Vec<_> = machine.run_query(r#"write(abcd), flush_output."#).collect();
assert!(succeeded(complete_answer));
assert_eq!(test_string.borrow().as_str(), "asdf\nabcd");
}
#[test]
#[cfg_attr(miri, ignore)]

View File

@@ -305,6 +305,8 @@ macro_rules! match_untyped_arena_ptr_pat {
| ArenaHeaderTag::ReadlineStream
| ArenaHeaderTag::StaticStringStream
| ArenaHeaderTag::ByteStream
| ArenaHeaderTag::CallbackStream
| ArenaHeaderTag::InputChannelStream
| ArenaHeaderTag::StandardOutputStream
| ArenaHeaderTag::StandardErrorStream
};