mod file;
mod identity;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Instant;
use axum::extract::{Path as AxumPath, Query, State};
use axum::http::StatusCode;
use axum::Json;
use serde::Deserialize;
use ferrox_core::cache::KvCache;
use ferrox_core::kv_signature::KvDtype;
use ferrox_models::Decoder;
use crate::{ActiveModel, ApiError, AppState};
use file::SlotPayload;
use identity::SlotIdentity;
fn slot_dir() -> Option<PathBuf> {
std::env::var("FERROX_SLOT_SAVE_PATH")
.ok()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.map(PathBuf::from)
}
#[derive(Debug, Deserialize)]
pub(crate) struct SlotQuery {
#[serde(default)]
action: String,
}
#[derive(Debug, Deserialize)]
struct SlotBody {
filename: Option<String>,
prompt: Option<String>,
}
fn error(status: StatusCode, message: String, kind: &str) -> ApiError {
(
status,
Json(serde_json::json!({
"error": {"message": message, "type": kind, "code": status.as_u16()}
})),
)
}
fn erase_refusal() -> ApiError {
error(
StatusCode::NOT_IMPLEMENTED,
"action=erase: ferrox holds slot KV in one shared prefix cache, which has no per-entry \
eviction, so erasing one slot would mean clearing all of them. Restart the server, or \
delete the slot file, to stop a slot being restorable"
.to_string(),
"unsupported_feature",
)
}
fn validate_filename(name: &str) -> Result<(), String> {
if name.is_empty() {
return Err("filename is empty".to_string());
}
if name.len() > 255 {
return Err(format!(
"filename is {} bytes; a slot name may be at most 255",
name.len()
));
}
if name.starts_with('.') {
return Err(
"filename starts with '.'; a slot name may not be a dotfile or a \
relative-path component"
.to_string(),
);
}
if let Some(bad) = name
.chars()
.find(|c| !(c.is_ascii_alphanumeric() || *c == '.' || *c == '_' || *c == '-'))
{
return Err(format!(
"filename contains {bad:?}; a slot name may use only ASCII letters, digits, '.', \
'_' and '-'"
));
}
Ok(())
}
fn serving_identity(active: &ActiveModel) -> Result<(SlotIdentity, Arc<Decoder>), ApiError> {
let model = active.generative()?;
let Some(decoder) = model.gguf_decoder() else {
return Err(error(
StatusCode::NOT_IMPLEMENTED,
"slots are implemented for the generic GGUF decoder only; this server is serving a \
dedicated engine whose KV layout this format does not describe"
.to_string(),
"unsupported_feature",
));
};
let Some(path) = active.checkpoint_path.as_ref() else {
return Err(error(
StatusCode::NOT_IMPLEMENTED,
identity::FingerprintError::NoCheckpoint.to_string(),
"unsupported_feature",
));
};
let fingerprint = identity::fingerprint_gguf(path).map_err(|e| {
error(
StatusCode::INTERNAL_SERVER_ERROR,
e.to_string(),
"checkpoint_unreadable",
)
})?;
let config = &decoder.config;
Ok((
SlotIdentity {
model_name: config.name.to_string(),
n_layers: decoder.layers.len(),
n_kv_heads: config.n_kv_heads,
head_dim: config.head_dim,
dtype: KvDtype::F32,
fingerprint,
},
Arc::clone(decoder),
))
}
pub(crate) async fn post_slot(
State(state): State<Arc<AppState>>,
AxumPath(id_slot): AxumPath<String>,
Query(query): Query<SlotQuery>,
body: String,
) -> Result<Json<serde_json::Value>, ApiError> {
let Some(dir) = slot_dir() else {
return Err(error(
StatusCode::NOT_IMPLEMENTED,
"This server does not support slots action. Start it with `--slot-save-path`"
.to_string(),
"unsupported_feature",
));
};
let Ok(id_slot) = id_slot.parse::<u32>() else {
return Err(error(
StatusCode::BAD_REQUEST,
"Invalid slot ID".to_string(),
"invalid_request_error",
));
};
match query.action.as_str() {
"save" => save(state, dir, id_slot, body).await,
"restore" => restore(state, dir, id_slot, body).await,
"erase" => Err(erase_refusal()),
other => Err(error(
StatusCode::BAD_REQUEST,
format!("Invalid action {other:?}: expected save, restore or erase"),
"invalid_request_error",
)),
}
}
fn parse_body(body: &str) -> Result<SlotBody, ApiError> {
serde_json::from_str::<SlotBody>(body).map_err(|e| {
error(
StatusCode::BAD_REQUEST,
format!("slot request body is not JSON: {e}"),
"invalid_request_error",
)
})
}
fn slot_path(dir: &std::path::Path, body: &SlotBody) -> Result<PathBuf, ApiError> {
let Some(filename) = body.filename.as_deref() else {
return Err(error(
StatusCode::BAD_REQUEST,
"slot request needs a \"filename\"".to_string(),
"invalid_request_error",
));
};
validate_filename(filename).map_err(|why| {
error(
StatusCode::BAD_REQUEST,
format!("Invalid filename: {why}"),
"invalid_request_error",
)
})?;
Ok(dir.join(filename))
}
fn require_prefix_cache(
state: &AppState,
) -> Result<Arc<std::sync::Mutex<ferrox_models::PrefixCache>>, ApiError> {
state.prefix_cache.clone().ok_or_else(|| {
error(
StatusCode::NOT_IMPLEMENTED,
"slots store and restore into this server's prefix cache, which is off: set \
FERROX_PREFIX_CACHE_ENTRIES to a positive number of entries"
.to_string(),
"unsupported_feature",
)
})
}
async fn save(
state: Arc<AppState>,
dir: PathBuf,
id_slot: u32,
body: String,
) -> Result<Json<serde_json::Value>, ApiError> {
let started = Instant::now();
let body = parse_body(&body)?;
let path = slot_path(&dir, &body)?;
let Some(prompt) = body.prompt.as_deref() else {
return Err(error(
StatusCode::BAD_REQUEST,
"slot save needs a \"prompt\": ferrox has no per-slot KV region holding one \
already, so the save names the prefix to warm"
.to_string(),
"invalid_request_error",
));
};
let active = state.require_active()?;
let cache = require_prefix_cache(&state)?;
let (identity, decoder) = serving_identity(&active)?;
let mut tokens = active.encode_any(prompt, ferrox_models::tokenizer::SpecialTokens::Parse);
ferrox_models::tokenizer::prepend_bos(
&mut tokens,
active.generative_opt().and_then(|m| m.bos_id()),
);
if tokens.is_empty() {
return Err(error(
StatusCode::BAD_REQUEST,
"slot save prompt encodes to no tokens; there is no KV state to save".to_string(),
"invalid_request_error",
));
}
let n_tokens = tokens.len();
let write = tokio::task::spawn_blocking(move || {
let payload = prefill_slot(&decoder, tokens);
let bytes = file::encode(&identity, &payload);
let written = publish(&path, &bytes)?;
cache.lock().unwrap_or_else(|p| p.into_inner()).store(
payload.tokens,
payload.layers,
payload.pending_logits,
);
Ok::<u64, std::io::Error>(written)
})
.await;
match write {
Ok(Ok(n_written)) => Ok(Json(serde_json::json!({
"id_slot": id_slot,
"filename": body.filename,
"n_saved": n_tokens,
"n_written": n_written,
"timings": {"save_ms": started.elapsed().as_secs_f64() * 1000.0},
}))),
Ok(Err(e)) => Err(error(
StatusCode::INTERNAL_SERVER_ERROR,
format!("writing the slot file: {e}"),
"slot_write_failed",
)),
Err(e) => Err(error(
StatusCode::INTERNAL_SERVER_ERROR,
format!("the slot save task did not finish: {e}"),
"slot_write_failed",
)),
}
}
async fn restore(
state: Arc<AppState>,
dir: PathBuf,
id_slot: u32,
body: String,
) -> Result<Json<serde_json::Value>, ApiError> {
let started = Instant::now();
let body = parse_body(&body)?;
let path = slot_path(&dir, &body)?;
let active = state.require_active()?;
let cache = require_prefix_cache(&state)?;
let (identity, _decoder) = serving_identity(&active)?;
let vocab = active.generative_opt().and_then(|m| m.vocab_size());
let bytes = match std::fs::read(&path) {
Ok(bytes) => bytes,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
return Err(error(
StatusCode::NOT_FOUND,
format!("no slot file {}", path.display()),
"not_found",
))
}
Err(e) => {
return Err(error(
StatusCode::INTERNAL_SERVER_ERROR,
format!("reading the slot file: {e}"),
"slot_read_failed",
))
}
};
let n_read = bytes.len();
let unverified = match file::decode(&bytes) {
Ok(slot) => slot,
Err(e) => {
return Err(error(
StatusCode::BAD_REQUEST,
format!("{}: {e}", path.display()),
"invalid_slot_file",
))
}
};
tracing::debug!(
"restoring {}: saved under model {} with {} layers, checkpoint {}",
path.display(),
unverified.identity().model_name,
unverified.identity().n_layers,
unverified.identity().fingerprint,
);
let payload = match unverified.verify(&identity) {
Ok(payload) => payload,
Err(mismatch) => {
return Err(error(
StatusCode::BAD_REQUEST,
format!(
"refusing to restore {}: {mismatch}. Restoring attention state computed by \
other weights does not fail, it answers wrongly",
path.display()
),
"slot_model_mismatch",
))
}
};
if let Some(vocab) = vocab {
if let Some(&bad) = payload.tokens.iter().find(|&&t| t >= vocab) {
return Err(error(
StatusCode::BAD_REQUEST,
format!(
"refusing to restore {}: it names token id {bad}, and this model's \
vocabulary has {vocab} tokens",
path.display()
),
"slot_model_mismatch",
));
}
}
let n_restored = payload.tokens.len();
cache.lock().unwrap_or_else(|p| p.into_inner()).store(
payload.tokens,
payload.layers,
payload.pending_logits,
);
Ok(Json(serde_json::json!({
"id_slot": id_slot,
"filename": body.filename,
"n_restored": n_restored,
"n_read": n_read,
"timings": {"restore_ms": started.elapsed().as_secs_f64() * 1000.0},
})))
}
fn prefill_slot(decoder: &Decoder, tokens: Vec<usize>) -> SlotPayload {
let mut layers: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let pending_logits =
crate::generate::forward_prompt_batch(decoder, &tokens, 0, &mut layers, true);
#[cfg(feature = "metal")]
decoder.sync_metal_attn_kv_to_host(&mut layers);
SlotPayload {
tokens,
pending_logits,
layers,
}
}
fn publish(path: &std::path::Path, bytes: &[u8]) -> std::io::Result<u64> {
use std::io::Write;
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let tmp = path.with_extension(format!("{}.tmp", std::process::id()));
{
let mut file = std::fs::File::create(&tmp)?;
file.write_all(bytes)?;
file.sync_all()?;
}
std::fs::rename(&tmp, path)?;
Ok(bytes.len() as u64)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_filename_that_could_leave_the_slot_directory_is_refused() {
for bad in [
"../escape",
"..",
".",
"sub/dir",
"back\\slash",
"nul\0byte",
".hidden",
"",
] {
assert!(
validate_filename(bad).is_err(),
"{bad:?} should not be a slot name"
);
}
}
#[test]
fn an_ordinary_slot_name_is_accepted() {
for good in ["system.fslot", "sys-prompt_v2.fslot", "a", "0"] {
assert!(validate_filename(good).is_ok(), "{good:?}");
}
}
#[test]
fn an_over_long_filename_is_refused_rather_than_truncated() {
let name = "a".repeat(256);
assert!(validate_filename(&name).is_err());
assert!(validate_filename(&name[..255]).is_ok());
}
}