use std::{collections::VecDeque, time::Duration};
use tokio::io::AsyncReadExt;
use crate::{
ContainerAsync, Image,
core::{
client::Client,
error::{Result, WaitContainerError, WaitLogError},
logs::LogSource,
},
};
const POLL_INTERVAL: Duration = Duration::from_millis(100);
pub(crate) const DRAIN_GRACE: Duration = Duration::from_secs(2);
const MAX_COLLECTED_LOG_BYTES: usize = 1024 * 1024;
pub(crate) struct CollectedLogs {
buf: VecDeque<u8>,
}
impl CollectedLogs {
pub(crate) fn new() -> Self {
Self {
buf: VecDeque::new(),
}
}
pub(crate) fn push(&mut self, bytes: &[u8]) {
if bytes.is_empty() {
return;
}
let keep = if bytes.len() > MAX_COLLECTED_LOG_BYTES {
&bytes[bytes.len() - MAX_COLLECTED_LOG_BYTES..]
} else {
bytes
};
self.buf.extend(keep.iter().copied());
while self.buf.len() > MAX_COLLECTED_LOG_BYTES {
self.buf.pop_front();
}
}
pub(crate) fn into_chunks(self) -> Vec<Vec<u8>> {
if self.buf.is_empty() {
Vec::new()
} else {
vec![self.buf.into_iter().collect()]
}
}
}
#[derive(Debug, Clone)]
pub struct LogWaitStrategy {
pub(crate) source: LogSource,
pub(crate) message: Vec<u8>,
pub(crate) times: usize,
}
impl LogWaitStrategy {
pub fn stdout(message: impl AsRef<[u8]>) -> Self {
Self::new(LogSource::StdOut, message)
}
pub fn stderr(message: impl AsRef<[u8]>) -> Self {
Self::new(LogSource::StdErr, message)
}
pub fn stdout_or_stderr(message: impl AsRef<[u8]>) -> Self {
Self::new(LogSource::BothStd, message)
}
pub fn new(source: LogSource, message: impl AsRef<[u8]>) -> Self {
Self {
source,
message: message.as_ref().to_vec(),
times: 1,
}
}
pub fn with_times(mut self, times: usize) -> Self {
self.times = times.max(1);
self
}
}
impl LogWaitStrategy {
pub(crate) async fn wait_until_ready<I: Image>(
self,
_client: &Client,
container: &ContainerAsync<I>,
) -> Result<()> {
if self.message.is_empty() {
return Ok(());
}
let mut readers = match self.source {
LogSource::StdOut => vec![container.stdout(true)],
LogSource::StdErr => vec![container.stderr(true)],
LogSource::BothStd => vec![container.stdout(true), container.stderr(true)],
};
let mut matchers: Vec<StreamMatcher> = readers
.iter()
.map(|_| StreamMatcher::new(self.message.clone()))
.collect();
let mut chunk = vec![0u8; 8192];
let mut exited_at: Option<tokio::time::Instant> = None;
let mut collected = CollectedLogs::new();
loop {
let mut progressed = false;
let mut total = 0usize;
for (reader, matcher) in readers.iter_mut().zip(matchers.iter_mut()) {
let n = match tokio::time::timeout(POLL_INTERVAL, reader.read(&mut chunk)).await {
Ok(Ok(n)) => n,
Ok(Err(e)) => return Err(WaitLogError::Io(e).into()),
Err(_elapsed) => 0,
};
let count = if n > 0 {
progressed = true;
collected.push(&chunk[..n]);
matcher.feed(&chunk[..n])
} else {
matcher.feed(&[])
};
total = total.saturating_add(count);
}
if total >= self.times {
return Ok(());
}
if progressed {
continue;
}
if container.exit_code_hint().is_some() || container.logs_terminated() {
let at = exited_at.get_or_insert_with(tokio::time::Instant::now);
if at.elapsed() >= DRAIN_GRACE {
return Err(WaitContainerError::WaitLog(WaitLogError::EndOfStream(
collected.into_chunks(),
))
.into());
}
}
tokio::time::sleep(POLL_INTERVAL).await;
}
}
}
pub(crate) struct StreamMatcher {
pattern: Vec<u8>,
carry: Vec<u8>,
count: usize,
}
impl StreamMatcher {
pub(crate) fn new(pattern: Vec<u8>) -> Self {
Self {
pattern,
carry: Vec::new(),
count: 0,
}
}
pub(crate) fn feed(&mut self, chunk: &[u8]) -> usize {
if self.pattern.is_empty() {
return usize::MAX;
}
let mut buf = std::mem::take(&mut self.carry);
buf.extend_from_slice(chunk);
if buf.len() >= self.pattern.len() {
self.count += buf
.windows(self.pattern.len())
.filter(|w| *w == self.pattern.as_slice())
.count();
let keep = self.pattern.len() - 1;
self.carry = buf[buf.len() - keep.min(buf.len())..].to_vec();
} else {
self.carry = buf;
}
self.count
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn collected_logs_trims_front_when_over_limit() {
let mut logs = CollectedLogs::new();
let first = vec![b'a'; MAX_COLLECTED_LOG_BYTES];
let second = b"TAIL";
logs.push(&first);
logs.push(second);
let chunks = logs.into_chunks();
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].len(), MAX_COLLECTED_LOG_BYTES);
assert!(chunks[0].ends_with(second));
}
#[test]
fn collected_logs_keeps_all_when_exactly_at_limit() {
let mut logs = CollectedLogs::new();
let exact = vec![b'x'; MAX_COLLECTED_LOG_BYTES];
logs.push(&exact);
let chunks = logs.into_chunks();
assert_eq!(chunks, vec![exact]);
}
#[test]
fn collected_logs_keeps_tail_of_oversized_single_push() {
let mut logs = CollectedLogs::new();
let mut oversized = vec![b'0'; MAX_COLLECTED_LOG_BYTES];
oversized.extend_from_slice(b"END!");
logs.push(&oversized);
let chunks = logs.into_chunks();
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].len(), MAX_COLLECTED_LOG_BYTES);
assert!(chunks[0].ends_with(b"END!"));
assert_eq!(
chunks[0],
oversized[oversized.len() - MAX_COLLECTED_LOG_BYTES..]
);
}
#[test]
fn collected_logs_survives_many_one_byte_pushes() {
let mut logs = CollectedLogs::new();
let total = MAX_COLLECTED_LOG_BYTES + 100;
for i in 0..total {
logs.push(&[((i % 256) as u8)]);
}
let chunks = logs.into_chunks();
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].len(), MAX_COLLECTED_LOG_BYTES);
let start = total - MAX_COLLECTED_LOG_BYTES;
let expected: Vec<u8> = (start..total).map(|i| (i % 256) as u8).collect();
assert_eq!(chunks[0], expected);
}
#[test]
fn collected_logs_into_chunks_is_empty_or_single() {
assert!(CollectedLogs::new().into_chunks().is_empty());
let mut logs = CollectedLogs::new();
logs.push(b"abc");
logs.push(b"def");
let chunks = logs.into_chunks();
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0], b"abcdef");
}
#[test]
fn stream_matcher_empty_pattern_always_matches() {
let mut m = StreamMatcher::new(Vec::new());
assert_eq!(m.feed(b"anything"), usize::MAX);
}
fn count_occurrences(pattern: &[u8], data: &[u8]) -> usize {
if pattern.is_empty() {
return usize::MAX;
}
data.windows(pattern.len())
.filter(|w| *w == pattern)
.count()
}
#[test]
fn stream_matcher_chunking_preserves_count_across_boundaries() {
let pattern = b"ab";
let data = b"xxababxxab";
let expected = count_occurrences(pattern, data);
let mut matcher = StreamMatcher::new(pattern.to_vec());
let mut actual = 0;
for chunk in [b"xxa".as_slice(), b"bab".as_slice(), b"xxab".as_slice()] {
actual = matcher.feed(chunk);
}
assert_eq!(actual, expected);
}
#[test]
fn stream_matcher_chunking_handles_split_pattern() {
let mut matcher = StreamMatcher::new(b"ready".to_vec());
assert_eq!(matcher.feed(b"re"), 0);
assert_eq!(matcher.feed(b"ady"), 1);
}
#[test]
fn with_times_zero_is_clamped_to_one() {
let strategy = LogWaitStrategy::stdout("ready").with_times(0);
assert_eq!(strategy.times, 1, "with_times(0) は 1 にクランプされること");
}
#[test]
fn with_times_positive_value_is_preserved() {
let strategy = LogWaitStrategy::stdout("ready").with_times(3);
assert_eq!(strategy.times, 3, "正の値は変更されないこと");
}
}