pub(crate) mod auth;
pub(crate) mod constrained;
pub(crate) mod lanes { pub use memra_lanes::*; }
mod toolcall;
mod worker;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::mpsc::Sender;
use axum::{
Json, Router,
extract::State,
http::StatusCode,
response::{sse::{Event as SseEvent, Sse}, IntoResponse, Response},
routing::{get, post},
};
use serde::{Deserialize, Serialize};
use serde_json::json;
use memra_engine::decode::GenParams;
use memra_engine::sampler::SamplerConfig;
use memra_tokenizer::chat::{ThinkMode, ToolCall as TmplToolCall, Turn as TmplTurn};
use toolcall::{ParsedToolCall, Piece, ToolStreamParser};
use worker::{Cmd, Event, ModelCaps, Request, SharedMetrics};
#[derive(Clone)]
struct AppState {
cmd_tx: Sender<Cmd>,
models: Arc<Vec<String>>,
caps: Arc<HashMap<String, ModelCaps>>,
metrics: SharedMetrics,
started: u64,
inflight: InflightCounts,
tenant_inflight: TenantGauge,
}
type InflightCounts = Arc<[std::sync::atomic::AtomicUsize; 3]>;
type TenantGauge = Arc<std::sync::Mutex<HashMap<String, usize>>>;
struct InflightGuard {
counts: InflightCounts,
idx: usize,
tenants: TenantGauge,
tenant: String,
}
impl InflightGuard {
fn acquire(counts: InflightCounts, lane: lanes::Lane, tenants: TenantGauge,
tenant: &str) -> (Self, usize, usize) {
let idx = lane.idx();
let n = counts[idx].fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
let nt = {
let mut m = tenants.lock().unwrap();
let e = m.entry(tenant.to_string()).or_insert(0);
*e += 1;
*e
};
(InflightGuard { counts, idx, tenants, tenant: tenant.to_string() }, n, nt)
}
}
impl Drop for InflightGuard {
fn drop(&mut self) {
self.counts[self.idx].fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
let mut m = self.tenants.lock().unwrap();
if let Some(e) = m.get_mut(&self.tenant) {
*e -= 1;
if *e == 0 {
m.remove(&self.tenant);
}
}
}
}
fn lane_cap(lane: lanes::Lane) -> usize {
static CAPS: std::sync::OnceLock<[usize; 3]> = std::sync::OnceLock::new();
CAPS.get_or_init(|| {
let batching = std::env::var("MEMRA_SERVE_BATCH").map(|v| v != "0").unwrap_or(true);
let interactive = if batching {
std::env::var("MEMRA_MAX_SESSIONS").ok()
.and_then(|v| v.parse().ok()).unwrap_or(64)
} else {
worker::MAX_ACTIVE
};
let p = lanes::LanePolicy::from_env();
[interactive, p.max_sessions[1], p.max_sessions[2]]
})[lane.idx()]
}
fn reset_estimate_s(m: &worker::Metrics) -> u64 {
if m.completed > 0 && m.step_p50_ms > 0.0 {
let mean_toks = m.tokens_out as f64 / m.completed as f64;
return ((mean_toks * m.step_p50_ms as f64 / 1000.0).ceil() as u64).clamp(1, 600);
}
static D: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
*D.get_or_init(|| std::env::var("MEMRA_RL_RESET_S").ok()
.and_then(|v| v.parse().ok()).unwrap_or(2))
}
static DRAINING: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
fn draining() -> bool {
DRAINING.load(std::sync::atomic::Ordering::SeqCst)
}
fn drain_deadline_s() -> u64 {
static D: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
*D.get_or_init(|| std::env::var("MEMRA_DRAIN_S").ok()
.and_then(|v| v.parse().ok()).unwrap_or(30))
}
fn drain_response() -> Response {
let mut resp = error_response(
StatusCode::SERVICE_UNAVAILABLE,
"server is draining (shutdown in progress); retry",
"server_error", None);
if let Ok(v) = axum::http::HeaderValue::from_str(&drain_deadline_s().to_string()) {
resp.headers_mut().insert(axum::http::header::RETRY_AFTER, v);
}
resp
}
struct RateLimit {
limit: usize,
remaining: usize,
reset_s: u64,
}
impl RateLimit {
fn at_admit(lane: lanes::Lane, n_inflight: usize, metrics: &SharedMetrics,
tenant: &auth::TenantCtx, n_tenant: usize) -> Self {
let global = lane_cap(lane);
let Some(t) = tenant.rate_limit.filter(|&t| t < global) else {
return Self::compute(global, n_inflight, metrics);
};
let headroom = t.saturating_sub(n_tenant)
.min(global.saturating_sub(n_inflight));
Self::compute(t, t - headroom, metrics)
}
fn compute(limit: usize, n_inflight: usize, metrics: &SharedMetrics) -> Self {
let remaining = limit.saturating_sub(n_inflight);
let reset_s = if remaining > 0 {
0
} else {
let m = metrics.lock().map(|m| m.clone()).unwrap_or_default();
reset_estimate_s(&m)
};
RateLimit { limit, remaining, reset_s }
}
fn attach(&self, mut resp: Response) -> Response {
let h = resp.headers_mut();
for (k, v) in [
("x-ratelimit-limit", self.limit as u64),
("x-ratelimit-remaining", self.remaining as u64),
("x-ratelimit-reset", self.reset_s),
] {
if let Ok(v) = axum::http::HeaderValue::from_str(&v.to_string()) {
h.insert(axum::http::HeaderName::from_static(k), v);
}
}
resp
}
}
#[derive(Deserialize)]
struct CompletionReq {
model: String,
#[serde(default)]
prompt: String,
#[serde(default)]
prompt_ids: Vec<u32>,
#[serde(default)]
max_tokens: Option<usize>,
#[serde(default = "default_temperature")]
temperature: f32,
#[serde(default = "one")]
top_p: f32,
#[serde(default)]
top_k: usize,
#[serde(default)]
min_p: f32,
#[serde(default)]
frequency_penalty: f32,
#[serde(default)]
presence_penalty: f32,
#[serde(default = "one")]
repetition_penalty: f32,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
stop: StopSequences,
#[serde(default)]
logit_bias: Option<serde_json::Value>,
#[serde(default)]
logprobs: Option<serde_json::Value>,
#[serde(default)]
n: Option<usize>,
#[serde(default)]
best_of: Option<usize>,
#[serde(default)]
chat: bool,
#[serde(default)]
stream: bool,
#[serde(default)]
max_ctx: Option<usize>,
#[serde(default)]
trace_id: Option<String>,
#[serde(default)]
cache_salt: Option<String>,
#[serde(default)]
session_id: Option<String>,
#[serde(default)]
user: Option<String>,
}
#[derive(Deserialize)]
struct ChatMessage {
role: String,
#[serde(default)]
content: serde_json::Value,
#[serde(default)]
tool_calls: Vec<ReqToolCall>,
#[serde(default)]
#[allow(dead_code)]
tool_call_id: Option<String>,
}
#[derive(Deserialize)]
struct ReqToolCall {
#[serde(default)]
#[allow(dead_code)]
id: Option<String>,
function: ReqToolFunction,
}
#[derive(Deserialize)]
struct ReqToolFunction {
name: String,
#[serde(default)]
arguments: serde_json::Value,
}
#[derive(Clone, Default, Deserialize)]
#[serde(untagged)]
enum StopSequences {
One(String),
Many(Vec<String>),
#[default]
None,
}
impl StopSequences {
fn into_vec(self) -> Vec<String> {
match self {
Self::One(stop) => vec![stop],
Self::Many(stops) => stops,
Self::None => Vec::new(),
}
}
}
#[derive(Deserialize)]
struct ChatCompletionReq {
model: String,
messages: Vec<ChatMessage>,
#[serde(default, alias = "max_completion_tokens")]
max_tokens: Option<usize>,
#[serde(default = "default_temperature")]
temperature: f32,
#[serde(default = "one")]
top_p: f32,
#[serde(default)]
top_k: usize,
#[serde(default)]
min_p: f32,
#[serde(default)]
frequency_penalty: f32,
#[serde(default)]
presence_penalty: f32,
#[serde(default = "one")]
repetition_penalty: f32,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
stop: StopSequences,
#[serde(default)]
stream: bool,
#[serde(default)]
max_ctx: Option<usize>,
#[serde(default)]
response_format: Option<serde_json::Value>,
#[serde(default)]
logit_bias: Option<serde_json::Value>,
#[serde(default)]
logprobs: Option<serde_json::Value>,
#[serde(default)]
top_logprobs: Option<usize>,
#[serde(default)]
n: Option<usize>,
#[serde(default)]
tools: Vec<serde_json::Value>,
#[serde(default)]
tool_choice: Option<serde_json::Value>,
#[serde(default)]
reasoning_effort: Option<String>,
#[serde(default)]
reasoning: Option<serde_json::Value>,
#[serde(default)]
include_reasoning: Option<bool>,
#[serde(default)]
cache_salt: Option<String>,
#[serde(default)]
session_id: Option<String>,
#[serde(default)]
user: Option<String>,
}
fn one() -> f32 { 1.0 }
fn default_temperature() -> f32 { 1.0 }
#[derive(Serialize)]
struct CompletionResp {
model: String,
text: String,
tokens: Vec<u32>,
stop_reason: String,
n_tokens: usize,
prompt_tokens: usize,
cached_tokens: usize,
elapsed_s: f64,
}
fn usage_json(n_prompt: usize, n_tokens: usize, n_cached: usize, elapsed_s: f64,
spec: Option<worker::SpecUsage>) -> serde_json::Value {
let mut u = json!({
"prompt_tokens": n_prompt,
"completion_tokens": n_tokens,
"total_tokens": n_prompt + n_tokens,
"prompt_tokens_details": { "cached_tokens": n_cached },
"elapsed_s": elapsed_s,
});
if let Some(sp) = spec {
u["spec"] = json!({
"rounds": sp.rounds,
"drafted": sp.drafted,
"accepted": sp.accepted,
"acceptance_rate": if sp.drafted > 0 {
sp.accepted as f64 / sp.drafted as f64 } else { 0.0 },
});
}
u
}
const SYSTEM_FINGERPRINT: &str = concat!("memra-", env!("MEMRA_BUILD_SHA"));
fn gen_hex128() -> String {
use std::hash::{BuildHasher, Hasher};
static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let n = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let t = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let mut h1 = std::collections::hash_map::RandomState::new().build_hasher();
h1.write_u64(n);
h1.write_u64(t);
let mut h2 = std::collections::hash_map::RandomState::new().build_hasher();
h2.write_u64(t.rotate_left(17));
h2.write_u64(n);
format!("{:016x}{:016x}", h1.finish(), h2.finish())
}
#[derive(Clone)]
struct Envelope {
id: String,
created: u64,
}
impl Envelope {
fn new(chat: bool) -> Self {
Envelope {
id: format!("{}-{}", if chat { "chatcmpl" } else { "cmpl" }, gen_hex128()),
created: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0),
}
}
fn stamp(&self, mut v: serde_json::Value) -> serde_json::Value {
v["id"] = json!(self.id);
v["created"] = json!(self.created);
v["system_fingerprint"] = json!(SYSTEM_FINGERPRINT);
v
}
}
fn with_request_id(id: &str, mut resp: Response) -> Response {
if let Ok(v) = axum::http::HeaderValue::from_str(id) {
resp.headers_mut()
.insert(axum::http::HeaderName::from_static("x-request-id"), v);
}
resp
}
fn openai_compat() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| {
match std::env::var("MEMRA_COMPAT").as_deref() {
Ok("openai") => true,
Ok(_) => false,
Err(_) => std::env::var("MEMRA_API_KEY").is_ok(),
}
})
}
fn cache_namespace(cache_salt: &Option<String>) -> String {
cache_salt.clone().unwrap_or_default()
}
fn affinity_key(
session_id: &Option<String>,
user: &Option<String>,
headers: &axum::http::HeaderMap,
) -> Option<String> {
let clean = |s: &str| -> Option<String> {
let t = s.trim();
if t.is_empty() { None } else { Some(t.to_string()) }
};
session_id.as_deref().and_then(clean)
.or_else(|| user.as_deref().and_then(clean))
.or_else(|| headers.get("x-session-id")
.and_then(|v| v.to_str().ok())
.and_then(clean))
}
fn error_body(message: &str, etype: &str, param: Option<&str>, code: Option<&str>)
-> serde_json::Value {
json!({ "error": {
"message": message,
"type": etype,
"param": param,
"code": code,
} })
}
fn error_response(status: StatusCode, message: &str, etype: &str, param: Option<&str>)
-> Response {
(status, Json(error_body(message, etype, param, None))).into_response()
}
fn bad_request(message: &str, param: Option<&str>) -> Response {
error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", param)
}
fn stop_reason_to_finish(r: &str) -> &'static str {
match r {
"Eos" | "Callback" => "stop",
"MaxNew" | "ContextFull" => "length",
_ => "stop",
}
}
fn content_to_text(v: &serde_json::Value) -> Result<String, String> {
match v {
serde_json::Value::Null => Ok(String::new()),
serde_json::Value::String(s) => Ok(s.clone()),
serde_json::Value::Array(parts) => {
let mut out = String::new();
for p in parts {
match p.get("type").and_then(|t| t.as_str()) {
Some("text") | None => match p.get("text").and_then(|t| t.as_str()) {
Some(t) => out.push_str(t),
None => return Err("content part has no text field".into()),
},
Some(other) => {
return Err(format!("unsupported content part type {other:?} (text only)"));
}
}
}
Ok(out)
}
_ => Err("content must be a string, null, or an array of text parts".into()),
}
}
fn pyjson(v: &serde_json::Value, out: &mut String) {
match v {
serde_json::Value::Object(m) => {
out.push('{');
for (i, (k, val)) in m.iter().enumerate() {
if i > 0 { out.push_str(", "); }
out.push_str(&serde_json::Value::String(k.clone()).to_string());
out.push_str(": ");
pyjson(val, out);
}
out.push('}');
}
serde_json::Value::Array(a) => {
out.push('[');
for (i, val) in a.iter().enumerate() {
if i > 0 { out.push_str(", "); }
pyjson(val, out);
}
out.push(']');
}
scalar => out.push_str(&scalar.to_string()),
}
}
fn pyjson_str(v: &serde_json::Value) -> String {
let mut s = String::new();
pyjson(v, &mut s);
s
}
fn sampler_config(temperature: f32, top_k: usize, top_p: f32, min_p: f32,
frequency_penalty: f32, presence_penalty: f32, repetition_penalty: f32,
seed: Option<u64>) -> SamplerConfig {
let penalties_on = frequency_penalty != 0.0 || presence_penalty != 0.0
|| repetition_penalty != 1.0;
SamplerConfig {
temperature,
top_k,
top_p,
min_p,
penalty_last_n: if penalties_on { usize::MAX } else { 0 },
penalty_repeat: repetition_penalty,
penalty_freq: frequency_penalty,
penalty_present: presence_penalty,
seed: seed.unwrap_or_else(fresh_seed),
}
}
fn fresh_seed() -> u64 {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let n = COUNTER.fetch_add(1, Ordering::Relaxed);
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let mut z = nanos
.wrapping_add(n.wrapping_mul(0x9E3779B97F4A7C15))
.wrapping_add(0x9E3779B97F4A7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^= z >> 31;
if z == 0 { 0x9E3779B97F4A7C15 } else { z }
}
fn reject_unsupported(fields: &[(&str, bool, &str)]) -> Result<(), (String, String)> {
for (param, present, why) in fields {
if *present {
return Err((format!("{param} is not supported{why}"), param.to_string()));
}
}
Ok(())
}
#[derive(PartialEq)]
enum ToolChoice { Auto, None }
fn parse_tool_choice(v: &Option<serde_json::Value>) -> Result<ToolChoice, String> {
match v {
None | Some(serde_json::Value::Null) => Ok(ToolChoice::Auto),
Some(serde_json::Value::String(s)) => match s.as_str() {
"auto" => Ok(ToolChoice::Auto),
"none" => Ok(ToolChoice::None),
"required" => Err("tool_choice \"required\" is not supported (no constrained \
decoding); use \"auto\"".into()),
other => Err(format!("bad tool_choice {other:?} (auto|none)")),
},
Some(serde_json::Value::Object(_)) =>
Err("named-function tool_choice is not supported; use \"auto\"".into()),
Some(other) => Err(format!("bad tool_choice: {other}")),
}
}
fn parse_think(reasoning_effort: &Option<String>, reasoning: &Option<serde_json::Value>)
-> Result<ThinkMode, String> {
let mut effort = reasoning_effort.clone();
if let Some(r) = reasoning {
match r {
serde_json::Value::Null => {}
serde_json::Value::Object(obj) => {
if obj.get("enabled").and_then(|v| v.as_bool()) == Some(false) {
return Ok(ThinkMode::NoThink);
}
if let Some(e) = obj.get("effort").and_then(|v| v.as_str()) {
effort = Some(e.to_string());
}
}
_ => return Err("reasoning must be an object".into()),
}
}
match effort.as_deref() {
None => Ok(ThinkMode::Default),
Some("none") | Some("minimal") | Some("low") => Ok(ThinkMode::NoThink),
Some("medium") | Some("high") => Ok(ThinkMode::Default),
Some(other) => Err(format!(
"bad reasoning_effort {other:?} (none|minimal|low|medium|high)")),
}
}
#[allow(clippy::type_complexity)]
fn prepare_tools(tools: &[serde_json::Value])
-> Result<(Vec<String>, HashMap<String, HashMap<String, String>>), String> {
let mut tools_json = Vec::with_capacity(tools.len());
let mut schemas: HashMap<String, HashMap<String, String>> = HashMap::new();
for t in tools {
let f = t.get("function").ok_or("each tool needs a function object")?;
let name = f.get("name").and_then(|n| n.as_str())
.ok_or("each tool needs function.name")?;
let mut params: HashMap<String, String> = HashMap::new();
if let Some(props) = f.get("parameters").and_then(|p| p.get("properties"))
.and_then(|p| p.as_object()) {
for (p, def) in props {
if let Some(ty) = def.get("type").and_then(|t| t.as_str()) {
params.insert(p.clone(), ty.to_string());
}
}
}
schemas.insert(name.to_string(), params);
tools_json.push(pyjson_str(t));
}
Ok((tools_json, schemas))
}
fn render_req_tool_call(tc: &ReqToolCall) -> Result<TmplToolCall, String> {
let parsed: serde_json::Value = match &tc.function.arguments {
serde_json::Value::Null => json!({}),
serde_json::Value::String(s) if s.trim().is_empty() => json!({}),
serde_json::Value::String(s) => serde_json::from_str(s)
.map_err(|e| format!("tool_calls arguments is not valid JSON: {e}"))?,
v @ serde_json::Value::Object(_) => v.clone(),
_ => return Err("tool_calls arguments must be a JSON object".into()),
};
let obj = parsed.as_object()
.ok_or("tool_calls arguments must decode to a JSON object")?;
let params = obj.iter().map(|(k, v)| {
let rendered = match v {
serde_json::Value::String(s) => s.clone(),
v @ (serde_json::Value::Object(_) | serde_json::Value::Array(_)) => pyjson_str(v),
scalar => scalar.to_string(),
};
(k.clone(), rendered)
}).collect();
Ok(TmplToolCall { name: tc.function.name.clone(), params })
}
fn tool_call_json(c: &ParsedToolCall) -> serde_json::Value {
json!({ "id": c.id, "type": "function",
"function": { "name": c.name, "arguments": c.arguments } })
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let args: Vec<String> = std::env::args().skip(1).collect();
if let Some(code) = auth::run_cli(&args) {
std::process::exit(code);
}
auth::init_from_env();
let models = parse_models_config();
eprintln!("[server] starting; models config = {models:?}");
let (cmd_tx, model_names, caps, metrics) = match worker::spawn(models) {
Ok(v) => v,
Err(err) => { eprintln!("[server] FATAL: worker init failed: {err}"); std::process::exit(1); }
};
eprintln!("[server] worker ready; serving models: {model_names:?}");
let state = AppState {
cmd_tx, models: model_names, caps, metrics,
started: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs()).unwrap_or(0),
inflight: Arc::new(Default::default()),
tenant_inflight: Arc::new(Default::default()),
};
let inflight_handle = state.inflight.clone();
let app = Router::new()
.route("/health", get(health))
.route("/models", get(list_models))
.route("/v1/models", get(list_models_v1))
.route("/v1/completions", post(completions))
.route("/v1/chat/completions", post(chat_completions))
.route("/metrics", get(get_metrics))
.route("/yield/metrics", get(yield_metrics))
.with_state(state);
let addr = std::env::var("MEMRA_ADDR").unwrap_or_else(|_| "127.0.0.1:8080".into());
let listener = tokio::net::TcpListener::bind(&addr).await?;
eprintln!("[server] listening on http://{addr}");
let inflight = inflight_handle;
axum::serve(listener, app)
.with_graceful_shutdown(async move {
let mut sigterm = match tokio::signal::unix::signal(
tokio::signal::unix::SignalKind::terminate()) {
Ok(s) => s,
Err(err) => {
eprintln!("[server] WARN: no SIGTERM handler ({err}); drain disabled");
std::future::pending::<()>().await;
unreachable!()
}
};
sigterm.recv().await;
DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
let n: usize = inflight.iter()
.map(|c| c.load(std::sync::atomic::Ordering::SeqCst)).sum();
eprintln!("[server] SIGTERM: draining ({n} in flight, deadline {}s)",
drain_deadline_s());
let deadline = std::time::Duration::from_secs(drain_deadline_s());
let t0 = std::time::Instant::now();
loop {
let n: usize = inflight.iter()
.map(|c| c.load(std::sync::atomic::Ordering::SeqCst)).sum();
if n == 0 {
eprintln!("[server] drain complete in {:.1}s; exiting",
t0.elapsed().as_secs_f64());
break;
}
if t0.elapsed() >= deadline {
eprintln!("[server] drain deadline ({}s) hit with {n} in flight; exiting",
drain_deadline_s());
break;
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
})
.await?;
Ok(())
}
fn validate_model_path(path: &str) -> Result<(), String> {
let p = std::path::Path::new(path);
if !p.exists() {
return Err(format!("model path {path:?} does not exist"));
}
if p.is_file() {
return Ok(()); }
if p.join("manifest.json").exists() {
return Ok(()); }
let has_st = p.join("model.safetensors").exists()
|| p.join("model.safetensors.index.json").exists();
if !has_st {
return Err(format!(
"model dir {path:?} is not a servable checkpoint: want model.safetensors or \
model.safetensors.index.json + config.json (HF safetensors dir), or \
manifest.json (memra repack dir)"));
}
if !p.join("config.json").exists() {
return Err(format!("model dir {path:?} has safetensors weights but no config.json"));
}
Ok(())
}
fn parse_models_config() -> Vec<(String, String, Option<String>)> {
if let Ok(spec) = std::env::var("MEMRA_MODELS") {
let mut out = Vec::new();
for entry in spec.split(',').filter(|s| !s.trim().is_empty()) {
if let Some((name, path)) = entry.split_once('=') {
let (mpath, dpath) = match path.trim().split_once('+') {
Some((m, d)) => (m.trim(), Some(d.trim())),
None => (path.trim(), None),
};
let resolve = |p: &str| memra_gguf::hf::resolve_arg(p).unwrap_or_else(|err| {
eprintln!("[server] FATAL: model {name:?}: {err}");
std::process::exit(1);
});
let mpath = resolve(mpath);
if let Err(err) = validate_model_path(&mpath) {
eprintln!("[server] FATAL: model {name:?}: {err}");
std::process::exit(1);
}
out.push((name.trim().to_string(), mpath, dpath.map(resolve)));
} else {
eprintln!("[server] WARN: bad MEMRA_MODELS entry {entry:?} (want name=/path[+/draft]); skipping");
}
}
if !out.is_empty() { return out; }
}
vec![
("main".into(), "/data/ai-ml/hf-models/qwen36-27b-nvfp4-mtp/Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf".into(), None),
("judge".into(), "/data/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf".into(), None),
]
}
async fn health(State(st): State<AppState>) -> impl IntoResponse {
let status = if draining() { "draining" } else { "ok" };
Json(json!({ "status": status, "models": *st.models }))
}
async fn get_metrics(State(st): State<AppState>) -> impl IntoResponse {
let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
let mut body = json!({
"admitted": m.admitted,
"completed": m.completed,
"tokens_out": m.tokens_out,
"step_p50_ms": m.step_p50_ms,
"step_p99_ms": m.step_p99_ms,
"prompt_tokens_in": m.prompt_tokens_in,
"cached_tokens_in": m.cached_tokens_in,
"prefix_cache_hits": m.prefix_hits,
"prefix_cache_entries": m.prefix_entries,
"prefix_cache_bytes": m.prefix_bytes,
});
let spec: serde_json::Map<String, serde_json::Value> = m.spec.iter().map(|(model, t)| {
let n_pos = t.pos_drafted.iter().rposition(|&d| d > 0).map_or(0, |p| p + 1);
(model.clone(), json!({
"rounds": t.rounds,
"drafted": t.drafted,
"accepted": t.accepted,
"acceptance_rate": if t.drafted > 0 {
t.accepted as f64 / t.drafted as f64 } else { 0.0 },
"tokens_per_round": if t.rounds > 0 {
(t.accepted + t.rounds) as f64 / t.rounds as f64 } else { 0.0 },
"pos_drafted": t.pos_drafted[..n_pos].to_vec(),
"pos_accepted": t.pos_accepted[..n_pos].to_vec(),
"accept_rate_per_pos": (0..n_pos).map(|j| if t.pos_drafted[j] > 0 {
t.pos_accepted[j] as f64 / t.pos_drafted[j] as f64 } else { 0.0 })
.collect::<Vec<f64>>(),
}))
}).collect();
if !spec.is_empty() {
body["spec"] = serde_json::Value::Object(spec);
}
Json(body)
}
async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
let data: Vec<_> = st.models.iter().map(|m| json!({ "id": m, "object": "model" })).collect();
Json(json!({ "object": "list", "data": data }))
}
fn model_entry_v1(name: &str, caps: Option<&ModelCaps>, created: u64) -> serde_json::Value {
let ctx = caps.map(|c| c.context_length).filter(|&c| c > 0);
let tokenizer = caps.map(|c| c.tokenizer.as_str()).filter(|t| !t.is_empty());
let instruct = caps.and_then(|c| c.instruct_type.as_deref());
json!({
"id": name,
"name": name,
"object": "model",
"created": created,
"context_length": ctx,
"architecture": {
"modality": "text->text",
"tokenizer": tokenizer,
"instruct_type": instruct,
},
"pricing": {
"prompt": "0",
"completion": "0",
"request": "0",
"image": "0",
},
"top_provider": {
"context_length": ctx,
"max_completion_tokens": serde_json::Value::Null,
},
})
}
async fn list_models_v1(State(st): State<AppState>) -> impl IntoResponse {
let data: Vec<_> = st.models.iter()
.map(|m| model_entry_v1(m, st.caps.get(m), st.started))
.collect();
Json(json!({ "object": "list", "data": data }))
}
async fn yield_metrics(State(st): State<AppState>) -> impl IntoResponse {
let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
let lane = |i: usize| json!({
"admitted": m.lane_admitted[i], "shed": m.lane_shed[i],
"completed": m.lane_completed[i], "tokens_out": m.lane_tokens[i],
});
Json(json!({
"lanes": {
"interactive": lane(0), "judge": lane(1), "harvest": lane(2),
},
"interactive_step_ms": { "p50": m.step_p50_ms, "p99": m.step_p99_ms },
"batch_size_last": m.batch_size_last,
}))
}
async fn peek_shed(
lane: lanes::Lane,
mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>,
) -> Result<tokio::sync::mpsc::UnboundedReceiver<Event>, Response> {
if lane == lanes::Lane::Interactive {
return Ok(rx);
}
match rx.recv().await {
Some(Event::Error(e)) if e.starts_with("shed:") => Err((
StatusCode::TOO_MANY_REQUESTS,
[(axum::http::header::RETRY_AFTER, "2")],
Json(json!({ "error": e })),
).into_response()),
first => {
let (tx2, rx2) = tokio::sync::mpsc::unbounded_channel();
if let Some(ev) = first {
let _ = tx2.send(ev);
}
tokio::spawn(async move {
while let Some(ev) = rx.recv().await {
if tx2.send(ev).is_err() { break; }
}
});
Ok(rx2)
}
}
}
fn build_request(req: &CompletionReq, tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>) -> Request {
let params = GenParams {
max_new: req.max_tokens.unwrap_or(worker::MAX_NEW_CTX_BOUNDED),
max_ctx: req.max_ctx,
eos: Vec::new(), };
let sampler_cfg = sampler_config(
req.temperature, req.top_k, req.top_p, req.min_p,
req.frequency_penalty, req.presence_penalty, req.repetition_penalty, req.seed);
Request {
model: req.model.clone(),
prompt_ids: req.prompt_ids.clone(),
prompt_text: req.prompt.clone(),
chat: req.chat,
chat_turns: Vec::new(),
tools_json: Vec::new(),
think: ThinkMode::Default,
params,
sampler_cfg,
stop_strings: req.stop.clone().into_vec(),
trace_id: req.trace_id.clone(),
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar: None, tx,
}
}
struct ChatPlan {
request: Request,
parser: Option<ToolStreamParser>,
}
fn build_chat_request(req: ChatCompletionReq, caps: Option<&ModelCaps>,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>)
-> Result<ChatPlan, String> {
let tool_choice = parse_tool_choice(&req.tool_choice)?;
if let Some(c) = caps {
if !c.chat_ok {
return Err(format!(
"model {:?} has no chat template (checkpoint carries neither \
tokenizer_config.json chat_template nor chat_template.jinja) — \
/v1/chat/completions unavailable; use /v1/completions with a raw prompt",
req.model));
}
}
let mut think = parse_think(&req.reasoning_effort, &req.reasoning)?;
let grammar = constrained::parse_response_format(req.response_format.as_ref())?;
if grammar.is_some() {
if let Some(c) = caps {
if c.qwen_think && think != ThinkMode::NoThink {
if c.think_switch {
think = ThinkMode::NoThink;
} else {
return Err("response_format requires disabling the model's think tail, \
but this chat template has no enable_thinking switch".into());
}
}
}
}
let (tools_json, schemas) = if !req.tools.is_empty() && tool_choice == ToolChoice::Auto {
prepare_tools(&req.tools)?
} else {
(Vec::new(), HashMap::new())
};
let mut turns: Vec<TmplTurn> = Vec::with_capacity(req.messages.len());
for msg in &req.messages {
let content = content_to_text(&msg.content)
.map_err(|e| format!("{} message: {e}", msg.role))?;
let tool_calls = msg.tool_calls.iter().map(render_req_tool_call)
.collect::<Result<Vec<_>, _>>()?;
if !tool_calls.is_empty() && msg.role != "assistant" {
return Err("tool_calls are only valid on assistant messages".into());
}
turns.push(TmplTurn { role: msg.role.clone(), content, tool_calls });
}
let has_tool_features = !tools_json.is_empty()
|| turns.iter().any(|t| t.role == "tool" || !t.tool_calls.is_empty());
if has_tool_features && !caps.map(|c| c.tools_branch).unwrap_or(false) {
return Err(format!("model {:?} chat template has no tools branch", req.model));
}
let think_open = caps.map(|c| c.qwen_think
&& !(think == ThinkMode::NoThink && c.think_switch)).unwrap_or(false);
let include_reasoning = req.include_reasoning.unwrap_or(true)
&& req.reasoning.as_ref()
.and_then(|r| r.get("exclude")).and_then(|v| v.as_bool()) != Some(true);
let parser = if !tools_json.is_empty() {
Some(ToolStreamParser::new(schemas, think_open)
.with_include_reasoning(include_reasoning))
} else if think_open {
Some(ToolStreamParser::reasoning_only(include_reasoning))
} else {
None
};
Ok(ChatPlan {
request: Request {
model: req.model,
prompt_ids: Vec::new(),
prompt_text: String::new(),
chat: false,
chat_turns: turns,
tools_json,
think,
params: GenParams {
max_new: req.max_tokens.unwrap_or(worker::MAX_NEW_CTX_BOUNDED),
max_ctx: req.max_ctx,
eos: Vec::new(),
},
sampler_cfg: sampler_config(
req.temperature, req.top_k, req.top_p, req.min_p,
req.frequency_penalty, req.presence_penalty, req.repetition_penalty,
req.seed),
stop_strings: req.stop.into_vec(),
trace_id: None,
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar,
tx,
},
parser,
})
}
fn authenticate(headers: &axum::http::HeaderMap) -> Result<auth::TenantCtx, Response> {
static SINGLE: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
let single = SINGLE.get_or_init(|| std::env::var("MEMRA_API_KEY").ok());
let bearer = headers.get("authorization")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "));
auth::authenticate_with(auth::global(), single.as_deref(), bearer).map_err(|why| {
match why {
auth::AuthDenied::Unknown => error_response(
StatusCode::UNAUTHORIZED, "invalid api key", "authentication_error", None),
auth::AuthDenied::Disabled => error_response(
StatusCode::FORBIDDEN, "api key is disabled", "authentication_error", None),
}
})
}
fn lane_for_tenant(headers: &axum::http::HeaderMap, tenant: &auth::TenantCtx)
-> Result<lanes::Lane, Response> {
let requested = match headers.get("x-lane").map(|v| v.to_str().unwrap_or("?")) {
None => None,
Some(v) => Some(lanes::Lane::parse(v).ok_or_else(|| {
(StatusCode::BAD_REQUEST,
Json(json!({ "error": format!("unknown x-lane {v:?}") }))).into_response()
})?),
};
match tenant.lane_class {
auth::LaneClass::Interactive => Ok(requested.unwrap_or(lanes::Lane::Interactive)),
auth::LaneClass::Batch => match requested {
None => Ok(lanes::Lane::Harvest),
Some(lanes::Lane::Interactive) => Err(error_response(
StatusCode::FORBIDDEN,
"this api key is batch-class: x-lane interactive is not permitted \
(use judge or harvest)",
"authentication_error", Some("x-lane"))),
Some(l) => Ok(l),
},
}
}
fn tenant_namespace(tenant: &auth::TenantCtx, cache_salt: &Option<String>) -> String {
let raw = cache_namespace(cache_salt);
if auth::global().is_some() {
auth::scope_namespace(&tenant.tenant, &raw)
} else {
raw
}
}
fn meter_admit(env: &Envelope, tenant: &auth::TenantCtx, model: &str, lane: lanes::Lane) {
eprintln!("[meter] admit id={} tenant={} lane={} model={:?}",
env.id, tenant.tenant, lane.as_str(), model);
}
async fn completions(State(st): State<AppState>, headers: axum::http::HeaderMap,
Json(req): Json<CompletionReq>) -> Response {
let env = Envelope::new(false);
let tenant = match authenticate(&headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
if let Err((msg, param)) = reject_unsupported(&[
("logit_bias", req.logit_bias.is_some(),
" (device-side sampling has no bias hook yet)"),
("logprobs", req.logprobs.is_some(), ""),
("n", req.n.is_some_and(|n| n != 1), " for n != 1 (single choice only)"),
("best_of", req.best_of.is_some_and(|n| n != 1), " (single choice only)"),
]) {
return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
}
let lane = match lane_for_tenant(&headers, &tenant) {
Ok(l) => l,
Err(resp) => return resp,
};
if draining() {
return with_request_id(&env.id, drain_response());
}
let (guard, n_inflight, n_tenant) = InflightGuard::acquire(
st.inflight.clone(), lane, st.tenant_inflight.clone(), &tenant.tenant);
let rl = RateLimit::at_admit(lane, n_inflight, &st.metrics, &tenant, n_tenant);
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Event>();
let model = req.model.clone();
let stream = req.stream;
let affinity = affinity_key(&req.session_id, &req.user, &headers);
let mut request = build_request(&req, tx, lane, affinity);
request.cache_ns = tenant_namespace(&tenant, &req.cache_salt);
meter_admit(&env, &tenant, &model, lane);
let stop_strings = request.stop_strings.clone();
if st.cmd_tx.send(Cmd::Generate(Box::new(request))).is_err() {
return rl.attach(with_request_id(&env.id, error_response(
StatusCode::SERVICE_UNAVAILABLE, "worker unavailable", "server_error", None)));
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => return rl.attach(with_request_id(&env.id, resp)),
};
let resp = if stream {
sse_response(rx, model, false, None, env.clone(), stop_strings, Some(guard))
.into_response()
} else {
let resp = blocking_response(rx, model, false, stop_strings, None, env.clone()).await
.into_response();
drop(guard); resp
};
rl.attach(with_request_id(&env.id, resp))
}
async fn chat_completions(State(st): State<AppState>, headers: axum::http::HeaderMap,
Json(req): Json<ChatCompletionReq>) -> Response {
let env = Envelope::new(true);
let tenant = match authenticate(&headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
if req.messages.is_empty() || req.messages.iter().any(|message| {
!matches!(message.role.as_str(), "system" | "user" | "assistant" | "tool")
}) {
return with_request_id(&env.id, bad_request(
"messages must use system/user/assistant/tool roles", Some("messages")));
}
if let Err((msg, param)) = reject_unsupported(&[
("logit_bias", req.logit_bias.is_some(),
" (device-side sampling has no bias hook yet)"),
("logprobs", req.logprobs.as_ref().is_some_and(|v| v.as_bool() != Some(false)), ""),
("top_logprobs", req.top_logprobs.is_some(), ""),
("n", req.n.is_some_and(|n| n != 1), " for n != 1 (single choice only)"),
]) {
return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
}
let lane = match lane_for_tenant(&headers, &tenant) {
Ok(l) => l,
Err(resp) => return resp,
};
let model = req.model.clone();
let stream = req.stream;
let cache_salt = req.cache_salt.clone();
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Event>();
let affinity = affinity_key(&req.session_id, &req.user, &headers);
let mut plan = match build_chat_request(req, st.caps.get(&model), tx, lane, affinity) {
Ok(plan) => plan,
Err(err) => {
return with_request_id(&env.id, bad_request(&err, None));
}
};
plan.request.cache_ns = tenant_namespace(&tenant, &cache_salt);
if draining() {
return with_request_id(&env.id, drain_response());
}
let (guard, n_inflight, n_tenant) = InflightGuard::acquire(
st.inflight.clone(), lane, st.tenant_inflight.clone(), &tenant.tenant);
let rl = RateLimit::at_admit(lane, n_inflight, &st.metrics, &tenant, n_tenant);
meter_admit(&env, &tenant, &model, lane);
let stop_strings = plan.request.stop_strings.clone();
if st.cmd_tx.send(Cmd::Generate(Box::new(plan.request))).is_err() {
return rl.attach(with_request_id(&env.id, error_response(
StatusCode::SERVICE_UNAVAILABLE, "worker unavailable", "server_error", None)));
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => return rl.attach(with_request_id(&env.id, resp)),
};
let resp = if stream {
sse_response(rx, model, true, plan.parser, env.clone(), stop_strings, Some(guard))
.into_response()
} else {
let resp = blocking_response(rx, model, true, stop_strings, plan.parser, env.clone())
.await.into_response();
drop(guard); resp
};
rl.attach(with_request_id(&env.id, resp))
}
fn sse_response(mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String, chat: bool,
mut parser: Option<ToolStreamParser>, env: Envelope,
stop_strings: Vec<String>, guard: Option<InflightGuard>)
-> Sse<impl futures_core::Stream<Item = Result<SseEvent, std::convert::Infallible>>> {
let mut scrub = (!stop_strings.is_empty() && (chat || openai_compat()))
.then(|| StopScrubber::new(stop_strings));
let stream = async_stream::stream! {
let _guard = guard;
let mut call_index: usize = 0;
let mut role_sent = false;
macro_rules! chat_chunk {
($delta:expr, $finish:expr) => {{
let mut delta = $delta;
if chat && !role_sent {
role_sent = true;
delta["role"] = json!("assistant");
}
env.stamp(json!({ "object": "chat.completion.chunk", "model": model,
"choices": [{ "index": 0, "delta": delta,
"finish_reason": $finish }] }))
.to_string()
}};
}
macro_rules! piece_chunks {
($piece:expr) => {{
let mut payloads: Vec<String> = Vec::new();
match $piece {
Piece::Content(text) => {
let text = match scrub.as_mut() {
Some(sc) => sc.push(&text),
None => text,
};
if !text.is_empty() {
payloads.push(chat_chunk!(json!({ "content": text }),
serde_json::Value::Null));
}
}
Piece::Reasoning(text) => payloads.push(
chat_chunk!(json!({ "reasoning": text }), serde_json::Value::Null)),
Piece::Call(call) => {
payloads.push(chat_chunk!(json!({ "tool_calls": [{
"index": call_index, "id": call.id, "type": "function",
"function": { "name": call.name, "arguments": "" } }] }),
serde_json::Value::Null));
payloads.push(chat_chunk!(json!({ "tool_calls": [{
"index": call_index,
"function": { "arguments": call.arguments } }] }),
serde_json::Value::Null));
call_index += 1;
}
}
payloads
}};
}
while let Some(ev) = rx.recv().await {
match ev {
Event::Token { id, text } => {
if let Some(p) = parser.as_mut() {
for piece in p.push(&text) {
for payload in piece_chunks!(piece) {
yield Ok(SseEvent::default().data(payload));
}
}
continue;
}
let text = match scrub.as_mut() {
Some(sc) => sc.push(&text),
None => text,
};
if text.is_empty() && scrub.is_some() {
continue; }
let payload = if chat {
chat_chunk!(json!({ "content": text }), serde_json::Value::Null)
} else if openai_compat() {
env.stamp(json!({ "object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": text, "finish_reason": null }] }))
.to_string()
} else {
json!({ "model": model, "id": id, "text": text }).to_string()
};
yield Ok(SseEvent::default().data(payload));
}
Event::Done { stop_reason, n_tokens, n_prompt, n_cached, elapsed_s, spec } => {
let mut finish = stop_reason_to_finish(&stop_reason);
if let Some(p) = parser.as_mut() {
for piece in p.finish() {
for payload in piece_chunks!(piece) {
yield Ok(SseEvent::default().data(payload));
}
}
if p.n_calls() > 0 { finish = "tool_calls"; }
}
if let Some(sc) = scrub.as_mut() {
let tail = sc.finish();
if !tail.is_empty() {
let payload = if chat {
chat_chunk!(json!({ "content": tail }),
serde_json::Value::Null)
} else {
env.stamp(json!({ "object": "text_completion",
"model": model,
"choices": [{ "index": 0, "text": tail,
"finish_reason": null }] })).to_string()
};
yield Ok(SseEvent::default().data(payload));
}
}
if chat || openai_compat() {
let usage = usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec);
let fin = if chat {
let mut v = env.stamp(json!({
"object": "chat.completion.chunk", "model": model,
"choices": [{ "index": 0, "delta": {},
"finish_reason": finish }],
"usage": usage }));
if !role_sent {
v["choices"][0]["delta"]["role"] = json!("assistant");
}
v
} else {
env.stamp(json!({ "object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": "",
"finish_reason": finish }],
"usage": usage }))
}.to_string();
yield Ok(SseEvent::default().data(fin));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
let payload = json!({
"stop_reason": stop_reason, "n_tokens": n_tokens,
"prompt_tokens": n_prompt, "cached_tokens": n_cached,
"elapsed_s": elapsed_s
}).to_string();
yield Ok(SseEvent::default().event("done").data(payload));
}
break;
}
Event::Error(msg) => {
if chat || openai_compat() {
let payload = error_body(&msg, "server_error", None, None).to_string();
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
let payload = json!({ "error": msg }).to_string();
yield Ok(SseEvent::default().event("error").data(payload));
}
break;
}
}
}
};
Sse::new(stream).keep_alive(
axum::response::sse::KeepAlive::new()
.interval(std::time::Duration::from_secs(5)),
)
}
fn truncate_at_stop(text: &mut String, stop_strings: &[String]) {
if let Some(offset) = stop_strings.iter().filter_map(|stop| text.find(stop)).min() {
text.truncate(offset);
}
}
fn partial_stop_suffix(s: &str, tag: &str) -> usize {
let mut best = 0;
for (k, _) in tag.char_indices().skip(1) {
if k <= s.len() && s.ends_with(&tag[..k]) {
best = k;
}
}
best
}
struct StopScrubber {
stops: Vec<String>,
buf: String,
done: bool,
}
impl StopScrubber {
fn new(stops: Vec<String>) -> Self {
Self { stops, buf: String::new(), done: false }
}
fn push(&mut self, text: &str) -> String {
if self.done {
return String::new();
}
self.buf.push_str(text);
if let Some(i) = self.stops.iter().filter_map(|s| self.buf.find(s.as_str())).min() {
self.done = true;
let out = self.buf[..i].to_string();
self.buf.clear();
return out;
}
let keep = self.stops.iter()
.map(|s| partial_stop_suffix(&self.buf, s)).max().unwrap_or(0);
let emit_to = self.buf.len() - keep;
let out = self.buf[..emit_to].to_string();
self.buf.drain(..emit_to);
out
}
fn finish(&mut self) -> String {
if self.done {
self.buf.clear();
return String::new();
}
std::mem::take(&mut self.buf)
}
}
async fn blocking_response(mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String,
chat: bool, stop_strings: Vec<String>,
mut parser: Option<ToolStreamParser>, env: Envelope) -> Response {
let mut text = String::new();
let mut reasoning = String::new();
let mut tokens: Vec<u32> = Vec::new();
let mut calls: Vec<ParsedToolCall> = Vec::new();
let consume = |pieces: Vec<Piece>, text: &mut String, reasoning: &mut String,
calls: &mut Vec<ParsedToolCall>| {
for piece in pieces {
match piece {
Piece::Content(t) => text.push_str(&t),
Piece::Reasoning(t) => reasoning.push_str(&t),
Piece::Call(c) => calls.push(c),
}
}
};
while let Some(ev) = rx.recv().await {
match ev {
Event::Token { id, text: delta } => {
tokens.push(id);
match parser.as_mut() {
Some(p) => consume(p.push(&delta), &mut text, &mut reasoning, &mut calls),
None => text.push_str(&delta),
}
}
Event::Done { stop_reason, n_tokens, n_prompt, n_cached, elapsed_s, spec } => {
if let Some(p) = parser.as_mut() {
consume(p.finish(), &mut text, &mut reasoning, &mut calls);
}
truncate_at_stop(&mut text, &stop_strings);
let finish = if calls.is_empty() { stop_reason_to_finish(&stop_reason) }
else { "tool_calls" };
if chat {
let content = if !calls.is_empty() && text.is_empty() {
serde_json::Value::Null
} else {
serde_json::Value::String(text)
};
let mut message = json!({ "role": "assistant", "content": content });
if !reasoning.is_empty() {
message["reasoning"] = json!(reasoning);
message["reasoning_details"] = json!([{
"type": "reasoning.text", "text": reasoning }]);
}
if !calls.is_empty() {
message["tool_calls"] = serde_json::Value::Array(
calls.iter().map(tool_call_json).collect());
}
return Json(env.stamp(json!({
"object": "chat.completion", "model": model,
"choices": [{ "index": 0,
"message": message,
"finish_reason": finish }],
"usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
}))).into_response();
}
if openai_compat() {
return Json(env.stamp(json!({
"object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": text,
"finish_reason": finish }],
"usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
}))).into_response();
}
return Json(CompletionResp {
model, text, tokens, stop_reason, n_tokens,
prompt_tokens: n_prompt, cached_tokens: n_cached, elapsed_s,
}).into_response();
}
Event::Error(msg) => {
return bad_request(&msg, None);
}
}
}
error_response(StatusCode::INTERNAL_SERVER_ERROR, "worker closed stream", "server_error", None)
}
#[cfg(test)]
mod tests {
use super::*;
fn tool_caps() -> ModelCaps {
ModelCaps {
tools_branch: true, qwen_think: true, think_switch: true, chat_ok: true,
..Default::default()
}
}
#[test]
fn chat_request_preserves_turns_and_openai_stop_forms() {
let payload = serde_json::json!({
"model": "plain_quant",
"messages": [
{"role": "system", "content": "rules"},
{"role": "user", "content": "task"},
{"role": "assistant", "content": "work"}
],
"max_tokens": 64,
"temperature": 0.0,
"stop": "<stop>"
});
let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
let request = plan.request;
assert!(plan.parser.is_none(), "no tools -> no parser (isolation contract)");
assert!(request.tools_json.is_empty());
assert_eq!(request.think, ThinkMode::Default);
assert_eq!(request.model, "plain_quant");
assert_eq!(request.params.max_new, 64);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}]
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
assert_eq!(plan.request.params.max_new, worker::MAX_NEW_CTX_BOUNDED);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"max_completion_tokens": 7
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap().request.params.max_new, 7);
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, None).params.max_new, worker::MAX_NEW_CTX_BOUNDED);
let turns: Vec<(String, String)> = request.chat_turns.iter()
.map(|t| (t.role.clone(), t.content.clone())).collect();
assert_eq!(turns, vec![
("system".into(), "rules".into()),
("user".into(), "task".into()),
("assistant".into(), "work".into()),
]);
assert!(request.chat_turns.iter().all(|t| t.tool_calls.is_empty()));
assert_eq!(request.stop_strings, vec!["<stop>"]);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"stop": ["a", "b"]
})).unwrap();
assert_eq!(req.stop.into_vec(), vec!["a", "b"]);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"stop": null
})).unwrap();
assert!(req.stop.into_vec().is_empty());
}
#[tokio::test]
async fn chat_response_has_openai_message_shape() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "hello".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 42, n_cached: 30, elapsed_s: 0.5,
spec: None,
}).unwrap();
drop(tx);
let response = blocking_response(rx, "plain_quant".into(), true, Vec::new(), None,
Envelope::new(true)).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["object"], "chat.completion");
assert!(payload["id"].as_str().unwrap().starts_with("chatcmpl-"));
assert!(payload["created"].as_u64().unwrap() > 1_700_000_000);
assert!(payload["system_fingerprint"].as_str().unwrap().starts_with("memra-"));
assert_eq!(payload["choices"][0]["message"]["role"], "assistant");
assert_eq!(payload["choices"][0]["message"]["content"], "hello");
assert_eq!(payload["choices"][0]["finish_reason"], "stop");
assert_eq!(payload["usage"]["prompt_tokens"], 42);
assert_eq!(payload["usage"]["completion_tokens"], 1);
assert_eq!(payload["usage"]["total_tokens"], 43);
assert_eq!(payload["usage"]["prompt_tokens_details"]["cached_tokens"], 30);
assert!(payload["usage"].get("spec").is_none());
}
#[tokio::test]
async fn chat_usage_carries_spec_acceptance_summary() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "hello".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 42, n_cached: 0, elapsed_s: 0.5,
spec: Some(worker::SpecUsage { rounds: 10, drafted: 30, accepted: 21 }),
}).unwrap();
drop(tx);
let response = blocking_response(rx, "plain_quant".into(), true, Vec::new(), None,
Envelope::new(true)).await;
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
let sp = &payload["usage"]["spec"];
assert_eq!(sp["rounds"], 10);
assert_eq!(sp["drafted"], 30);
assert_eq!(sp["accepted"], 21);
assert!((sp["acceptance_rate"].as_f64().unwrap() - 0.7).abs() < 1e-9);
assert_eq!(payload["usage"]["total_tokens"], 43);
}
fn weather_request(extra: serde_json::Value) -> ChatCompletionReq {
let mut payload = serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "Weather in Paris?"}],
"tools": [{"type": "function", "function": {
"name": "get_weather",
"description": "Get current weather",
"parameters": {"type": "object",
"properties": {"city": {"type": "string"},
"days": {"type": "integer"}},
"required": ["city"]}}}],
});
if let Some(obj) = extra.as_object() {
for (k, v) in obj { payload[k] = v.clone(); }
}
serde_json::from_value(payload).unwrap()
}
#[test]
fn tools_request_renders_client_key_order_and_arms_parser() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(json!({})), Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
assert!(plan.parser.is_some());
assert_eq!(plan.request.tools_json.len(), 1);
assert_eq!(plan.request.tools_json[0],
"{\"type\": \"function\", \"function\": {\"name\": \"get_weather\", \
\"description\": \"Get current weather\", \"parameters\": {\"type\": \"object\", \
\"properties\": {\"city\": {\"type\": \"string\"}, \"days\": {\"type\": \
\"integer\"}}, \"required\": [\"city\"]}}}");
}
#[test]
fn tool_choice_none_strips_tools_and_parser() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(json!({"tool_choice": "none"})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
let mut p = plan.parser.expect("think-open chat arms the reasoning splitter");
let pieces = p.push("x</think>\n\n<tool_call> stays prose");
assert_eq!(pieces, vec![
Piece::Reasoning("x".into()),
Piece::Content("<tool_call> stays prose".into()),
]);
assert!(plan.request.tools_json.is_empty());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({"tool_choice": "required"})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).is_err());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({"tool_choice":
{"type": "function", "function": {"name": "get_weather"}}})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).is_err());
}
#[test]
fn model_plan_accepts_st_dir_and_rejects_bogus_dir() {
let root = std::env::temp_dir().join(format!("memra_plan_test_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&root);
let st = root.join("st_single");
std::fs::create_dir_all(&st).unwrap();
std::fs::write(st.join("config.json"), "{}").unwrap();
std::fs::write(st.join("model.safetensors"), b"x").unwrap();
assert!(validate_model_path(st.to_str().unwrap()).is_ok());
let sh = root.join("st_sharded");
std::fs::create_dir_all(&sh).unwrap();
std::fs::write(sh.join("config.json"), "{}").unwrap();
std::fs::write(sh.join("model.safetensors.index.json"), "{}").unwrap();
assert!(validate_model_path(sh.to_str().unwrap()).is_ok());
let rp = root.join("repack");
std::fs::create_dir_all(&rp).unwrap();
std::fs::write(rp.join("manifest.json"), "{}").unwrap();
assert!(validate_model_path(rp.to_str().unwrap()).is_ok());
let bogus = root.join("bogus");
std::fs::create_dir_all(&bogus).unwrap();
let err = validate_model_path(bogus.to_str().unwrap()).unwrap_err();
assert!(err.contains("model.safetensors"), "error should say what is missing: {err}");
assert!(err.contains("manifest.json"), "error should mention the repack form: {err}");
let nc = root.join("no_config");
std::fs::create_dir_all(&nc).unwrap();
std::fs::write(nc.join("model.safetensors"), b"x").unwrap();
let err = validate_model_path(nc.to_str().unwrap()).unwrap_err();
assert!(err.contains("config.json"), "error should name config.json: {err}");
let err = validate_model_path(root.join("nowhere").to_str().unwrap()).unwrap_err();
assert!(err.contains("does not exist"), "{err}");
let f = root.join("model.gguf");
std::fs::write(&f, b"g").unwrap();
assert!(validate_model_path(f.to_str().unwrap()).is_ok());
let _ = std::fs::remove_dir_all(&root);
}
#[test]
fn chat_on_templateless_dir_checkpoint_is_rejected_with_clear_message() {
let caps = ModelCaps {
tools_branch: false, qwen_think: false, think_switch: false, chat_ok: false,
..Default::default() };
let payload = serde_json::json!({
"model": "st_model",
"messages": [{"role": "user", "content": "hello"}],
});
let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let err = match build_chat_request(req, Some(&caps), tx, lanes::Lane::Interactive, None) {
Err(e) => e,
Ok(_) => panic!("templateless dir checkpoint must reject chat"),
};
assert!(err.contains("no chat template"), "message should name the cause: {err}");
assert!(err.contains("/v1/completions"), "message should point at the raw-prompt escape hatch: {err}");
}
#[test]
fn tools_on_model_without_tools_branch_is_rejected() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let caps = ModelCaps { chat_ok: true, ..Default::default() };
assert!(build_chat_request(weather_request(json!({})), Some(&caps), tx, lanes::Lane::Interactive, None).is_err());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({})), None, tx, lanes::Lane::Interactive, None).is_err());
}
#[test]
fn reasoning_effort_maps_to_think_switch() {
for (extra, want) in [
(json!({}), ThinkMode::Default),
(json!({"reasoning_effort": "low"}), ThinkMode::NoThink),
(json!({"reasoning_effort": "none"}), ThinkMode::NoThink),
(json!({"reasoning_effort": "minimal"}), ThinkMode::NoThink),
(json!({"reasoning_effort": "high"}), ThinkMode::Default),
(json!({"reasoning_effort": "medium"}), ThinkMode::Default),
(json!({"reasoning": {"enabled": false}}), ThinkMode::NoThink),
(json!({"reasoning": {"effort": "low"}}), ThinkMode::NoThink),
(json!({"reasoning": {"enabled": true}}), ThinkMode::Default),
] {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(extra.clone()),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
assert_eq!(plan.request.think, want, "extra={extra}");
}
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({"reasoning_effort": "extreme"})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).is_err());
}
#[test]
fn assistant_history_tool_calls_and_tool_role_render_into_turns() {
let payload = serde_json::json!({
"model": "m",
"messages": [
{"role": "user", "content": "Weather in Paris?"},
{"role": "assistant", "content": null, "tool_calls": [
{"id": "call_x", "type": "function", "function": {
"name": "get_weather",
"arguments": "{\"city\": \"Paris\", \"days\": 3}"}}]},
{"role": "tool", "tool_call_id": "call_x", "content": "{\"temp_c\": 21}"}
],
});
let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(req, Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
let turns = &plan.request.chat_turns;
assert_eq!(turns[1].tool_calls, vec![TmplToolCall {
name: "get_weather".into(),
params: vec![("city".into(), "Paris".into()), ("days".into(), "3".into())],
}]);
assert_eq!(turns[2].role, "tool");
assert_eq!(turns[2].content, "{\"temp_c\": 21}");
let mut p = plan.parser.expect("think-open chat arms the reasoning splitter");
let pieces = p.push("thought</think>\n\nanswer <tool_call> is prose here");
assert_eq!(pieces, vec![
Piece::Reasoning("thought".into()),
Piece::Content("answer <tool_call> is prose here".into()),
]);
}
#[tokio::test]
async fn blocking_tools_response_carries_tool_calls_and_finish_reason() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "plan</think>\n\n".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "<tool_call>\n<function=get_weather>\n\
<parameter=city>\nParis\n</parameter>\n</function>\n</tool_call>".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 2, n_prompt: 40, n_cached: 0, elapsed_s: 0.5,
spec: None,
}).unwrap();
drop(tx);
let parser = ToolStreamParser::new(HashMap::new(), true);
let response = blocking_response(rx, "m".into(), true, Vec::new(), Some(parser),
Envelope::new(true)).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["choices"][0]["finish_reason"], "tool_calls");
assert_eq!(payload["choices"][0]["message"]["content"], serde_json::Value::Null);
assert_eq!(payload["choices"][0]["message"]["reasoning"], "plan");
assert_eq!(payload["choices"][0]["message"]["reasoning_details"][0]["text"], "plan");
let call = &payload["choices"][0]["message"]["tool_calls"][0];
assert_eq!(call["type"], "function");
assert_eq!(call["function"]["name"], "get_weather");
assert_eq!(call["function"]["arguments"], "{\"city\":\"Paris\"}");
assert_eq!(payload["usage"]["prompt_tokens"], 40);
assert_eq!(payload["usage"]["completion_tokens"], 2);
assert_eq!(payload["usage"]["total_tokens"], 42);
assert_eq!(payload["usage"]["prompt_tokens_details"]["cached_tokens"], 0);
}
#[test]
fn cache_salt_plumbs_to_the_worker_namespace() {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task", "cache_salt": "tenant-a"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, None).cache_ns, "tenant-a");
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "task"}],
"cache_salt": "tenant-b"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap().request.cache_ns, "tenant-b");
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, None).cache_ns, "");
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "task"}]
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap().request.cache_ns, "");
}
#[test]
fn affinity_key_honors_both_client_conventions_in_priority_order() {
use axum::http::HeaderMap;
let hdr = |v: &str| {
let mut h = HeaderMap::new();
h.insert("x-session-id", v.parse().unwrap());
h
};
let empty = HeaderMap::new();
let s = |v: &str| Some(v.to_string());
assert_eq!(affinity_key(&s("explicit"), &None, &empty), s("explicit"));
assert_eq!(affinity_key(&None, &s("openai-user"), &empty), s("openai-user"));
assert_eq!(affinity_key(&None, &None, &hdr("hdr-id")), s("hdr-id"));
assert_eq!(affinity_key(&s("a"), &s("b"), &hdr("c")), s("a"));
assert_eq!(affinity_key(&None, &s("b"), &hdr("c")), s("b"));
assert_eq!(affinity_key(&s(" "), &s(""), &hdr(" ")), None);
assert_eq!(affinity_key(&s(""), &s("real"), &empty), s("real"));
assert_eq!(affinity_key(&s(" padded "), &None, &empty), s("padded"));
assert_eq!(affinity_key(&None, &None, &empty), None);
}
#[test]
fn affinity_key_plumbs_to_the_worker_request_on_both_bodies() {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task", "session_id": "conv-1"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new());
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, key).affinity.as_deref(),
Some("conv-1"));
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "task"}],
"user": "conv-2"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new());
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, key)
.unwrap().request.affinity.as_deref(), Some("conv-2"));
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_request(&req, tx, lanes::Lane::Interactive, None).affinity.is_none());
}
async fn sse_data_lines(resp: Response) -> Vec<String> {
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
String::from_utf8(bytes.to_vec()).unwrap()
.lines()
.filter_map(|l| l.strip_prefix("data: ").map(str::to_string))
.collect()
}
#[tokio::test]
async fn stream_chunks_carry_envelope_and_first_delta_role() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "he".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "llo".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 2, n_prompt: 10, n_cached: 0, elapsed_s: 0.1,
spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true), Vec::new(), None)
.into_response();
let lines = sse_data_lines(resp).await;
assert_eq!(lines.last().map(String::as_str), Some("[DONE]"));
let chunks: Vec<serde_json::Value> = lines[..lines.len() - 1].iter()
.map(|l| serde_json::from_str(l).unwrap()).collect();
let id = chunks[0]["id"].as_str().unwrap().to_string();
assert!(id.starts_with("chatcmpl-"));
for c in &chunks {
assert_eq!(c["id"], id.as_str());
assert!(c["created"].as_u64().unwrap() > 1_700_000_000);
assert!(c["system_fingerprint"].as_str().unwrap().starts_with("memra-"));
assert_eq!(c["object"], "chat.completion.chunk");
}
assert_eq!(chunks[0]["choices"][0]["delta"]["role"], "assistant");
assert_eq!(chunks[0]["choices"][0]["delta"]["content"], "he");
assert!(chunks[1]["choices"][0]["delta"].get("role").is_none());
let fin = chunks.last().unwrap();
assert_eq!(fin["choices"][0]["finish_reason"], "stop");
assert_eq!(fin["usage"]["prompt_tokens"], 10);
}
#[tokio::test]
async fn stream_excludes_stop_text_like_non_stream_does() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "answer\nPro".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "blem: leaked prompt".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Callback".into(), n_tokens: 2, n_prompt: 8, n_cached: 0,
elapsed_s: 0.1, spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true),
vec!["Problem:".into()], None).into_response();
let lines = sse_data_lines(resp).await;
let content: String = lines.iter()
.filter(|l| *l != "[DONE]")
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.filter_map(|c| c["choices"][0]["delta"]["content"].as_str()
.map(str::to_string))
.collect();
assert_eq!(content, "answer\n");
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "ends in Pro".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 8, n_cached: 0, elapsed_s: 0.1,
spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true),
vec!["Problem:".into()], None).into_response();
let lines = sse_data_lines(resp).await;
let content: String = lines.iter()
.filter(|l| *l != "[DONE]")
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.filter_map(|c| c["choices"][0]["delta"]["content"].as_str()
.map(str::to_string))
.collect();
assert_eq!(content, "ends in Pro");
}
#[tokio::test]
async fn stream_worker_error_is_a_data_chunk_not_a_named_event() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error("boom".into())).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true), Vec::new(), None)
.into_response();
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
let body = String::from_utf8(bytes.to_vec()).unwrap();
assert!(!body.contains("event: error"), "named SSE event leaked: {body}");
let lines: Vec<&str> = body.lines()
.filter_map(|l| l.strip_prefix("data: ")).collect();
let err: serde_json::Value = serde_json::from_str(lines[0]).unwrap();
assert_eq!(err["error"]["message"], "boom");
assert_eq!(err["error"]["type"], "server_error");
assert_eq!(lines.last(), Some(&"[DONE]"));
}
#[tokio::test]
async fn error_bodies_use_the_openai_object_shape() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error("unknown model \"x\"".into())).unwrap();
drop(tx);
let response = blocking_response(rx, "m".into(), true, Vec::new(), None,
Envelope::new(true)).await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["error"]["message"], "unknown model \"x\"");
assert_eq!(payload["error"]["type"], "invalid_request_error");
assert!(payload["error"].get("param").is_some());
assert!(payload["error"].get("code").is_some());
}
#[test]
fn penalties_plumb_from_http_to_sampler_config() {
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "task"}],
"frequency_penalty": 0.5, "presence_penalty": 0.25, "repetition_penalty": 1.1
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let cfg = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap().request.sampler_cfg;
assert_eq!(cfg.penalty_freq, 0.5);
assert_eq!(cfg.penalty_present, 0.25);
assert_eq!(cfg.penalty_repeat, 1.1);
assert_eq!(cfg.penalty_last_n, usize::MAX);
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task", "frequency_penalty": 1.5
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
assert_eq!(cfg.penalty_freq, 1.5);
assert_eq!(cfg.penalty_last_n, usize::MAX);
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
assert_eq!(cfg.penalty_last_n, 0);
assert_eq!(cfg.penalty_repeat, 1.0);
}
#[test]
fn omitted_temperature_is_openai_default_not_greedy() {
let chat_temp = |body: serde_json::Value| {
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
.unwrap().request.sampler_cfg.temperature
};
let comp_temp = |body: serde_json::Value| {
let req: CompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg.temperature
};
assert_eq!(chat_temp(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]})), 1.0,
"omitted chat temperature must be the OpenAI 1.0 default, not 0.0/greedy");
assert_eq!(comp_temp(serde_json::json!({
"model": "m", "prompt": "t"})), 1.0,
"omitted completions temperature must be the OpenAI 1.0 default");
assert_eq!(chat_temp(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"temperature": 0.0})), 0.0, "explicit temperature 0 must stay greedy");
assert_eq!(comp_temp(serde_json::json!({
"model": "m", "prompt": "t", "temperature": 0})), 0.0,
"explicit temperature 0 must stay greedy");
assert!(memra_engine::sampler::Sampler::new(
sampler_config(0.0, 0, 1.0, 0.0, 0.0, 0.0, 1.0, Some(0))).is_greedy());
assert!(!memra_engine::sampler::Sampler::new(
sampler_config(1.0, 0, 1.0, 0.0, 0.0, 0.0, 1.0, Some(0))).is_greedy());
assert_eq!(chat_temp(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"temperature": 0.7})), 0.7);
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "t"})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
assert_eq!(cfg.top_p, 1.0, "omitted top_p = OpenAI 1.0 = disabled");
assert_eq!(cfg.top_k, 0, "omitted top_k = disabled");
assert_eq!(cfg.min_p, 0.0, "omitted min_p = disabled");
assert_eq!(cfg.penalty_last_n, 0, "omitted penalties = window off");
assert!(memra_engine::sampler::Sampler::new(cfg).is_spec_sampling(),
"the omitted-temperature default must ride sampled spec's pure-temp regime");
}
#[test]
fn omitted_seed_is_fresh_entropy_not_a_pinned_zero() {
let comp_seed = |body: serde_json::Value| {
let req: CompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg.seed
};
let chat_seed = |body: serde_json::Value| {
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
.unwrap().request.sampler_cfg.seed
};
let a = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
let b = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
let c = chat_seed(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]}));
assert_ne!(a, 0, "omitted seed must not be the pinned 0 that caused the loop");
assert_ne!(b, 0);
assert_ne!(c, 0);
assert_ne!(a, b, "two seed-omitting requests must get DIFFERENT streams");
assert_ne!(a, c);
assert_eq!(comp_seed(serde_json::json!({
"model": "m", "prompt": "t", "seed": 0})), 0,
"explicit seed 0 must stay 0 — the determinism gates depend on it");
assert_eq!(comp_seed(serde_json::json!({
"model": "m", "prompt": "t", "seed": 12345})), 12345);
assert_eq!(chat_seed(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"seed": 777})), 777);
assert_eq!(comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42})),
comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42})));
let seeds: std::collections::HashSet<u64> = (0..256).map(|_| fresh_seed()).collect();
assert_eq!(seeds.len(), 256, "fresh_seed must not collide across rapid calls");
assert!(!seeds.contains(&0));
}
#[test]
fn response_format_builds_grammar_only_when_present() {
let mk = |rf: Option<serde_json::Value>| {
let mut body = serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]});
if let Some(rf) = rf { body["response_format"] = rf; }
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
};
assert!(mk(None).unwrap().request.grammar.is_none());
assert!(mk(Some(serde_json::json!({"type": "text"})))
.unwrap().request.grammar.is_none());
assert!(matches!(mk(Some(serde_json::json!({"type": "json_object"})))
.unwrap().request.grammar,
Some(constrained::GrammarSpec::JsonObject)));
assert!(matches!(mk(Some(serde_json::json!({"type": "json_schema",
"json_schema": {"schema": {"type": "object"}}})))
.unwrap().request.grammar,
Some(constrained::GrammarSpec::JsonSchema(_))));
assert!(mk(Some(serde_json::json!({"type": "yaml"}))).is_err());
}
#[test]
fn unsupported_semantic_params_are_named_rejections() {
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"response_format": {"type": "json_object"}
})).unwrap();
assert!(req.response_format.is_some());
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"response_format": {"type": "text"}, "logprobs": false, "n": 1,
"user": "u-1", "stream_options": {"include_usage": true}
})).unwrap();
assert_eq!(req.response_format.as_ref().unwrap()["type"], "text");
assert_eq!(req.logprobs.as_ref().unwrap().as_bool(), Some(false));
assert_eq!(req.n, Some(1));
assert!(reject_unsupported(&[("logit_bias", false, "")]).is_ok());
let (msg, param) = reject_unsupported(&[("logit_bias", true, " (why)")]).unwrap_err();
assert_eq!(param, "logit_bias");
assert_eq!(msg, "logit_bias is not supported (why)");
}
#[test]
fn completions_accept_openai_stop_forms() {
for (value, expected) in [
(serde_json::json!("Problem:"), vec!["Problem:"]),
(serde_json::json!(["Question:", "Problem:"]), vec!["Question:", "Problem:"]),
(serde_json::Value::Null, Vec::<&str>::new()),
] {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "prompt": "task", "stop": value
})).unwrap();
assert_eq!(req.stop.into_vec(), expected);
}
}
fn fake_worker_state() -> AppState {
let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
std::thread::spawn(move || {
while let Ok(Cmd::Generate(req)) = cmd_rx.recv() {
let _ = req.tx.send(Event::Token { id: 1, text: "ok".into() });
let _ = req.tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 1, n_cached: 0,
elapsed_s: 0.01, spec: None,
});
}
});
AppState {
cmd_tx,
models: Arc::new(vec!["m".into()]),
caps: Arc::new(HashMap::new()),
metrics: SharedMetrics::default(),
started: 1,
inflight: Arc::new(Default::default()),
tenant_inflight: Arc::new(Default::default()),
}
}
#[test]
fn rate_limit_math_remaining_hits_zero_at_cap_and_reset_arms() {
let metrics = SharedMetrics::default();
let rl = RateLimit::compute(4, 1, &metrics);
assert_eq!((rl.limit, rl.remaining, rl.reset_s), (4, 3, 0));
let rl = RateLimit::compute(4, 3, &metrics);
assert_eq!(rl.remaining, 1);
let rl = RateLimit::compute(4, 4, &metrics);
assert_eq!(rl.remaining, 0);
assert!(rl.reset_s > 0, "reset must arm when no slots are free");
assert_eq!(RateLimit::compute(4, 9, &metrics).remaining, 0);
let m = worker::Metrics {
completed: 2, tokens_out: 200, step_p50_ms: 20.0, ..Default::default()
};
assert_eq!(reset_estimate_s(&m), 2); }
#[test]
fn inflight_guard_counts_up_and_frees_on_drop() {
let counts: InflightCounts = Arc::new(Default::default());
let tenants: TenantGauge = Arc::new(Default::default());
let (g1, n1, t1) = InflightGuard::acquire(
counts.clone(), lanes::Lane::Interactive, tenants.clone(), "acme");
let (g2, n2, t2) = InflightGuard::acquire(
counts.clone(), lanes::Lane::Interactive, tenants.clone(), "acme");
assert_eq!((n1, n2), (1, 2));
assert_eq!((t1, t2), (1, 2));
let (gj, nj, tj) = InflightGuard::acquire(
counts.clone(), lanes::Lane::Judge, tenants.clone(), "blue");
assert_eq!((nj, tj), (1, 1));
drop(g1);
drop(gj);
assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 1);
assert_eq!(counts[1].load(std::sync::atomic::Ordering::SeqCst), 0);
assert_eq!(tenants.lock().unwrap().get("acme"), Some(&1));
assert!(tenants.lock().unwrap().get("blue").is_none());
drop(g2);
assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 0);
assert!(tenants.lock().unwrap().is_empty());
}
#[test]
fn tenant_rate_limit_override_is_min_with_global_cap() {
let metrics = SharedMetrics::default();
let unlimited = auth::TenantCtx::default_tenant();
let capped = auth::TenantCtx {
tenant: "acme".into(),
lane_class: auth::LaneClass::Interactive,
rate_limit: Some(2),
};
let global = lane_cap(lanes::Lane::Interactive);
let rl = RateLimit::at_admit(lanes::Lane::Interactive, 1, &metrics, &unlimited, 1);
assert_eq!((rl.limit, rl.remaining), (global, global - 1));
let rl = RateLimit::at_admit(lanes::Lane::Interactive, 5, &metrics, &capped, 1);
assert_eq!((rl.limit, rl.remaining), (2, 1));
let rl = RateLimit::at_admit(lanes::Lane::Interactive, 5, &metrics, &capped, 2);
assert_eq!(rl.remaining, 0);
assert!(rl.reset_s > 0, "reset must arm at the tenant cap too");
let rl = RateLimit::at_admit(lanes::Lane::Interactive, global, &metrics, &capped, 0);
assert_eq!(rl.remaining, 0);
let wide = auth::TenantCtx { rate_limit: Some(global + 100), ..capped.clone() };
let rl = RateLimit::at_admit(lanes::Lane::Interactive, 1, &metrics, &wide, 1);
assert_eq!((rl.limit, rl.remaining), (global, global - 1));
}
#[test]
fn batch_class_keys_default_to_harvest_and_cannot_claim_interactive() {
let batch = auth::TenantCtx {
tenant: "bulk".into(),
lane_class: auth::LaneClass::Batch,
rate_limit: None,
};
let interactive = auth::TenantCtx::default_tenant();
let hdr = |v: Option<&str>| {
let mut h = axum::http::HeaderMap::new();
if let Some(v) = v {
h.insert("x-lane", axum::http::HeaderValue::from_str(v).unwrap());
}
h
};
assert_eq!(lane_for_tenant(&hdr(None), &interactive).unwrap(),
lanes::Lane::Interactive);
assert_eq!(lane_for_tenant(&hdr(Some("judge")), &interactive).unwrap(),
lanes::Lane::Judge);
assert_eq!(lane_for_tenant(&hdr(None), &batch).unwrap(), lanes::Lane::Harvest);
assert_eq!(lane_for_tenant(&hdr(Some("judge")), &batch).unwrap(),
lanes::Lane::Judge);
let resp = lane_for_tenant(&hdr(Some("interactive")), &batch).unwrap_err();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
let resp = lane_for_tenant(&hdr(Some("turbo")), &interactive).unwrap_err();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
static DRAIN_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[tokio::test]
async fn responses_carry_rate_limit_headers_and_slot_frees() {
let _l = DRAIN_LOCK.lock().unwrap();
let st = fake_worker_state();
let resp = chat_completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
let h = resp.headers();
let limit: usize = h["x-ratelimit-limit"].to_str().unwrap().parse().unwrap();
let remaining: usize = h["x-ratelimit-remaining"].to_str().unwrap().parse().unwrap();
assert_eq!(remaining, limit - 1);
assert_eq!(h["x-ratelimit-reset"], "0");
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 0,
"slot must free at completion");
let resp = completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "t", "stream": true
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().contains_key("x-ratelimit-limit"));
assert!(resp.headers().contains_key("x-ratelimit-remaining"));
assert!(resp.headers().contains_key("x-ratelimit-reset"));
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 1,
"stream in flight holds the slot");
let _ = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 0,
"slot must free when the stream completes");
}
#[tokio::test]
async fn draining_rejects_new_requests_with_503_and_retry_after() {
let _l = DRAIN_LOCK.lock().unwrap();
let st = fake_worker_state();
DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
let resp = chat_completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert!(resp.headers().contains_key("retry-after"));
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert!(payload["error"]["message"].as_str().unwrap().contains("draining"));
let resp = completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "t"
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert!(resp.headers().contains_key("retry-after"));
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 0,
"rejected requests must not hold slots");
let resp = health(State(st.clone())).await.into_response();
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["status"], "draining");
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let resp = chat_completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
}
#[test]
fn v1_models_entry_matches_or_schema_with_honest_nulls() {
let caps = ModelCaps {
tools_branch: true, qwen_think: true, think_switch: true, chat_ok: true,
context_length: 262144,
tokenizer: "qwen2".into(),
instruct_type: Some("chatml".into()),
};
let e = model_entry_v1("main", Some(&caps), 1_754_000_000);
assert_eq!(e["id"], "main");
assert_eq!(e["name"], "main");
assert_eq!(e["object"], "model");
assert_eq!(e["created"], 1_754_000_000u64);
assert_eq!(e["context_length"], 262144);
assert_eq!(e["architecture"]["modality"], "text->text");
assert_eq!(e["architecture"]["tokenizer"], "qwen2");
assert_eq!(e["architecture"]["instruct_type"], "chatml");
assert_eq!(e["pricing"]["prompt"], "0");
assert_eq!(e["pricing"]["completion"], "0");
assert_eq!(e["top_provider"]["context_length"], 262144);
assert!(e["top_provider"]["max_completion_tokens"].is_null());
let e = model_entry_v1("m", None, 7);
assert!(e["context_length"].is_null());
assert!(e["architecture"]["tokenizer"].is_null());
assert!(e["architecture"]["instruct_type"].is_null());
assert!(e["top_provider"]["context_length"].is_null());
let bare = ModelCaps::default(); let e = model_entry_v1("m", Some(&bare), 7);
assert!(e["context_length"].is_null());
assert!(e["architecture"]["tokenizer"].is_null());
assert!(e["architecture"]["instruct_type"].is_null());
}
#[tokio::test]
async fn blocking_response_excludes_stop_text_across_token_events() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "answer\nPro".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "blem: leaked prompt".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Callback".into(), n_tokens: 2, n_prompt: 8, n_cached: 0, elapsed_s: 0.5,
spec: None,
}).unwrap();
drop(tx);
let response = blocking_response(
rx, "plain_quant".into(), false, vec!["Problem:".into()], None, Envelope::new(false)
).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["text"], "answer\n");
assert_eq!(payload["stop_reason"], "Callback");
}
}