diff --git a/src/arena.rs b/src/arena.rs index e5aceb4d..113ff6f0 100644 --- a/src/arena.rs +++ b/src/arena.rs @@ -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, 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); } diff --git a/src/machine/config.rs b/src/machine/config.rs index 2981899d..b9d3f703 100644 --- a/src/machine/config.rs +++ b/src/machine/config.rs @@ -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(&"").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>), +} + +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) -> 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, stderr: Option) -> (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>, +} + +impl Write for UserInput { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + 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, diff --git a/src/machine/lib_machine/mod.rs b/src/machine/lib_machine/mod.rs index 24f22fe7..87b64074 100644 --- a/src/machine/lib_machine/mod.rs +++ b/src/machine/lib_machine/mod.rs @@ -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!(); } ); } diff --git a/src/machine/streams.rs b/src/machine/streams.rs index b1ecba03..500ed679 100644 --- a/src/machine/streams.rs +++ b/src/machine/streams.rs @@ -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>)>; + +pub struct CallbackStream { + pub(crate) inner: Cursor>, + 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 { + 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>, + pub eof: bool, + channel: Receiver>, +} + +impl Read for InputChannelStream { + #[inline] + fn read(&mut self, buf: &mut [u8]) -> std::io::Result { + 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); +arena_allocated_impl_for_stream!(CharReader, InputChannelStream); #[derive(Debug, Copy, Clone)] pub enum Stream { @@ -518,6 +606,8 @@ pub enum Stream { Readline(TypedArenaPtr), StandardOutput(TypedArenaPtr), StandardError(TypedArenaPtr), + Callback(TypedArenaPtr), + InputChannel(TypedArenaPtr), } impl From> for Stream { @@ -548,6 +638,19 @@ impl Stream { )) } + #[inline] + pub fn input_channel(channel: Receiver>, 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>) -> 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)] diff --git a/src/macros.rs b/src/macros.rs index 30a863ca..547ccc00 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -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 };