use axum::response::sse::Event;
use serde_json::{json, Value};
use crate::generate::{FinishReason, GenerationParams};
pub(super) fn stop_type(finish: &FinishReason) -> &'static str {
match finish {
FinishReason::Stop => "eos",
FinishReason::StopSequence(_) => "word",
FinishReason::Length => "limit",
FinishReason::Cancelled => "cancelled",
}
}
pub(super) fn stopping_word(finish: &FinishReason) -> &str {
match finish {
FinishReason::StopSequence(word) => word,
_ => "",
}
}
pub(super) fn timings(usage: &ferrox_api::Usage) -> Value {
let per_token = |ms: Option<f64>, n: usize| match (ms, n) {
(Some(ms), n) if n > 0 => json!(ms / n as f64),
_ => Value::Null,
};
json!({
"cache_n": usage.cached_tokens.map(|n| n as i64).unwrap_or(-1),
"prompt_n": usage.prompt_tokens,
"prompt_ms": usage.prompt_eval_duration_ms,
"prompt_per_token_ms": per_token(usage.prompt_eval_duration_ms, usage.prompt_tokens),
"prompt_per_second": usage.prompt_per_second,
"predicted_n": usage.completion_tokens,
"predicted_ms": usage.generation_duration_ms,
"predicted_per_token_ms":
per_token(usage.generation_duration_ms, usage.completion_tokens),
"predicted_per_second": usage.predicted_per_second,
})
}
pub(super) fn generation_settings(params: &GenerationParams, model: &str) -> Value {
let s = ¶ms.sampling;
json!({
"model": model,
"n_predict": params.max_tokens,
"seed": params.seed,
"temperature": s.temperature,
"top_p": s.top_p,
"min_p": s.min_p,
"top_k": s.top_k,
"repeat_penalty": s.repetition_penalty,
"repeat_last_n": s.penalty_last_n,
"presence_penalty": s.presence_penalty,
"frequency_penalty": s.frequency_penalty,
"stop": params.stop,
"ignore_eos": params.ignore_eos,
"grammar": params.grammar.is_some(),
})
}
pub(super) fn final_body(
content: &str,
finish: &FinishReason,
usage: &ferrox_api::Usage,
params: &GenerationParams,
model: &str,
prompt: &str,
) -> Value {
json!({
"index": 0,
"content": content,
"tokens": Vec::<u32>::new(),
"id_slot": -1,
"stop": true,
"model": model,
"tokens_predicted": usage.completion_tokens,
"tokens_evaluated": usage.prompt_tokens,
"generation_settings": generation_settings(params, model),
"prompt": prompt,
"has_new_line": content.contains('\n'),
"truncated": false,
"stop_type": stop_type(finish),
"stopping_word": stopping_word(finish),
"tokens_cached": usage.cached_tokens.unwrap_or(0),
"timings": timings(usage),
})
}
pub(super) fn partial_body(content: &str) -> Value {
json!({
"index": 0,
"content": content,
"tokens": Vec::<u32>::new(),
"stop": false,
"id_slot": -1,
})
}
pub(super) fn frame(body: &Value) -> Event {
Event::default().data(body.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stop_type_speaks_llama_cpps_vocabulary() {
assert_eq!(stop_type(&FinishReason::Stop), "eos");
assert_eq!(stop_type(&FinishReason::Length), "limit");
assert_eq!(stop_type(&FinishReason::StopSequence("END".into())), "word");
assert_eq!(stop_type(&FinishReason::Cancelled), "cancelled");
assert_eq!(
stopping_word(&FinishReason::StopSequence("END".into())),
"END"
);
assert_eq!(stopping_word(&FinishReason::Stop), "");
}
#[test]
fn an_unmeasured_timing_is_null_rather_than_zero() {
let mut usage = ferrox_api::Usage::new(10, 5);
let untimed = timings(&usage);
assert!(untimed["prompt_ms"].is_null());
assert!(untimed["prompt_per_token_ms"].is_null());
assert!(untimed["predicted_per_second"].is_null());
assert_eq!(untimed["cache_n"], -1);
assert_eq!(untimed["prompt_n"], 10);
assert_eq!(untimed["predicted_n"], 5);
usage.prompt_eval_duration_ms = Some(100.0);
usage.generation_duration_ms = Some(50.0);
usage.cached_tokens = Some(0);
let timed = timings(&usage);
assert_eq!(timed["prompt_per_token_ms"], 10.0);
assert_eq!(timed["predicted_per_token_ms"], 10.0);
assert_eq!(timed["cache_n"], 0);
}
#[test]
fn a_partial_frame_is_the_documented_shape() {
let body = partial_body("tok");
assert_eq!(body["content"], "tok");
assert_eq!(body["stop"], false);
assert!(body["tokens"].as_array().unwrap().is_empty());
assert!(body.get("timings").is_none());
assert!(body.get("generation_settings").is_none());
}
}