skippy-server 0.76.1

Embedded Skippy staged runtime server
use crate::frontend::generation::PhaseTimer;
use crate::frontend::util::openai_io_error;
use anyhow::Context;
use anyhow::Result;
use anyhow::anyhow;
use anyhow::bail;
use openai_frontend::OpenAiError;
use openai_frontend::OpenAiResult;
use skippy_protocol::binary::StageReplyStats;
use skippy_protocol::binary::WireReplyKind;
use skippy_protocol::binary::recv_reply;
use std::net::TcpStream;

#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct PrefillChunkSchedule {
    pub(super) sizes: Vec<usize>,
}

impl PrefillChunkSchedule {
    pub(super) fn parse(spec: Option<&str>) -> Result<Option<Self>> {
        let Some(spec) = spec else {
            return Ok(None);
        };
        let spec = spec.trim();
        if spec.is_empty() {
            return Ok(None);
        }
        let mut sizes = Vec::new();
        for part in spec.split(',') {
            let part = part.trim();
            if part.is_empty() {
                bail!("empty chunk size in schedule");
            }
            let size = part
                .parse::<usize>()
                .with_context(|| format!("invalid chunk size '{part}'"))?;
            if size == 0 {
                bail!("chunk sizes must be greater than zero");
            }
            sizes.push(size);
        }
        Ok(Some(Self { sizes }))
    }

    pub(super) fn chunk_size_for(&self, chunk_index: usize) -> usize {
        self.sizes
            .get(chunk_index)
            .copied()
            .or_else(|| self.sizes.last().copied())
            .expect("schedule has at least one size")
    }

    pub(super) fn label(&self) -> String {
        self.sizes
            .iter()
            .map(usize::to_string)
            .collect::<Vec<_>>()
            .join(",")
    }
}

#[derive(Clone, Debug, PartialEq)]
pub(super) enum PrefillChunkPolicy {
    Fixed {
        chunk_size: usize,
    },
    Schedule {
        fixed_chunk_size: usize,
        schedule: PrefillChunkSchedule,
    },
    AdaptiveRamp {
        fixed_chunk_size: usize,
        start: usize,
        step: usize,
        max: usize,
        target_ms: f64,
    },
}

pub(super) struct PrefillChunkPolicyArgs<'a> {
    pub(super) policy: &'a str,
    pub(super) schedule: Option<&'a str>,
    pub(super) fixed_chunk_size: usize,
    pub(super) adaptive_start: usize,
    pub(super) adaptive_step: usize,
    pub(super) adaptive_max: usize,
    pub(super) adaptive_target_ms: f64,
    pub(super) schedule_arg: &'static str,
    pub(super) policy_arg: &'static str,
}

#[derive(Clone, Copy, Debug)]
pub(super) struct PrefillChunkObservation {
    pub(super) compute_ms: f64,
    pub(super) forward_write_ms: f64,
    pub(super) downstream_wait_ms: f64,
}

pub(super) fn representative_prefill_compute_sample(
    current_ms: f64,
    current_tokens: usize,
    candidate_ms: f64,
    candidate_tokens: usize,
) -> (f64, usize) {
    if candidate_tokens > current_tokens
        || (candidate_tokens == current_tokens && candidate_ms > current_ms)
    {
        (candidate_ms, candidate_tokens)
    } else {
        (current_ms, current_tokens)
    }
}

#[derive(Clone, Debug)]
pub(super) struct PrefillChunkPlanner {
    pub(super) policy: PrefillChunkPolicy,
    pub(super) next_adaptive_size: usize,
    duration_ceiling: Option<usize>,
}

impl PrefillChunkPolicy {
    pub(super) fn parse(args: PrefillChunkPolicyArgs<'_>) -> Result<Self> {
        if args.fixed_chunk_size == 0 {
            bail!("prefill chunk size must be greater than zero");
        }
        let normalized = args.policy.trim().to_ascii_lowercase();
        match normalized.as_str() {
            "fixed" => {
                if let Some(schedule) = PrefillChunkSchedule::parse(args.schedule)
                    .with_context(|| format!("invalid {} value", args.schedule_arg))?
                {
                    return Ok(Self::Schedule {
                        fixed_chunk_size: args.fixed_chunk_size,
                        schedule,
                    });
                }
                Ok(Self::Fixed {
                    chunk_size: args.fixed_chunk_size,
                })
            }
            "schedule" => {
                let schedule = PrefillChunkSchedule::parse(args.schedule)
                    .with_context(|| format!("invalid {} value", args.schedule_arg))?
                    .ok_or_else(|| anyhow!("{} requires {}", args.policy_arg, args.schedule_arg))?;
                Ok(Self::Schedule {
                    fixed_chunk_size: args.fixed_chunk_size,
                    schedule,
                })
            }
            "adaptive" | "adaptive-ramp" => {
                if args.adaptive_start == 0
                    || args.adaptive_step == 0
                    || args.adaptive_max == 0
                    || args.adaptive_start > args.adaptive_max
                    || !args.adaptive_target_ms.is_finite()
                    || args.adaptive_target_ms <= 0.0
                {
                    bail!(
                        "{} adaptive-ramp requires positive start/step/max/target-ms with start <= max",
                        args.policy_arg
                    );
                }
                Ok(Self::AdaptiveRamp {
                    fixed_chunk_size: args.fixed_chunk_size,
                    start: args.adaptive_start,
                    step: args.adaptive_step,
                    max: args.adaptive_max,
                    target_ms: args.adaptive_target_ms,
                })
            }
            other => bail!(
                "invalid {} '{}'; expected fixed, schedule, or adaptive-ramp",
                args.policy_arg,
                other
            ),
        }
    }

    pub(super) fn planner(&self) -> PrefillChunkPlanner {
        let next_adaptive_size = match self {
            Self::AdaptiveRamp { start, .. } => *start,
            _ => 0,
        };
        PrefillChunkPlanner {
            policy: self.clone(),
            next_adaptive_size,
            duration_ceiling: None,
        }
    }

    pub(super) fn policy_label(&self) -> &'static str {
        match self {
            Self::Fixed { .. } => "fixed",
            Self::Schedule { .. } => "schedule",
            Self::AdaptiveRamp { .. } => "adaptive-ramp",
        }
    }

    pub(super) fn fixed_chunk_size(&self) -> usize {
        match self {
            Self::Fixed { chunk_size } => *chunk_size,
            Self::Schedule {
                fixed_chunk_size, ..
            }
            | Self::AdaptiveRamp {
                fixed_chunk_size, ..
            } => *fixed_chunk_size,
        }
    }

    pub(super) fn schedule(&self) -> Option<&PrefillChunkSchedule> {
        match self {
            Self::Schedule { schedule, .. } => Some(schedule),
            _ => None,
        }
    }

    pub(super) fn adaptive_params(&self) -> Option<(usize, usize, usize, f64)> {
        match self {
            Self::AdaptiveRamp {
                start,
                step,
                max,
                target_ms,
                ..
            } => Some((*start, *step, *max, *target_ms)),
            _ => None,
        }
    }
}

impl PrefillChunkPlanner {
    pub(super) fn chunk_size_for(
        &mut self,
        chunk_index: usize,
        prefill_token_count: usize,
    ) -> usize {
        match &self.policy {
            PrefillChunkPolicy::Fixed { chunk_size } => *chunk_size,
            PrefillChunkPolicy::Schedule { schedule, .. } => schedule.chunk_size_for(chunk_index),
            PrefillChunkPolicy::AdaptiveRamp {
                fixed_chunk_size, ..
            } if chunk_index == 0 && prefill_token_count <= *fixed_chunk_size => *fixed_chunk_size,
            PrefillChunkPolicy::AdaptiveRamp { .. } => self
                .duration_ceiling
                .map_or(self.next_adaptive_size, |ceiling| {
                    self.next_adaptive_size.min(ceiling)
                }),
        }
    }

    pub(super) fn calibrate_slowest_stage_rate(&mut self, compute_ms_per_token: f64) {
        let PrefillChunkPolicy::AdaptiveRamp {
            start,
            step,
            max,
            target_ms,
            ..
        } = &self.policy
        else {
            return;
        };
        if !compute_ms_per_token.is_finite() || compute_ms_per_token <= 0.0 {
            return;
        }
        let predicted_tokens = (*target_ms / compute_ms_per_token).floor();
        let predicted_tokens = if predicted_tokens.is_finite() && predicted_tokens > 0.0 {
            predicted_tokens as usize
        } else {
            *start
        };
        let stepped = predicted_tokens
            .saturating_sub(*start)
            .checked_div(*step)
            .unwrap_or(0)
            .saturating_mul(*step)
            .saturating_add(*start);
        let ceiling = stepped.clamp(*start, *max);
        self.duration_ceiling = Some(ceiling);
        self.next_adaptive_size = self.next_adaptive_size.min(ceiling);
    }

    pub(super) fn observe(&mut self, observation: PrefillChunkObservation) {
        let PrefillChunkPolicy::AdaptiveRamp {
            start, step, max, ..
        } = &self.policy
        else {
            return;
        };
        let compute_ms = observation.compute_ms.max(0.001);
        let downstream_hidden = observation.downstream_wait_ms <= compute_ms * 0.75
            && observation.forward_write_ms <= compute_ms * 0.25;
        let downstream_exposed = observation.downstream_wait_ms > compute_ms * 1.25
            || observation.forward_write_ms > compute_ms * 0.75;
        if downstream_hidden {
            self.next_adaptive_size = self.next_adaptive_size.saturating_add(*step).min(*max);
        } else if downstream_exposed {
            self.next_adaptive_size = self.next_adaptive_size.saturating_sub(*step).max(*start);
        }
        if let Some(ceiling) = self.duration_ceiling {
            self.next_adaptive_size = self.next_adaptive_size.min(ceiling);
        }
    }

    #[cfg(test)]
    pub(super) fn advance_without_observation(&mut self) {
        let PrefillChunkPolicy::AdaptiveRamp { step, max, .. } = &self.policy else {
            return;
        };
        self.next_adaptive_size = self.next_adaptive_size.saturating_add(*step).min(*max);
    }
}

#[derive(Clone, Copy, Debug, Default)]
pub(super) struct EmbeddedPrefillDrain {
    pub(super) drained_replies: usize,
    pub(super) downstream_wait_ms: f64,
    pub(super) downstream_wait_max_ms: f64,
}

impl EmbeddedPrefillDrain {
    pub(super) fn absorb(&mut self, current: Self) {
        self.drained_replies = self.drained_replies.saturating_add(current.drained_replies);
        self.downstream_wait_ms += current.downstream_wait_ms;
        self.downstream_wait_max_ms = self
            .downstream_wait_max_ms
            .max(current.downstream_wait_max_ms);
    }
}

pub(super) fn drain_one_embedded_prefill_reply(
    downstream: &mut TcpStream,
    pending_prefill_replies: &mut usize,
    stats: &mut StageReplyStats,
) -> OpenAiResult<EmbeddedPrefillDrain> {
    if *pending_prefill_replies == 0 {
        return Ok(EmbeddedPrefillDrain::default());
    }
    let wait_timer = PhaseTimer::start();
    let reply = recv_reply(&mut *downstream).map_err(openai_io_error)?;
    let downstream_wait_ms = wait_timer.elapsed_ms();
    if reply.kind != WireReplyKind::Ack {
        return Err(OpenAiError::backend(format!(
            "expected deferred prefill ACK from downstream, got {:?}",
            reply.kind
        )));
    }
    stats.merge(reply.stats);
    *pending_prefill_replies = pending_prefill_replies.saturating_sub(1);
    Ok(EmbeddedPrefillDrain {
        drained_replies: 1,
        downstream_wait_ms,
        downstream_wait_max_ms: downstream_wait_ms,
    })
}

pub(super) fn drain_embedded_prefill_replies(
    downstream: &mut TcpStream,
    pending_prefill_replies: &mut usize,
    stats: &mut StageReplyStats,
) -> OpenAiResult<EmbeddedPrefillDrain> {
    let mut drained = EmbeddedPrefillDrain::default();
    while *pending_prefill_replies > 0 {
        let current = drain_one_embedded_prefill_reply(downstream, pending_prefill_replies, stats)?;
        drained.absorb(current);
    }
    Ok(drained)
}