use anyhow::{Context, Result, anyhow, bail};
use skippy_ffi::{
ACTIVATION_MAX_PARTS, ActivationBoundaryDesc as RawActivationBoundaryDesc,
ActivationDesc as RawActivationDesc, ActivationPartDesc as RawActivationPartDesc,
GenerationSignalWindow as RawGenerationSignalWindow,
KvPageComponentDesc as RawKvPageComponentDesc, KvPageDesc as RawKvPageDesc,
LogitBias as RawLogitBias, MAX_DRY_SEQUENCE_BREAKER_BYTES, MAX_DRY_SEQUENCE_BREAKERS,
MAX_SAMPLERS, SamplingConfig as RawSamplingConfig, TensorRole, TokenSignal as RawTokenSignal,
};
pub const MAX_LOGIT_BIAS: usize = 256;
pub const ACTIVATION_BOUNDARY_DESC_VERSION: u32 = skippy_ffi::ACTIVATION_BOUNDARY_DESC_VERSION;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelStateKind {
Dense,
Recurrent,
Hybrid,
Diffusion,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoadedModelCapability {
pub state_kind: ModelStateKind,
pub has_indexer_memory: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ActivationPartDesc {
pub identity: [u8; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
pub ggml_type: u32,
pub rank: u32,
pub token_axis: i32,
pub flags: u32,
pub dimensions: [i64; skippy_ffi::ACTIVATION_MAX_DIMS],
pub byte_strides: [u64; skippy_ffi::ACTIVATION_MAX_DIMS],
pub payload_offset: u64,
pub payload_bytes: u64,
}
impl From<RawActivationPartDesc> for ActivationPartDesc {
fn from(raw: RawActivationPartDesc) -> Self {
Self {
identity: raw.identity,
ggml_type: raw.ggml_type,
rank: raw.rank,
token_axis: raw.token_axis,
flags: raw.flags,
dimensions: raw.dimensions,
byte_strides: raw.byte_strides,
payload_offset: raw.payload_offset,
payload_bytes: raw.payload_bytes,
}
}
}
impl ActivationPartDesc {
fn as_raw(self) -> RawActivationPartDesc {
RawActivationPartDesc {
identity: self.identity,
ggml_type: self.ggml_type,
rank: self.rank,
token_axis: self.token_axis,
flags: self.flags,
dimensions: self.dimensions,
byte_strides: self.byte_strides,
payload_offset: self.payload_offset,
payload_bytes: self.payload_bytes,
}
}
pub fn is_optional(self) -> bool {
self.flags & skippy_ffi::ACTIVATION_PART_OPTIONAL != 0
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ActivationBoundaryDesc {
pub version: u32,
pub part_count: u32,
pub frontier_identity: [u8; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
pub parts: [ActivationPartDesc; ACTIVATION_MAX_PARTS],
}
impl From<RawActivationBoundaryDesc> for ActivationBoundaryDesc {
fn from(raw: RawActivationBoundaryDesc) -> Self {
Self {
version: raw.version,
part_count: raw.part_count,
frontier_identity: raw.frontier_identity,
parts: raw.parts.map(ActivationPartDesc::from),
}
}
}
impl ActivationBoundaryDesc {
pub fn parts(&self) -> Result<&[ActivationPartDesc]> {
if self.version != ACTIVATION_BOUNDARY_DESC_VERSION {
bail!(
"unsupported activation boundary descriptor version {}",
self.version
);
}
let count =
usize::try_from(self.part_count).context("activation part count exceeds usize")?;
if count == 0 || count > self.parts.len() {
bail!("activation boundary part count is invalid");
}
Ok(&self.parts[..count])
}
pub fn raw_f32_width(self, edge: &str) -> Result<i32> {
let primary = *self
.parts()?
.first()
.context("activation boundary has no primary part")?;
if primary.ggml_type != crate::GGML_TYPE_F32 || primary.token_axis < 0 {
bail!("{edge} primary activation part is not token-indexed F32");
}
let rank = usize::try_from(primary.rank)
.with_context(|| format!("{edge} primary activation rank exceeds usize"))?;
let token_axis = usize::try_from(primary.token_axis)
.with_context(|| format!("{edge} primary activation token axis is negative"))?;
if rank == 0 || rank > primary.dimensions.len() || token_axis >= rank {
bail!("{edge} primary activation part has an invalid rank or token axis");
}
let mut elements = 1_u64;
for (axis, dimension) in primary.dimensions.iter().copied().enumerate().take(rank) {
if axis == token_axis {
continue;
}
let dimension = u64::try_from(dimension)
.with_context(|| format!("{edge} primary activation dimension is dynamic"))?;
elements = elements
.checked_mul(dimension)
.context("activation element count overflow")?;
}
i32::try_from(elements)
.with_context(|| format!("graph-observed {edge} activation width exceeds i32"))
}
pub fn payload_bytes_hint(self, edge: &str, token_count: u32) -> Result<Option<u64>> {
let mut total = 0_u64;
for part in self.parts()? {
if part.token_axis < 0
|| part.token_axis as u32 >= part.rank
|| part.rank == 0
|| part.rank > 4
{
bail!("{edge} activation part has an invalid token axis");
}
let mut elements = 1_u64;
for (axis, dimension) in part
.dimensions
.iter()
.copied()
.enumerate()
.take(part.rank as usize)
{
let dimension = if axis == part.token_axis as usize {
u64::from(token_count)
} else {
match u64::try_from(dimension) {
Ok(dimension) => dimension,
Err(_) => return Ok(None),
}
};
elements = elements
.checked_mul(dimension)
.context("activation element count overflow")?;
}
let bytes = elements
.checked_mul(ggml_type_bytes(part.ggml_type)?)
.context("activation payload bytes overflow")?;
total = total
.checked_add(bytes)
.context("activation payload bytes overflow")?;
}
Ok(Some(total))
}
pub fn payload_bytes(self, edge: &str, token_count: u32) -> Result<u64> {
self.payload_bytes_hint(edge, token_count)?
.with_context(|| format!("{edge} activation part has an unresolved dynamic dimension"))
}
}
fn ggml_type_bytes(ggml_type: u32) -> Result<u64> {
match ggml_type {
crate::GGML_TYPE_F32 | crate::GGML_TYPE_I32 => Ok(4),
1 | 30 => Ok(2),
other => bail!("activation transport does not support ggml type {other}"),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TensorInfo {
pub name: String,
pub layer_index: Option<u32>,
pub role: TensorRole,
pub ggml_type: u32,
pub byte_size: u64,
pub element_count: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ActivationDesc {
pub version: u32,
pub producer_stage_index: i32,
pub layer_start: i32,
pub layer_end: i32,
pub token_count: u32,
pub sequence_count: u32,
pub part_count: u32,
pub payload_bytes: u64,
pub frontier_identity: [u8; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
pub parts: [ActivationPartDesc; ACTIVATION_MAX_PARTS],
}
impl ActivationDesc {
pub fn parts(&self) -> Result<&[ActivationPartDesc]> {
let count =
usize::try_from(self.part_count).context("activation part count exceeds usize")?;
if count > self.parts.len() {
bail!("activation part count exceeds maximum");
}
Ok(&self.parts[..count])
}
pub(crate) fn as_raw(&self) -> RawActivationDesc {
RawActivationDesc {
version: self.version,
producer_stage_index: self.producer_stage_index,
layer_start: self.layer_start,
layer_end: self.layer_end,
token_count: self.token_count,
sequence_count: self.sequence_count,
part_count: self.part_count,
reserved: 0,
payload_bytes: self.payload_bytes,
frontier_identity: self.frontier_identity,
parts: self.parts.map(ActivationPartDesc::as_raw),
}
}
}
impl From<RawActivationDesc> for ActivationDesc {
fn from(raw: RawActivationDesc) -> Self {
Self {
version: raw.version,
producer_stage_index: raw.producer_stage_index,
layer_start: raw.layer_start,
layer_end: raw.layer_end,
token_count: raw.token_count,
sequence_count: raw.sequence_count,
part_count: raw.part_count,
payload_bytes: raw.payload_bytes,
frontier_identity: raw.frontier_identity,
parts: raw.parts.map(ActivationPartDesc::from),
}
}
}
pub(crate) fn empty_raw_activation_desc() -> RawActivationDesc {
RawActivationDesc {
version: 0,
producer_stage_index: -1,
layer_start: 0,
layer_end: 0,
token_count: 0,
sequence_count: 0,
part_count: 0,
reserved: 0,
payload_bytes: 0,
frontier_identity: [0; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
parts: [RawActivationPartDesc::default(); ACTIVATION_MAX_PARTS],
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ActivationFrame {
pub desc: ActivationDesc,
pub payload: Vec<u8>,
}
impl ActivationFrame {
pub(crate) fn validate_payload_len(&self) -> Result<()> {
let payload_len = u64::try_from(self.payload.len())
.map_err(|_| anyhow!("activation payload length exceeds u64"))?;
if self.desc.payload_bytes != payload_len {
return Err(anyhow!(
"activation payload length {} does not match descriptor payload_bytes {}",
self.payload.len(),
self.desc.payload_bytes
));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct RuntimeKvPageComponentDesc {
pub version: u32,
pub role: u32,
pub token_start: u64,
pub token_count: u64,
pub layer_count: u32,
pub k_type: u32,
pub v_type: u32,
pub k_row_bytes: u32,
pub v_row_bytes: u32,
pub v_element_bytes: u32,
#[serde(default)]
pub k_idx_row_bytes: u32,
pub payload_offset: u64,
pub payload_bytes: u64,
pub flags: u64,
}
impl RuntimeKvPageComponentDesc {
fn as_raw(self) -> RawKvPageComponentDesc {
RawKvPageComponentDesc {
version: self.version,
role: self.role,
token_start: self.token_start,
token_count: self.token_count,
layer_count: self.layer_count,
k_type: self.k_type,
v_type: self.v_type,
k_row_bytes: self.k_row_bytes,
v_row_bytes: self.v_row_bytes,
v_element_bytes: self.v_element_bytes,
k_idx_row_bytes: self.k_idx_row_bytes,
payload_offset: self.payload_offset,
payload_bytes: self.payload_bytes,
flags: self.flags,
}
}
}
impl From<RawKvPageComponentDesc> for RuntimeKvPageComponentDesc {
fn from(raw: RawKvPageComponentDesc) -> Self {
Self {
version: raw.version,
role: raw.role,
token_start: raw.token_start,
token_count: raw.token_count,
layer_count: raw.layer_count,
k_type: raw.k_type,
v_type: raw.v_type,
k_row_bytes: raw.k_row_bytes,
v_row_bytes: raw.v_row_bytes,
v_element_bytes: raw.v_element_bytes,
k_idx_row_bytes: raw.k_idx_row_bytes,
payload_offset: raw.payload_offset,
payload_bytes: raw.payload_bytes,
flags: raw.flags,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct RuntimeKvPageDesc {
pub version: u32,
pub layer_start: i32,
pub layer_end: i32,
pub token_start: u64,
pub token_count: u64,
pub layer_count: u32,
pub k_type: u32,
pub v_type: u32,
pub k_row_bytes: u32,
pub v_row_bytes: u32,
pub v_element_bytes: u32,
#[serde(default)]
pub k_idx_row_bytes: u32,
pub payload_bytes: u64,
pub flags: u64,
#[serde(default)]
pub codec: u32,
#[serde(default)]
pub component_count: u32,
#[serde(default)]
pub components: Box<[RuntimeKvPageComponentDesc; 2]>,
}
impl RuntimeKvPageDesc {
pub fn validate_payload(&self, payload_len: usize) -> Result<()> {
let payload_len = u64::try_from(payload_len)?;
if self.payload_bytes != payload_len {
return Err(anyhow!("KV page payload length does not match descriptor"));
}
let has_k_idx_flag = self.flags & skippy_ffi::KV_PAGE_FLAG_HAS_K_IDX;
if (self.k_idx_row_bytes > 0) != (has_k_idx_flag != 0) {
return Err(anyhow!(
"KV page k_idx flag does not match its indexer row size"
));
}
match self.codec {
0 | skippy_ffi::KV_PAGE_CODEC_SINGLE_V1 if self.component_count == 0 => Ok(()),
skippy_ffi::KV_PAGE_CODEC_ISWA_COMPOSITE_V1
if self.version == 2 && self.component_count == 2 =>
{
let [base, swa] = *self.components;
let base_k_idx_flag = base.flags & skippy_ffi::KV_PAGE_FLAG_HAS_K_IDX;
let swa_k_idx_flag = swa.flags & skippy_ffi::KV_PAGE_FLAG_HAS_K_IDX;
if base.version != 1
|| base.role != 1
|| base.payload_offset != 0
|| base.token_start != self.token_start
|| base.token_count != self.token_count
|| (base.k_idx_row_bytes > 0) != (base_k_idx_flag != 0)
|| swa.version != 1
|| swa.role != 2
|| swa.token_start < self.token_start
|| swa.token_start.checked_add(swa.token_count)
!= self.token_start.checked_add(self.token_count)
|| (swa.k_idx_row_bytes > 0) != (swa_k_idx_flag != 0)
|| swa.payload_offset != base.payload_bytes
|| swa.payload_offset.checked_add(swa.payload_bytes) != Some(payload_len)
{
return Err(anyhow!("invalid composite ISWA KV page components"));
}
Ok(())
}
_ => Err(anyhow!("unsupported KV page codec {}", self.codec)),
}
}
pub(crate) fn as_raw(&self) -> RawKvPageDesc {
RawKvPageDesc {
version: self.version,
layer_start: self.layer_start,
layer_end: self.layer_end,
token_start: self.token_start,
token_count: self.token_count,
layer_count: self.layer_count,
k_type: self.k_type,
v_type: self.v_type,
k_row_bytes: self.k_row_bytes,
v_row_bytes: self.v_row_bytes,
v_element_bytes: self.v_element_bytes,
k_idx_row_bytes: self.k_idx_row_bytes,
payload_bytes: self.payload_bytes,
flags: self.flags,
codec: self.codec,
component_count: self.component_count,
components: self
.components
.as_ref()
.map(RuntimeKvPageComponentDesc::as_raw),
}
}
}
impl From<RawKvPageDesc> for RuntimeKvPageDesc {
fn from(raw: RawKvPageDesc) -> Self {
Self {
version: raw.version,
layer_start: raw.layer_start,
layer_end: raw.layer_end,
token_start: raw.token_start,
token_count: raw.token_count,
layer_count: raw.layer_count,
k_type: raw.k_type,
v_type: raw.v_type,
k_row_bytes: raw.k_row_bytes,
v_row_bytes: raw.v_row_bytes,
v_element_bytes: raw.v_element_bytes,
k_idx_row_bytes: raw.k_idx_row_bytes,
payload_bytes: raw.payload_bytes,
flags: raw.flags,
codec: raw.codec,
component_count: raw.component_count,
components: Box::new(raw.components.map(Into::into)),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RuntimeKvPage {
pub desc: RuntimeKvPageDesc,
pub payload: Vec<u8>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct TokenSignal {
pub entropy: f32,
pub top_logprob: f32,
pub second_logprob: f32,
pub margin: f32,
pub top_token: i32,
pub second_token: i32,
}
impl From<RawTokenSignal> for TokenSignal {
fn from(raw: RawTokenSignal) -> Self {
Self {
entropy: raw.entropy,
top_logprob: raw.top_logprob,
second_logprob: raw.second_logprob,
margin: raw.margin,
top_token: raw.top_token,
second_token: raw.second_token,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct GenerationSignalWindow {
pub token_count: u32,
pub mean_entropy: f32,
pub max_entropy: f32,
pub mean_margin: f32,
pub min_margin: f32,
pub high_entropy_count: u32,
pub repetition_count: u32,
}
impl From<RawGenerationSignalWindow> for GenerationSignalWindow {
fn from(raw: RawGenerationSignalWindow) -> Self {
Self {
token_count: raw.token_count,
mean_entropy: raw.mean_entropy,
max_entropy: raw.max_entropy,
mean_margin: raw.mean_margin,
min_margin: raw.min_margin,
high_entropy_count: raw.high_entropy_count,
repetition_count: raw.repetition_count,
}
}
}
pub struct DecodeFrameBatchOutput {
pub predicted_token: i32,
pub output: ActivationFrame,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct IterationSample {
pub request_index: usize,
pub predicted_token: i32,
}
pub struct IterationBatchOutput {
pub request_outputs: Vec<ActivationFrame>,
pub samples: Vec<IterationSample>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MediaInput {
pub bytes: Vec<u8>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MediaPrefill {
pub token_count: usize,
pub position: u64,
pub first_token: i32,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MediaPrefillChunkFrame {
pub token_count: usize,
pub tokens: Vec<i32>,
pub positions: Vec<i32>,
pub output: ActivationFrame,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MediaPrefillFrame {
pub token_count: usize,
pub position: u64,
pub positions: Vec<i32>,
pub output: ActivationFrame,
pub chunks: Vec<MediaPrefillChunkFrame>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct LogitBias {
pub token_id: i32,
pub bias: f32,
}
#[derive(Debug, Clone, PartialEq)]
pub struct SamplingConfig {
pub enabled: bool,
pub ignore_eos: bool,
pub seed: u32,
pub temperature: f32,
pub top_p: f32,
pub top_k: i32,
pub min_p: f32,
pub presence_penalty: f32,
pub frequency_penalty: f32,
pub repeat_penalty: f32,
pub penalty_last_n: i32,
pub logit_bias: Vec<LogitBias>,
pub typical_p: f32,
pub top_nsigma: f32,
pub dynatemp_range: f32,
pub dynatemp_exponent: f32,
pub dry: DrySamplingConfig,
pub xtc: XtcSamplingConfig,
pub mirostat_mode: i32,
pub mirostat_entropy: f32,
pub mirostat_learning_rate: f32,
pub samplers: Vec<String>,
pub reasoning_budget: ReasoningBudget,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum ReasoningBudget {
#[default]
Unrestricted,
Explicit(u32),
Capped(u32),
Resolved(i32),
}
impl ReasoningBudget {
pub fn resolve_for_output(&mut self, max_output_tokens: usize) {
let resolved = match *self {
Self::Unrestricted => -1,
Self::Explicit(tokens) => i32::try_from(tokens).unwrap_or(i32::MAX),
Self::Capped(tokens) => {
let reserved = max_output_tokens / 2;
i32::try_from((tokens as usize).min(reserved)).unwrap_or(i32::MAX)
}
Self::Resolved(tokens) => tokens,
};
*self = Self::Resolved(resolved);
}
fn native_tokens(self) -> Result<i32> {
match self {
Self::Unrestricted => Ok(-1),
Self::Explicit(tokens) => Ok(i32::try_from(tokens).unwrap_or(i32::MAX)),
Self::Resolved(tokens) => Ok(tokens),
Self::Capped(_) => Err(anyhow!(
"reasoning budget must be resolved against the output limit before sampling"
)),
}
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct DrySamplingConfig {
pub multiplier: f32,
pub base: f32,
pub allowed_length: i32,
pub penalty_last_n: i32,
pub sequence_breakers: Vec<String>,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct XtcSamplingConfig {
pub probability: f32,
pub threshold: f32,
}
impl Default for SamplingConfig {
fn default() -> Self {
Self {
enabled: false,
ignore_eos: false,
seed: 0,
temperature: 1.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
presence_penalty: 0.0,
frequency_penalty: 0.0,
repeat_penalty: 1.0,
penalty_last_n: -1,
logit_bias: Vec::new(),
typical_p: 1.0,
top_nsigma: -1.0,
dynatemp_range: 0.0,
dynatemp_exponent: 1.0,
dry: DrySamplingConfig {
multiplier: 0.0,
base: 1.75,
allowed_length: 2,
penalty_last_n: 64,
sequence_breakers: vec!["\n".into(), ":".into(), "\"".into(), "*".into()],
},
xtc: XtcSamplingConfig {
probability: 0.0,
threshold: 0.1,
},
mirostat_mode: 0,
mirostat_entropy: 5.0,
mirostat_learning_rate: 0.1,
samplers: vec![
"penalties".into(),
"dry".into(),
"top_n_sigma".into(),
"top_k".into(),
"typical_p".into(),
"top_p".into(),
"min_p".into(),
"xtc".into(),
"temperature".into(),
],
reasoning_budget: ReasoningBudget::Unrestricted,
}
}
}
impl SamplingConfig {
pub fn resolve_reasoning_budget(&mut self, max_output_tokens: u32) {
self.reasoning_budget
.resolve_for_output(max_output_tokens as usize);
}
pub(crate) fn as_raw(&self) -> Result<RawSamplingConfig> {
if self.logit_bias.len() > MAX_LOGIT_BIAS {
return Err(anyhow!("sampling logit_bias exceeds the native limit"));
}
if self.samplers.len() > MAX_SAMPLERS {
return Err(anyhow!("sampling sampler order exceeds the native limit"));
}
if self.dry.sequence_breakers.len() > MAX_DRY_SEQUENCE_BREAKERS {
return Err(anyhow!(
"sampling DRY sequence breakers exceed the native count limit"
));
}
if self
.dry
.sequence_breakers
.iter()
.any(|value| value.len() >= MAX_DRY_SEQUENCE_BREAKER_BYTES)
{
return Err(anyhow!(
"sampling DRY sequence breaker exceeds the native byte limit"
));
}
let mut logit_bias = [RawLogitBias {
token_id: 0,
bias: 0.0,
}; MAX_LOGIT_BIAS];
for (target, source) in logit_bias.iter_mut().zip(&self.logit_bias) {
*target = RawLogitBias {
token_id: source.token_id,
bias: source.bias,
};
}
let mut samplers = [0_u32; MAX_SAMPLERS];
for (target, source) in samplers.iter_mut().zip(&self.samplers) {
*target = sampler_id(source)
.ok_or_else(|| anyhow!("sampling sampler order contains {source:?}"))?;
}
let mut dry_sequence_breakers =
[[0_u8; MAX_DRY_SEQUENCE_BREAKER_BYTES]; MAX_DRY_SEQUENCE_BREAKERS];
for (target, source) in dry_sequence_breakers
.iter_mut()
.zip(self.dry.sequence_breakers.iter())
{
let bytes = source.as_bytes();
let length = bytes.len();
target[..length].copy_from_slice(&bytes[..length]);
}
Ok(RawSamplingConfig {
version: 3,
flags: u32::from(self.enabled) | (u32::from(self.ignore_eos) << 1),
seed: self.seed,
top_k: self.top_k,
penalty_last_n: self.penalty_last_n,
temperature: self.temperature,
top_p: self.top_p,
presence_penalty: self.presence_penalty,
frequency_penalty: self.frequency_penalty,
repeat_penalty: self.repeat_penalty,
logit_bias_count: self.logit_bias.len() as u32,
min_p: self.min_p,
typical_p: self.typical_p,
top_nsigma: self.top_nsigma,
dynatemp_range: self.dynatemp_range,
dynatemp_exponent: self.dynatemp_exponent,
dry_multiplier: self.dry.multiplier,
dry_base: self.dry.base,
dry_allowed_length: self.dry.allowed_length,
dry_penalty_last_n: self.dry.penalty_last_n,
xtc_probability: self.xtc.probability,
xtc_threshold: self.xtc.threshold,
mirostat_mode: self.mirostat_mode,
mirostat_entropy: self.mirostat_entropy,
mirostat_learning_rate: self.mirostat_learning_rate,
sampler_count: self.samplers.len() as u32,
samplers,
ignore_eos: u32::from(self.ignore_eos),
dry_sequence_breaker_count: self.dry.sequence_breakers.len() as u32,
dry_sequence_breakers,
logit_bias,
reasoning_budget_tokens: self.reasoning_budget.native_tokens()?,
})
}
}
fn sampler_id(name: &str) -> Option<u32> {
match name {
"penalties" => Some(1),
"dry" => Some(2),
"top_n_sigma" => Some(3),
"top_k" => Some(4),
"typical_p" | "typ_p" => Some(5),
"top_p" => Some(6),
"min_p" => Some(7),
"xtc" => Some(8),
"temperature" | "temp" => Some(9),
_ => None,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChatTemplateMessage {
pub role: String,
pub content: String,
}
impl ChatTemplateMessage {
pub fn new(role: impl Into<String>, content: impl Into<String>) -> Self {
Self {
role: role.into(),
content: content.into(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ChatReasoningFormat {
Auto,
None,
Deepseek,
DeepseekLegacy,
Hidden,
}
impl ChatReasoningFormat {
pub const fn parser_name(self) -> &'static str {
match self {
Self::Auto | Self::Hidden => "auto",
Self::None => "none",
Self::Deepseek => "deepseek",
Self::DeepseekLegacy => "deepseek-legacy",
}
}
pub const fn parses_reasoning(self) -> bool {
!matches!(self, Self::None)
}
pub const fn exposes_reasoning(self) -> bool {
!matches!(self, Self::None | Self::Hidden)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChatTemplateOptions {
pub add_assistant: bool,
pub enable_thinking: Option<bool>,
pub reasoning_format: Option<ChatReasoningFormat>,
pub chat_template_kwargs: Option<String>,
pub chat_template: Option<String>,
pub use_jinja: bool,
pub grammar: Option<String>,
pub json_schema: Option<String>,
pub skip_chat_parsing: bool,
}
impl Default for ChatTemplateOptions {
fn default() -> Self {
Self {
add_assistant: true,
enable_thinking: None,
reasoning_format: None,
chat_template_kwargs: None,
chat_template: None,
use_jinja: true,
grammar: None,
json_schema: None,
skip_chat_parsing: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChatTemplateJsonOptions {
pub add_assistant: bool,
pub enable_thinking: Option<bool>,
pub reasoning_format: Option<ChatReasoningFormat>,
pub chat_template_kwargs: Option<String>,
pub tools_json: Option<String>,
pub tool_choice_json: Option<String>,
pub parallel_tool_calls: bool,
pub chat_template: Option<String>,
pub use_jinja: bool,
pub grammar: Option<String>,
pub json_schema: Option<String>,
pub skip_chat_parsing: bool,
}
impl Default for ChatTemplateJsonOptions {
fn default() -> Self {
Self {
add_assistant: true,
enable_thinking: None,
reasoning_format: None,
chat_template_kwargs: None,
tools_json: None,
tool_choice_json: None,
parallel_tool_calls: true,
chat_template: None,
use_jinja: true,
grammar: None,
json_schema: None,
skip_chat_parsing: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChatTemplateJsonResult {
pub prompt: String,
pub metadata_json: String,
}
#[cfg(test)]
mod activation_boundary_descriptor_tests {
use super::*;
fn f32_boundary(elements_per_token: u64) -> ActivationBoundaryDesc {
let mut parts = [ActivationPartDesc::default(); ACTIVATION_MAX_PARTS];
parts[0] = ActivationPartDesc {
identity: [1; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
ggml_type: crate::GGML_TYPE_F32,
rank: 2,
token_axis: 1,
dimensions: [elements_per_token as i64, -1, 0, 0],
byte_strides: [4, elements_per_token.saturating_mul(4), 0, 0],
..ActivationPartDesc::default()
};
ActivationBoundaryDesc {
version: ACTIVATION_BOUNDARY_DESC_VERSION,
part_count: 1,
frontier_identity: [2; skippy_ffi::ACTIVATION_IDENTITY_BYTES],
parts,
}
}
#[test]
fn raw_f32_boundary_accepts_exact_graph_contract() {
let boundary = f32_boundary(1024);
assert_eq!(boundary.raw_f32_width("output").unwrap(), 1024);
assert_eq!(boundary.payload_bytes("output", 128).unwrap(), 524_288);
}
#[test]
fn raw_f32_boundary_sizes_the_full_graph_observed_width() {
let boundary = f32_boundary(3072);
assert_eq!(boundary.raw_f32_width("input").unwrap(), 3072);
assert_eq!(boundary.payload_bytes("input", 2).unwrap(), 24_576);
}
#[test]
fn raw_f32_boundary_rejects_each_invalid_semantic() {
let mut cases = Vec::new();
let mut unsupported_version = f32_boundary(1024);
unsupported_version.version += 1;
cases.push((unsupported_version, "descriptor version"));
let mut unsupported_type = f32_boundary(1024);
unsupported_type.parts[0].ggml_type = crate::GGML_TYPE_F16;
cases.push((unsupported_type, "not token-indexed F32"));
let mut invalid_axis = f32_boundary(1024);
invalid_axis.parts[0].token_axis = -1;
cases.push((invalid_axis, "not token-indexed F32"));
let mut axis_outside_rank = f32_boundary(1024);
axis_outside_rank.parts[0].token_axis = 2;
cases.push((axis_outside_rank, "invalid rank or token axis"));
let mut zero_rank = f32_boundary(1024);
zero_rank.parts[0].rank = 0;
cases.push((zero_rank, "invalid rank or token axis"));
let mut excessive_rank = f32_boundary(1024);
excessive_rank.parts[0].rank = 5;
cases.push((excessive_rank, "invalid rank or token axis"));
let mut invalid_count = f32_boundary(1024);
invalid_count.part_count = 0;
cases.push((invalid_count, "part count is invalid"));
let too_wide = f32_boundary(i32::MAX as u64 + 1);
cases.push((too_wide, "width exceeds i32"));
for (boundary, expected) in cases {
let error = boundary
.raw_f32_width("input")
.expect_err("invalid graph contract must fail closed");
assert!(
error.to_string().contains(expected),
"expected {expected:?} in {error:#}"
);
}
}
#[test]
fn payload_size_overflow_fails_closed() {
let mut boundary = f32_boundary(1);
boundary.parts[0].dimensions[0] = i64::MAX;
let error = boundary
.payload_bytes("output", u32::MAX)
.expect_err("overflowing payload size must fail");
assert!(error.to_string().contains("overflow"));
}
#[test]
fn payload_size_hint_defers_dynamic_non_token_dimensions() {
let mut boundary = f32_boundary(1024);
boundary.parts[0].rank = 3;
boundary.parts[0].token_axis = 2;
boundary.parts[0].dimensions = [1024, -1, -1, 0];
assert_eq!(boundary.payload_bytes_hint("output", 2).unwrap(), None);
assert!(
boundary
.payload_bytes("output", 2)
.unwrap_err()
.to_string()
.contains("unresolved dynamic dimension")
);
}
}
#[cfg(test)]
mod kv_page_descriptor_tests {
use super::*;
fn composite() -> RuntimeKvPageDesc {
RuntimeKvPageDesc {
version: 2,
payload_bytes: 12,
codec: skippy_ffi::KV_PAGE_CODEC_ISWA_COMPOSITE_V1,
component_count: 2,
components: Box::new([
RuntimeKvPageComponentDesc {
version: 1,
role: 1,
token_start: 0,
token_count: 3,
payload_bytes: 5,
..Default::default()
},
RuntimeKvPageComponentDesc {
version: 1,
role: 2,
token_start: 1,
token_count: 2,
payload_offset: 5,
payload_bytes: 7,
..Default::default()
},
]),
layer_start: 0,
layer_end: 2,
token_start: 0,
token_count: 3,
layer_count: 0,
k_type: 0,
v_type: 0,
k_row_bytes: 0,
v_row_bytes: 0,
v_element_bytes: 0,
k_idx_row_bytes: 0,
flags: 0,
}
}
#[test]
fn accepts_single_v1_and_legacy_codec() {
for codec in [0, skippy_ffi::KV_PAGE_CODEC_SINGLE_V1] {
let desc = RuntimeKvPageDesc {
payload_bytes: 3,
codec,
component_count: 0,
version: 1,
layer_start: 0,
layer_end: 1,
token_start: 0,
token_count: 1,
layer_count: 1,
k_type: 0,
v_type: 0,
k_row_bytes: 0,
v_row_bytes: 0,
v_element_bytes: 0,
k_idx_row_bytes: 0,
flags: 0,
components: Box::default(),
};
assert!(desc.validate_payload(3).is_ok());
}
}
#[test]
fn composite_roundtrips_through_json() {
let desc = composite();
let encoded = serde_json::to_vec(&desc).unwrap();
let decoded: RuntimeKvPageDesc = serde_json::from_slice(&encoded).unwrap();
assert_eq!(decoded, desc);
assert!(decoded.validate_payload(12).is_ok());
}
#[test]
fn rejects_composite_corruption() {
let mut gap = composite();
gap.components[1].payload_offset = 6;
assert!(gap.validate_payload(12).is_err());
let mut wrong_component_version = composite();
wrong_component_version.components[1].version = 2;
assert!(wrong_component_version.validate_payload(12).is_err());
let mut wrong_base_range = composite();
wrong_base_range.components[0].token_count = 2;
assert!(wrong_base_range.validate_payload(12).is_err());
let mut wrong_swa_end = composite();
wrong_swa_end.components[1].token_count = 1;
assert!(wrong_swa_end.validate_payload(12).is_err());
let mut overflowing_swa = composite();
overflowing_swa.components[1].token_start = u64::MAX;
overflowing_swa.components[1].token_count = 2;
assert!(overflowing_swa.validate_payload(12).is_err());
assert!(composite().validate_payload(11).is_err());
}
#[test]
fn sampling_raw_accepts_exact_native_dry_breaker_payload() {
let sampling = SamplingConfig {
dry: DrySamplingConfig {
sequence_breakers: vec!["123456789012345".to_string()],
..DrySamplingConfig::default()
},
..SamplingConfig::default()
};
assert!(sampling.as_raw().is_ok());
}
#[test]
fn sampling_raw_rejects_dry_breakers_that_cannot_fit_native_payload() {
for breaker in ["1234567890123456", "🐍🐍🐍🐍"] {
let sampling = SamplingConfig {
dry: DrySamplingConfig {
sequence_breakers: vec![breaker.to_string()],
..DrySamplingConfig::default()
},
..SamplingConfig::default()
};
assert!(sampling.as_raw().is_err(), "breaker={breaker:?}");
}
}
#[test]
fn sampling_raw_rejects_unknown_sampler_names() {
let sampling = SamplingConfig {
samplers: vec!["unknown".to_string()],
..SamplingConfig::default()
};
assert!(sampling.as_raw().is_err());
}
#[test]
fn reasoning_budget_resolution_distinguishes_explicit_levels_and_unrestricted() {
let mut explicit = ReasoningBudget::Explicit(8_192);
explicit.resolve_for_output(4_096);
assert_eq!(explicit, ReasoningBudget::Resolved(8_192));
let mut semantic = ReasoningBudget::Capped(8_192);
semantic.resolve_for_output(4_096);
assert_eq!(semantic, ReasoningBudget::Resolved(2_048));
let mut small_fallback = ReasoningBudget::Capped(4_096);
small_fallback.resolve_for_output(31);
assert_eq!(small_fallback, ReasoningBudget::Resolved(15));
let mut disabled = ReasoningBudget::Explicit(0);
disabled.resolve_for_output(8_192);
assert_eq!(disabled, ReasoningBudget::Resolved(0));
let mut unrestricted = ReasoningBudget::Unrestricted;
unrestricted.resolve_for_output(8_192);
assert_eq!(unrestricted, ReasoningBudget::Resolved(-1));
}
#[test]
fn resolved_reasoning_budget_reaches_native_sampling_abi() {
for (budget, expected) in [
(ReasoningBudget::Resolved(-1), -1),
(ReasoningBudget::Resolved(0), 0),
(ReasoningBudget::Resolved(1_024), 1_024),
] {
let raw = SamplingConfig {
reasoning_budget: budget,
..SamplingConfig::default()
}
.as_raw()
.unwrap();
assert_eq!(raw.reasoning_budget_tokens, expected);
}
}
}