use chrono::{DateTime, Utc};
use concepts::{
ExecutionId,
prefixed_ulid::RunId,
storage::{LogEntry, LogInfoAppendRow, LogStreamType},
};
use std::{
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
};
use tokio::{io::AsyncWrite, sync::mpsc};
use tracing::{debug, instrument};
use wasmtime_wasi::p2::{StreamError, StreamResult};
#[derive(Clone)]
pub enum StdOutputConfigWithSender {
Stdout,
Stderr,
Db {
sender: mpsc::Sender<LogInfoAppendRow>,
forwarding_from: LogStreamType,
},
}
impl StdOutputConfigWithSender {
#[must_use]
pub fn new(
config: Option<StdOutputConfig>,
log_forwarder_sender: &mpsc::Sender<LogInfoAppendRow>,
forwarding_from: LogStreamType,
) -> Option<Self> {
config.map(|config| match config {
StdOutputConfig::Stdout => StdOutputConfigWithSender::Stdout,
StdOutputConfig::Stderr => StdOutputConfigWithSender::Stderr,
StdOutputConfig::Db => StdOutputConfigWithSender::Db {
sender: log_forwarder_sender.clone(),
forwarding_from,
},
})
}
#[must_use]
pub fn build(&self, execution_id: &ExecutionId, run_id: RunId) -> StdOutput {
match self {
StdOutputConfigWithSender::Stdout => StdOutput::Stdout,
StdOutputConfigWithSender::Stderr => StdOutput::Stderr,
StdOutputConfigWithSender::Db {
sender,
forwarding_from,
} => StdOutput::Db(DbOutput {
execution_id: execution_id.clone(),
run_id,
sender: sender.clone(),
forwarding_from: *forwarding_from,
}),
}
}
}
#[derive(Clone, Copy, Debug)]
pub enum StdOutputConfig {
Stdout,
Stderr,
Db,
}
#[derive(Clone)]
pub enum StdOutput {
Stdout,
Stderr,
Db(DbOutput),
}
#[derive(Clone)]
pub struct DbOutput {
pub sender: mpsc::Sender<LogInfoAppendRow>,
pub execution_id: ExecutionId,
pub run_id: RunId,
pub forwarding_from: LogStreamType,
}
impl DbOutput {
#[instrument(skip_all, fields(execution_id = %self.execution_id, run_id = %self.run_id))]
fn write(&mut self, buf: &[u8]) {
let res = self.sender.try_send(LogInfoAppendRow {
execution_id: self.execution_id.clone(),
run_id: self.run_id,
log_entry: LogEntry::Stream {
created_at: Utc::now(),
payload: Vec::from(buf),
stream_type: self.forwarding_from,
},
});
if res.is_err() {
debug!("Dropping stream message");
}
}
}
#[derive(Clone)]
pub struct OutputEvent {
pub buf: Vec<u8>,
pub created_at: DateTime<Utc>,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum OutOrErr {
Stdout,
Stderr,
}
impl OutOrErr {
fn write_all(&self, buf: &[u8]) -> Result<(), std::io::Error> {
use std::io::Write;
match self {
OutOrErr::Stdout => std::io::stdout().write_all(buf),
OutOrErr::Stderr => std::io::stderr().write_all(buf),
}
}
}
#[derive(Clone)]
pub(crate) struct LogStream {
output: StdOutput,
state: Arc<LogStreamState>,
}
struct LogStreamState {
prefix: String,
needs_prefix_on_next_write: AtomicBool,
}
impl LogStream {
pub(crate) fn new(prefix: String, output: StdOutput) -> LogStream {
LogStream {
output,
state: Arc::new(LogStreamState {
prefix,
needs_prefix_on_next_write: AtomicBool::new(true),
}),
}
}
}
impl wasmtime_wasi::cli::StdoutStream for LogStream {
fn p2_stream(&self) -> Box<dyn wasmtime_wasi::p2::OutputStream> {
Box::new(self.clone())
}
fn async_stream(&self) -> Box<dyn tokio::io::AsyncWrite + Send + Sync> {
Box::new(self.clone())
}
}
impl wasmtime_wasi::cli::IsTerminal for LogStream {
fn is_terminal(&self) -> bool {
match &self.output {
StdOutput::Stdout => std::io::stdout().is_terminal(),
StdOutput::Stderr => std::io::stderr().is_terminal(),
StdOutput::Db { .. } => false,
}
}
}
impl wasmtime_wasi::p2::OutputStream for LogStream {
fn write(&mut self, bytes: bytes::Bytes) -> StreamResult<()> {
self.write_all(&bytes)
.map_err(|e| StreamError::LastOperationFailed(e.into()))?;
Ok(())
}
fn flush(&mut self) -> StreamResult<()> {
Ok(())
}
fn check_write(&mut self) -> StreamResult<usize> {
Ok(1024 * 1024)
}
}
impl LogStream {
fn write_all(&mut self, mut bytes: &[u8]) -> std::io::Result<()> {
let our_or_err = match &mut self.output {
StdOutput::Db(db_output) => {
db_output.write(bytes);
return Ok(());
}
StdOutput::Stdout => OutOrErr::Stdout,
StdOutput::Stderr => OutOrErr::Stderr,
};
while !bytes.is_empty() {
if self
.state
.needs_prefix_on_next_write
.load(Ordering::Relaxed)
{
our_or_err.write_all(self.state.prefix.as_bytes())?;
self.state
.needs_prefix_on_next_write
.store(false, Ordering::Relaxed);
}
if let Some(i) = bytes.iter().position(|b| *b == b'\n') {
let (a, b) = bytes.split_at(i + 1);
bytes = b;
our_or_err.write_all(a)?;
self.state
.needs_prefix_on_next_write
.store(true, Ordering::Relaxed);
} else {
our_or_err.write_all(bytes)?;
break;
}
}
Ok(())
}
}
impl AsyncWrite for LogStream {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Poll::Ready(self.write_all(buf).map(|()| buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[async_trait::async_trait]
impl wasmtime_wasi::p2::Pollable for LogStream {
async fn ready(&mut self) {}
}