mod admin;
pub(crate) mod auth;
pub(crate) mod constrained;
pub(crate) mod darklane;
pub(crate) mod health;
pub(crate) mod lanes { pub use memra_lanes::*; }
mod ledger;
mod toolcall;
mod ttft;
mod worker;
use std::collections::HashMap;
use std::net::ToSocketAddrs;
use std::sync::Arc;
use std::sync::mpsc::Sender;
use axum::{
body::Body,
Extension, Json, Router,
extract::{Query, Request as AxumRequest, State},
http::{header::CONTENT_TYPE, HeaderMap, StatusCode},
middleware::{self, Next},
response::{sse::{Event as SseEvent, Sse}, IntoResponse, Response},
routing::{get, post},
};
use futures_core::Stream as _;
use serde::{Deserialize, Serialize};
use serde_json::json;
use memra_engine::decode::GenParams;
use memra_engine::sampler::SamplerConfig;
use memra_tokenizer::{
Tokenizer,
chat::{ThinkMode, ToolCall as TmplToolCall, Turn as TmplTurn},
};
use toolcall::{ParsedToolCall, Piece, ToolStreamParser};
use worker::{Cmd, Event, ModelCaps, Request, SharedMetrics};
#[derive(Clone, Default)]
struct TtftRequestTrace(Option<Arc<ttft::Trace>>);
fn is_sse_data_frame(bytes: &[u8]) -> bool {
bytes
.windows(b"data:".len())
.any(|window| window == b"data:")
}
async fn ttft_request_start(mut req: AxumRequest, next: Next) -> Response {
let trace = ttft::start(req.uri().path());
req.extensions_mut().insert(TtftRequestTrace(trace.clone()));
let response = next.run(req).await;
let Some(trace) = trace else {
return response;
};
let is_sse = response
.headers()
.get(CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.starts_with("text/event-stream"));
if !is_sse {
return response;
}
let (parts, body) = response.into_parts();
let mut body = Box::pin(body.into_data_stream());
let stream = async_stream::stream! {
while let Some(frame) =
std::future::poll_fn(|cx| body.as_mut().poll_next(cx)).await
{
if frame
.as_ref()
.is_ok_and(|bytes| is_sse_data_frame(bytes))
{
trace.mark_first_sse_byte();
}
yield frame;
}
};
Response::from_parts(parts, Body::from_stream(stream))
}
const OPENROUTER_SCHEMA_VERSION: &str = "2.4";
const JSON_SAFE_INTEGER_MAX: u64 = 9_007_199_254_740_991;
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
struct OpenRouterMetadataFile {
#[serde(default)]
models: HashMap<String, OpenRouterModelMetadata>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
struct OpenRouterModelMetadata {
#[serde(default)]
hugging_face_id: Option<String>,
#[serde(default)]
created: Option<u64>,
#[serde(default)]
quantization: Option<String>,
#[serde(default)]
description: Option<String>,
#[serde(default)]
max_prompt_length: Option<u64>,
#[serde(default)]
max_output_length: Option<u64>,
#[serde(default)]
pricing: OpenRouterPricing,
#[serde(default)]
capacity: OpenRouterCapacity,
#[serde(default)]
is_ready: Option<bool>,
#[serde(default)]
is_free: Option<bool>,
#[serde(default)]
discount_to_user: Option<f64>,
#[serde(default)]
openrouter_slug: Option<String>,
#[serde(default)]
datacenters: Vec<OpenRouterDatacenter>,
#[serde(default)]
zdr: Option<bool>,
#[serde(default)]
hipaa: Option<bool>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
struct OpenRouterPricing {
#[serde(default)]
prompt: Option<String>,
#[serde(default)]
cached_prompt: Option<String>,
#[serde(default)]
cache_write: Option<String>,
#[serde(default)]
completion: Option<String>,
#[serde(default)]
internal_reasoning: Option<String>,
#[serde(default)]
request: Option<String>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
struct OpenRouterCapacity {
#[serde(default)]
prompt_tpm: Option<u64>,
#[serde(default)]
cached_prompt_tpm: Option<u64>,
#[serde(default)]
completion_tpm: Option<u64>,
#[serde(default)]
request_rpm: Option<u64>,
#[serde(default)]
concurrency: Option<u64>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
struct OpenRouterDatacenter {
country_code: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
region: Option<String>,
}
impl OpenRouterMetadataFile {
fn from_toml(text: &str) -> Result<HashMap<String, OpenRouterModelMetadata>, String> {
let file: Self = toml::from_str(text)
.map_err(|e| format!("models metadata TOML parse: {e}"))?;
for (alias, metadata) in &file.models {
validate_openrouter_metadata(alias, metadata)?;
}
Ok(file.models)
}
}
fn valid_price_string(value: &str) -> bool {
let mut parts = value.split('.');
let whole = parts.next().unwrap_or_default();
let fraction = parts.next();
!whole.is_empty()
&& whole.bytes().all(|b| b.is_ascii_digit())
&& fraction.is_none_or(|v| !v.is_empty() && v.bytes().all(|b| b.is_ascii_digit()))
&& parts.next().is_none()
}
fn validate_openrouter_metadata(
alias: &str,
metadata: &OpenRouterModelMetadata,
) -> Result<(), String> {
if alias.is_empty() {
return Err("models metadata contains an empty model alias".into());
}
if let Some(q) = metadata.quantization.as_deref()
&& !matches!(
q,
"int4" | "int8" | "fp4" | "mxfp4" | "nvfp4" | "fp6" | "fp8" | "mxfp8"
| "fp16" | "bf16" | "fp32"
)
{
return Err(format!(
"model {alias:?}: quantization {q:?} is not in the OpenRouter schema 2.4 enum"
));
}
for (field, value) in [
("pricing.prompt", metadata.pricing.prompt.as_deref()),
(
"pricing.cached_prompt",
metadata.pricing.cached_prompt.as_deref(),
),
(
"pricing.cache_write",
metadata.pricing.cache_write.as_deref(),
),
(
"pricing.completion",
metadata.pricing.completion.as_deref(),
),
(
"pricing.internal_reasoning",
metadata.pricing.internal_reasoning.as_deref(),
),
("pricing.request", metadata.pricing.request.as_deref()),
] {
if let Some(value) = value
&& !valid_price_string(value)
{
return Err(format!(
"model {alias:?}: {field} must be a non-negative per-unit USD decimal string"
));
}
}
for (field, value) in [
("created", metadata.created),
("max_prompt_length", metadata.max_prompt_length),
("max_output_length", metadata.max_output_length),
("capacity.prompt_tpm", metadata.capacity.prompt_tpm),
(
"capacity.cached_prompt_tpm",
metadata.capacity.cached_prompt_tpm,
),
(
"capacity.completion_tpm",
metadata.capacity.completion_tpm,
),
("capacity.request_rpm", metadata.capacity.request_rpm),
("capacity.concurrency", metadata.capacity.concurrency),
] {
if let Some(value) = value
&& value > JSON_SAFE_INTEGER_MAX
{
return Err(format!(
"model {alias:?}: {field} exceeds OpenRouter's JSON safe-integer maximum"
));
}
}
for (field, value) in [
("max_prompt_length", metadata.max_prompt_length),
("max_output_length", metadata.max_output_length),
("capacity.prompt_tpm", metadata.capacity.prompt_tpm),
(
"capacity.cached_prompt_tpm",
metadata.capacity.cached_prompt_tpm,
),
(
"capacity.completion_tpm",
metadata.capacity.completion_tpm,
),
("capacity.request_rpm", metadata.capacity.request_rpm),
("capacity.concurrency", metadata.capacity.concurrency),
] {
if value == Some(0) {
return Err(format!(
"model {alias:?}: {field} must be greater than zero when declared"
));
}
}
if let Some(discount) = metadata.discount_to_user
&& (!discount.is_finite() || discount >= 1.0)
{
return Err(format!(
"model {alias:?}: discount_to_user must be finite and less than 1"
));
}
if metadata
.openrouter_slug
.as_deref()
.is_some_and(str::is_empty)
{
return Err(format!(
"model {alias:?}: openrouter_slug must not be empty when declared"
));
}
for dc in &metadata.datacenters {
if dc.country_code.len() != 2
|| !dc.country_code.bytes().all(|b| b.is_ascii_uppercase())
{
return Err(format!(
"model {alias:?}: datacenter country_code {:?} must be two uppercase ASCII letters",
dc.country_code
));
}
}
Ok(())
}
fn load_openrouter_metadata(
models: &[(String, String, Option<String>)],
) -> Result<HashMap<String, OpenRouterModelMetadata>, String> {
let path = match std::env::var("MEMRA_MODEL_METADATA") {
Ok(path) => path,
Err(_) => return Ok(HashMap::new()),
};
let p = std::path::Path::new(&path);
if !p.is_file() {
return Err(format!(
"MEMRA_MODEL_METADATA={path:?} is not an existing TOML file"
));
}
let text = std::fs::read_to_string(p)
.map_err(|e| format!("MEMRA_MODEL_METADATA {path:?}: {e}"))?;
let metadata = OpenRouterMetadataFile::from_toml(&text)
.map_err(|e| format!("MEMRA_MODEL_METADATA {path:?}: {e}"))?;
for alias in metadata.keys() {
if !models.iter().any(|(name, _, _)| name == alias) {
return Err(format!(
"MEMRA_MODEL_METADATA {path:?}: model alias {alias:?} is not present in MEMRA_MODELS"
));
}
}
eprintln!(
"[server] OpenRouter metadata loaded: {} model(s) from {path}",
metadata.len()
);
Ok(metadata)
}
#[derive(Clone)]
struct AppState {
cmd_tx: Sender<Cmd>,
models: Arc<Vec<String>>,
caps: Arc<HashMap<String, ModelCaps>>,
openrouter_metadata: Arc<HashMap<String, OpenRouterModelMetadata>>,
request_ledger: Option<ledger::Ledger>,
tenant_budgets: Option<ledger::TenantBudgets>,
budget_tokenizers: Option<Arc<HashMap<String, Arc<Tokenizer>>>>,
api_auth: ApiAuth,
metrics_auth: MetricsAuth,
metrics: SharedMetrics,
started: u64,
inflight: InflightCounts,
tenant_inflight: TenantGauge,
health: health::SharedHealth,
bg: Option<(Arc<darklane::BgJobState>, &'static str)>,
}
#[derive(Clone, Default)]
struct ApiAuth {
keyring: Option<&'static auth::KeyStore>,
single_key: Option<Arc<str>>,
}
impl ApiAuth {
fn from_env() -> Result<ApiAuth, String> {
let single_key = match std::env::var("MEMRA_API_KEY") {
Ok(key) if key.is_empty() => return Err("MEMRA_API_KEY must not be empty".into()),
Ok(key) => Some(Arc::from(key)),
Err(std::env::VarError::NotPresent) => None,
Err(std::env::VarError::NotUnicode(_)) => {
return Err("MEMRA_API_KEY must be valid UTF-8".into());
}
};
Ok(ApiAuth {
keyring: auth::global(),
single_key,
})
}
fn configured(&self) -> bool {
self.keyring.is_some() || self.single_key.is_some()
}
}
#[derive(Clone, Default)]
struct MetricsAuth {
required: bool,
token: Option<Arc<str>>,
}
impl MetricsAuth {
fn new(bind_loopback: bool, api_auth_configured: bool, token: Option<String>) -> MetricsAuth {
let token = token.map(Arc::from);
MetricsAuth {
required: !bind_loopback || api_auth_configured || token.is_some(),
token,
}
}
}
fn bind_is_loopback(addr: &str) -> Result<bool, String> {
let mut resolved = addr.to_socket_addrs()
.map_err(|e| format!("MEMRA_ADDR={addr:?} cannot be resolved: {e}"))?;
let first = resolved.next()
.ok_or_else(|| format!("MEMRA_ADDR={addr:?} resolved to no socket addresses"))?;
let mut loopback = first.ip().to_canonical().is_loopback();
for socket in resolved {
loopback &= socket.ip().to_canonical().is_loopback();
}
Ok(loopback)
}
fn validate_bind_security(
addr: &str,
api_auth_configured: bool,
allow_open_bind: bool,
) -> Result<bool, String> {
let loopback = bind_is_loopback(addr)?;
if !loopback && !api_auth_configured && !allow_open_bind {
return Err(format!(
"refusing unauthenticated non-loopback bind {addr:?}; configure MEMRA_API_KEY or \
MEMRA_API_KEYS, or set MEMRA_ALLOW_OPEN_BIND=1 for an explicit development override"
));
}
Ok(loopback)
}
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 try_acquire(counts: InflightCounts, lane: lanes::Lane, tenants: TenantGauge,
tenant: &str, tenant_cap: Option<usize>)
-> Result<(Self, usize, usize), usize>
{
let idx = lane.idx();
let nt = {
let mut m = tenants.lock().unwrap();
let e = m.entry(tenant.to_string()).or_insert(0);
if tenant_cap.is_some_and(|cap| *e >= cap) {
return Err(*e);
}
*e += 1;
*e
};
let n = counts[idx].fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
Ok((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 resp = (StatusCode::SERVICE_UNAVAILABLE,
Json(error_body("server is draining (shutdown in progress); retry",
"server_error", None, Some("draining")))).into_response();
retry_contract_response(resp, Some(drain_deadline_s()))
}
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
}
}
fn acquire_request_slot(st: &AppState, lane: lanes::Lane, tenant: &auth::TenantCtx,
env: &Envelope) -> Result<(InflightGuard, RateLimit), Response> {
let global = lane_cap(lane);
let tenant_cap = tenant.rate_limit.filter(|&cap| cap < global);
match InflightGuard::try_acquire(
st.inflight.clone(), lane, st.tenant_inflight.clone(), &tenant.tenant, tenant_cap,
) {
Ok((guard, n_inflight, n_tenant)) => {
let rl = RateLimit::at_admit(lane, n_inflight, &st.metrics, tenant, n_tenant);
Ok((guard, rl))
}
Err(n_tenant) => {
let n_inflight =
st.inflight[lane.idx()].load(std::sync::atomic::Ordering::SeqCst);
let rl = RateLimit::at_admit(
lane, n_inflight, &st.metrics, tenant, n_tenant);
let error = worker::EngineError::rate_limit(
"api key concurrent request limit reached; retry");
Err(rl.attach(with_request_id(&env.id, engine_error_response(&error))))
}
}
}
#[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)]
temperature: Option<f32>,
#[serde(default)]
top_p: Option<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 }
fn chat_sampling_values(req: &ChatCompletionReq, caps: Option<&ModelCaps>) -> (f32, f32) {
let temperature = req.temperature
.or_else(|| caps.and_then(|c| c.chat_temperature_default))
.unwrap_or_else(default_temperature);
let top_p = req.top_p
.or_else(|| caps.and_then(|c| c.chat_top_p_default))
.unwrap_or_else(one);
(temperature, top_p)
}
#[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()
}
const CACHE_SALT_MAX_BYTES: usize = 64;
fn validate_cache_namespace(
cache_salt: &Option<String>,
keyring_configured: bool,
) -> Result<String, &'static str> {
let raw = cache_namespace(cache_salt);
if raw.len() > CACHE_SALT_MAX_BYTES {
return Err("cache_salt must be at most 64 bytes");
}
if !keyring_configured && raw.starts_with("t:") {
return Err("cache_salt must not use the reserved t: prefix without a keyring");
}
if !raw.bytes().all(|b| {
b.is_ascii_alphanumeric() || matches!(b, b'-' | b'_' | b'.' | b'+' | b'/' | b'=')
}) {
return Err("cache_salt contains unsupported characters");
}
Ok(raw)
}
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 {
error_response_coded(status, message, etype, param, None)
}
fn error_response_coded(status: StatusCode, message: &str, etype: &str,
param: Option<&str>, code: Option<&str>) -> Response {
let mut resp = (status, Json(error_body(message, etype, param, code))).into_response();
if status.is_client_error() && status != StatusCode::TOO_MANY_REQUESTS
&& status != StatusCode::REQUEST_TIMEOUT && status != StatusCode::CONFLICT {
resp.headers_mut()
.insert("x-should-retry", axum::http::HeaderValue::from_static("false"));
}
resp
}
fn bad_request(message: &str, param: Option<&str>) -> Response {
error_response(StatusCode::BAD_REQUEST, message, "invalid_request_error", param)
}
const RETRY_AFTER_S_RATE_LIMIT: u64 = 2; const RETRY_AFTER_S_OVERLOADED: u64 = 5;
fn class_http(class: worker::ErrClass) -> (StatusCode, &'static str, Option<&'static str>) {
use worker::ErrClass as C;
match class {
C::InvalidRequest => (StatusCode::BAD_REQUEST, "invalid_request_error", None),
C::ContextLength => (StatusCode::BAD_REQUEST, "invalid_request_error",
Some("context_length_exceeded")),
C::ModelNotFound => (StatusCode::BAD_REQUEST, "invalid_request_error",
Some("model_not_found")),
C::RateLimit => (StatusCode::TOO_MANY_REQUESTS, "rate_limit_error", Some("rate_limit_exceeded")),
C::Overloaded => (StatusCode::SERVICE_UNAVAILABLE, "server_error", Some("overloaded")),
C::Engine => (StatusCode::INTERNAL_SERVER_ERROR, "server_error", Some("engine_error")),
}
}
fn class_retry_after_s(class: worker::ErrClass) -> Option<u64> {
use worker::ErrClass as C;
match class {
C::RateLimit => Some(RETRY_AFTER_S_RATE_LIMIT),
C::Overloaded => Some(RETRY_AFTER_S_OVERLOADED),
C::Engine | C::InvalidRequest | C::ContextLength | C::ModelNotFound => None,
}
}
fn engine_error_body(e: &worker::EngineError) -> serde_json::Value {
let (_, etype, code) = class_http(e.class);
error_body(&e.message, etype, e.param, code)
}
fn engine_error_response(e: &worker::EngineError) -> Response {
engine_error_response_with_retry_after(e, class_retry_after_s(e.class))
}
fn engine_error_response_with_retry_after(
e: &worker::EngineError,
retry_after_s: Option<u64>,
) -> Response {
let (status, _, _) = class_http(e.class);
let resp = (status, Json(engine_error_body(e))).into_response();
retry_contract_response(resp, retry_after_s)
}
fn retry_contract_response(mut resp: Response, retry_after_s: Option<u64>) -> Response {
let status = resp.status();
let h = resp.headers_mut();
match retry_after_s {
Some(secs) => {
let secs = secs.clamp(1, 60);
if let Ok(v) = axum::http::HeaderValue::from_str(&secs.to_string()) {
h.insert(axum::http::header::RETRY_AFTER, v);
}
if let Ok(v) = axum::http::HeaderValue::from_str(&(secs * 1000).to_string()) {
h.insert("retry-after-ms", v);
}
}
None if status.is_client_error() => {
h.insert("x-should-retry", axum::http::HeaderValue::from_static("false"));
}
None => {}
}
resp
}
fn worker_unavailable_response() -> Response {
engine_error_response_with_retry_after(
&worker::EngineError::overloaded("worker unavailable"),
Some(worker::WORKER_RESPAWN_BACKOFF_BASE_S),
)
}
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, Option<String>), String> {
let mut effort = reasoning_effort.clone();
let mut enabled = None;
if let Some(r) = reasoning {
match r {
serde_json::Value::Null => {}
serde_json::Value::Object(obj) => {
enabled = obj.get("enabled").and_then(|v| v.as_bool());
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()),
}
}
if enabled == Some(false) {
return Ok((ThinkMode::NoThink, Some("low".to_string())));
}
let (think, level) = match effort.as_deref() {
None => (
if enabled == Some(true) { ThinkMode::Think } else { ThinkMode::Default },
None,
),
Some("none") | Some("minimal") => (ThinkMode::NoThink, Some("low".to_string())),
Some("low") => (ThinkMode::Think, Some("low".to_string())),
Some("medium") => (ThinkMode::Think, Some("medium".to_string())),
Some("high") => (ThinkMode::Think, Some("high".to_string())),
Some(other) => return Err(format!(
"bad reasoning_effort {other:?} (none|minimal|low|medium|high)")),
};
Ok((think, level))
}
#[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 api_auth = match ApiAuth::from_env() {
Ok(auth) => auth,
Err(err) => {
eprintln!("[server] FATAL: {err}");
std::process::exit(1);
}
};
let addr = std::env::var("MEMRA_ADDR").unwrap_or_else(|_| "127.0.0.1:8080".into());
let allow_open_bind = std::env::var("MEMRA_ALLOW_OPEN_BIND").as_deref() == Ok("1");
let bind_loopback = match validate_bind_security(
&addr, api_auth.configured(), allow_open_bind)
{
Ok(loopback) => loopback,
Err(err) => {
eprintln!("[server] FATAL: {err}");
std::process::exit(1);
}
};
if !bind_loopback && !api_auth.configured() {
eprintln!(
"[server] WARNING: MEMRA_ALLOW_OPEN_BIND=1 permits open completion routes on {addr}; \
metrics remain bearer-protected"
);
}
let metrics_token = match std::env::var("MEMRA_METRICS_TOKEN") {
Ok(token) if token.is_empty() => {
eprintln!("[server] FATAL: MEMRA_METRICS_TOKEN must not be empty");
std::process::exit(1);
}
Ok(token) => Some(token),
Err(std::env::VarError::NotPresent) => None,
Err(std::env::VarError::NotUnicode(_)) => {
eprintln!("[server] FATAL: MEMRA_METRICS_TOKEN must be valid UTF-8");
std::process::exit(1);
}
};
let metrics_auth = MetricsAuth::new(bind_loopback, api_auth.configured(), metrics_token);
let models = parse_models_config();
let openrouter_metadata = match load_openrouter_metadata(&models) {
Ok(metadata) => metadata,
Err(err) => {
eprintln!("[server] FATAL: {err}");
std::process::exit(1);
}
};
let request_ledger = match ledger::Ledger::from_env(&models, &openrouter_metadata) {
Ok(ledger) => ledger,
Err(err) => {
eprintln!("[server] FATAL: MEMRA_REQUEST_LEDGER: {err}");
std::process::exit(1);
}
};
let tenant_budgets = request_ledger.as_ref().and_then(ledger::Ledger::budgets);
let budget_tokenizers = if tenant_budgets.is_some() {
match load_budget_tokenizers(&models) {
Ok(tokenizers) => Some(tokenizers),
Err(err) => {
eprintln!("[server] FATAL: prepaid reservation tokenizers: {err}");
std::process::exit(1);
}
}
} else {
None
};
let admin_config = match admin::Config::from_env(request_ledger.as_ref(), auth::global()) {
Ok(config) => config,
Err(err) => {
eprintln!("[server] FATAL: admin configuration: {err}");
std::process::exit(1);
}
};
let admin_listener = if let Some(config) = admin_config.as_ref() {
match tokio::net::TcpListener::bind(config.addr()).await {
Ok(listener) => Some(listener),
Err(err) => {
eprintln!(
"[server] FATAL: bind MEMRA_ADMIN_ADDR {:?}: {err}",
config.addr(),
);
std::process::exit(1);
}
}
} else {
None
};
eprintln!("[server] starting; models config = {models:?}");
let health_state = health::WorkerHealth::new();
health::spawn_gpu_watch(health_state.clone());
health::spawn_sd_watchdog(health_state.clone());
let (cmd_tx, model_names, caps, metrics, worker_thread) =
match worker::spawn(models, health_state.clone()) {
Ok(v) => v,
Err(err) => {
eprintln!("[server] FATAL: worker init failed: {err}");
health_state.mark_dead(format!("worker init failed: {err}"));
health::sd_notify(&format!("STATUS=worker init failed: {err}"));
std::process::exit(1);
}
};
eprintln!("[server] worker ready; serving models: {model_names:?}");
let bg_handle = darklane::spawn_from_env(health_state.clone());
let bg_state = bg_handle.as_ref().map(|h| {
let mode = darklane::BgConfig::from_env()
.map(|c| c.yield_mode.as_str()).unwrap_or("stop");
(h.state.clone(), mode)
});
let state = AppState {
cmd_tx, models: model_names, caps,
openrouter_metadata: Arc::new(openrouter_metadata),
request_ledger,
tenant_budgets,
budget_tokenizers,
api_auth,
metrics_auth,
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()),
health: health_state.clone(),
bg: bg_state,
};
let inflight_handle = state.inflight.clone();
let app = Router::new()
.route("/health", get(health_live))
.route("/livez", get(health_live))
.route("/readyz", get(health_ready))
.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 app = if ttft::enabled() {
app.layer(middleware::from_fn(ttft_request_start))
} else {
app
};
let listener = tokio::net::TcpListener::bind(&addr).await?;
eprintln!("[server] listening on http://{addr}");
let (admin_shutdown_tx, admin_shutdown_rx) = tokio::sync::watch::channel(false);
let admin_task = match (admin_config, admin_listener) {
(Some(config), Some(listener)) => {
let admin_addr = config.addr().to_string();
let admin_app = config.router();
eprintln!("[admin] listening on http://{admin_addr}");
let mut shutdown = admin_shutdown_rx;
Some(tokio::spawn(async move {
axum::serve(listener, admin_app)
.with_graceful_shutdown(async move {
while !*shutdown.borrow() {
if shutdown.changed().await.is_err() {
break;
}
}
})
.await
}))
}
(None, None) => None,
_ => unreachable!("admin config and pre-bound listener are created together"),
};
health::sd_notify("READY=1\nSTATUS=serving");
let inflight = inflight_handle;
let signal_admin_shutdown = admin_shutdown_tx.clone();
let serve_result = 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 _ = signal_admin_shutdown.send(true);
health::sd_notify(&format!(
"STOPPING=1\nSTATUS=draining\nEXTEND_TIMEOUT_USEC={}",
(drain_deadline_s() + 5) * 1_000_000));
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;
let _ = admin_shutdown_tx.send(true);
if let Some(task) = admin_task {
match task.await {
Ok(Ok(())) => {}
Ok(Err(err)) => return Err(format!("admin listener failed: {err}").into()),
Err(err) => return Err(format!("admin listener task failed: {err}").into()),
}
}
serve_result?;
if let Some(h) = bg_handle {
h.shutdown();
}
worker_thread.join().map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::Other,
"GPU worker thread panicked during graceful shutdown",
)
})?;
eprintln!("[server] GPU worker shutdown complete");
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);
}
let dpath = dpath.map(|d| {
let d = resolve(d);
let p = std::path::Path::new(&d);
if !p.exists() {
eprintln!("[server] FATAL: model {name:?}: drafter path {d:?} does not \
exist (MEMRA_MODELS '+draft' attach). Refusing to start \
rather than serving plain decode under a config that asked \
for speculative decoding.");
std::process::exit(1);
}
if !p.is_file() {
eprintln!("[server] FATAL: model {name:?}: drafter path {d:?} is not a \
file — a '+draft' attach must be a NextN/MTP GGUF file.");
std::process::exit(1);
}
d
});
out.push((name.trim().to_string(), mpath, dpath));
} 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),
]
}
fn load_budget_tokenizers(
models: &[(String, String, Option<String>)],
) -> Result<Arc<HashMap<String, Arc<Tokenizer>>>, String> {
let mut tokenizers = HashMap::new();
for (alias, path, _) in models {
let path = std::path::Path::new(path);
let tokenizer = if path.is_dir() {
let tokenizer_dir = if path.join("manifest.json").exists() {
let repack = memra_gguf::source::Hy3RepackSource::open(path).map_err(|err| {
format!("model {alias:?}: open repack tokenizer source: {err}")
})?;
repack
.source_dir()
.filter(|source| source.join("tokenizer.json").exists())
.unwrap_or(path)
.to_path_buf()
} else {
path.to_path_buf()
};
Tokenizer::from_hf_dir(&tokenizer_dir)
.map_err(|err| format!("model {alias:?}: reservation tokenizer: {err}"))?
} else {
let gguf = memra_gguf::GgufFile::open(path)
.map_err(|err| format!("model {alias:?}: open reservation tokenizer: {err}"))?;
Tokenizer::from_gguf(&gguf)
.map_err(|err| format!("model {alias:?}: reservation tokenizer: {err}"))?
};
tokenizers.insert(alias.clone(), Arc::new(tokenizer));
}
Ok(Arc::new(tokenizers))
}
fn health_payload(st: &AppState, status: &str, detail: Option<&str>) -> serde_json::Value {
let s = st.health.snapshot();
let mut v = json!({
"status": status,
"models": *st.models,
"worker": {
"phase": health::phase_name(s.phase),
"beat_age_ms": s.beat_age_ms,
"tick_max_ms": s.tick_max_ms,
"stall_threshold_ms": s.stall_threshold_ms,
"generation": s.generation,
"xid_warnings": s.xid_warns,
},
});
if let Some(d) = detail {
v["detail"] = json!(d);
}
v
}
fn readiness_payload(st: &AppState, status: &str, detail: Option<&str>) -> serde_json::Value {
let mut v = health_payload(st, status, detail);
v["peer_probe_integrity"] = json!(st.health.peer_probe_integrity().detail());
v
}
async fn health_live(State(st): State<AppState>) -> impl IntoResponse {
if draining() {
return (StatusCode::OK, Json(health_payload(&st, "draining", None))).into_response();
}
match st.health.live() {
Ok(()) => (StatusCode::OK, Json(health_payload(&st, "ok", None))).into_response(),
Err(why) => retry_contract_response(
(StatusCode::SERVICE_UNAVAILABLE,
Json(health_payload(&st, "unhealthy", Some(&why)))).into_response(),
Some(worker::WORKER_RESPAWN_BACKOFF_BASE_S),
),
}
}
async fn health_ready(State(st): State<AppState>) -> impl IntoResponse {
let is_draining = draining();
match st.health.ready(is_draining) {
Ok(()) => (StatusCode::OK, Json(readiness_payload(&st, "ready", None))).into_response(),
Err(why) => retry_contract_response(
(StatusCode::SERVICE_UNAVAILABLE,
Json(readiness_payload(&st, "not_ready", Some(&why)))).into_response(),
Some(if is_draining {
drain_deadline_s()
} else {
worker::WORKER_RESPAWN_BACKOFF_BASE_S
}),
),
}
}
#[derive(Clone, Copy)]
struct DualPpMetricsSnapshot {
stage_ns: [u64; 4],
stage_samples: [usize; 4],
dropped_timing_samples: usize,
overlaps: usize,
slot_pairs: usize,
slot_uses: [usize; 2],
slot_collisions: usize,
}
impl DualPpMetricsSnapshot {
fn current() -> Self {
let (stage_ns, stage_samples) = memra_engine::pp::dual_pp_timing_snapshot();
let (slot_pairs, slot_uses, slot_collisions) =
memra_engine::pp::dual_pp_slot_snapshot();
Self {
stage_ns,
stage_samples,
dropped_timing_samples: memra_engine::pp::dual_pp_timing_dropped(),
overlaps: memra_engine::pp::dual_pp_overlaps(),
slot_pairs,
slot_uses,
slot_collisions,
}
}
fn populated(self) -> bool {
self.stage_samples.iter().any(|&n| n > 0)
|| self.dropped_timing_samples > 0
|| self.slot_pairs > 0
|| self.slot_collisions > 0
}
}
fn insert_dual_pp_metrics(
body: &mut serde_json::Value,
metrics_scope: &MetricsScope,
snapshot: impl FnOnce() -> DualPpMetricsSnapshot,
) {
if !metrics_scope.operator() {
return;
}
let snapshot = snapshot();
if !snapshot.populated() {
return;
}
let timings: serde_json::Map<String, serde_json::Value> =
memra_engine::pp::DUAL_PP_STAGE_NAMES.iter().enumerate().map(|(i, name)| {
let total_ms = snapshot.stage_ns[i] as f64 / 1_000_000.0;
(name.to_string(), json!({
"samples": snapshot.stage_samples[i],
"total_ms": total_ms,
"mean_ms": if snapshot.stage_samples[i] > 0 {
total_ms / snapshot.stage_samples[i] as f64
} else { 0.0 },
}))
}).collect();
body["dual_pp"] = json!({
"overlaps": snapshot.overlaps,
"slot_pairs": snapshot.slot_pairs,
"slot_uses": snapshot.slot_uses,
"slot_collisions": snapshot.slot_collisions,
"cuda_event_spans": timings,
"dropped_timing_samples": snapshot.dropped_timing_samples,
});
}
fn insert_spec_acceptance_metrics(
body: &mut serde_json::Value,
metrics_scope: &MetricsScope,
snapshot: impl FnOnce() -> HashMap<String, memra_engine::spec::SpecTelemetry>,
) {
if !metrics_scope.operator() {
return;
}
let snapshot = snapshot();
if snapshot.is_empty() {
return;
}
let mut tau = serde_json::Map::new();
let mut by_position = serde_json::Map::new();
for (model, telemetry) in snapshot {
if telemetry.rounds == 0 {
continue;
}
let n_pos = telemetry.pos_drafted.iter().rposition(|&n| n > 0)
.map_or(0, |position| position + 1);
tau.insert(model.clone(), json!(telemetry.tau()));
by_position.insert(model, json!({
"window_seconds": worker::SPEC_METRICS_WINDOW_S,
"rounds": telemetry.rounds,
"offered": telemetry.pos_drafted[..n_pos].to_vec(),
"accepted": telemetry.pos_accepted[..n_pos].to_vec(),
"accept_rate": (0..n_pos).map(|position| {
let offered = telemetry.pos_drafted[position];
if offered > 0 {
telemetry.pos_accepted[position] as f64 / offered as f64
} else {
0.0
}
}).collect::<Vec<f64>>(),
}));
}
if !tau.is_empty() {
body["spec_tau"] = serde_json::Value::Object(tau);
body["spec_accept_by_position"] = serde_json::Value::Object(by_position);
}
}
fn insert_peer_probe_metrics(
body: &mut serde_json::Value,
metrics_scope: &MetricsScope,
snapshot: impl FnOnce() -> memra_engine::pp::PeerProbeMetrics,
) {
if !metrics_scope.operator() {
return;
}
let snapshot = snapshot();
body["peer_probe_bypassed"] = json!(snapshot.bypassed);
body["peer_probe_boundary_copies"] = json!(snapshot.boundary_copies);
body["peer_probe_runtime_reprobes"] = json!(snapshot.runtime_probes);
body["peer_probe_runtime_failures"] = json!(snapshot.runtime_failures);
body["peer_probe_deferred_total"] = json!(snapshot.deferred_total);
body["peer_probe_integrity_degraded"] = json!(snapshot.integrity_degraded);
body["peer_probe_degraded_to_host_bounce"] = json!(snapshot.degraded_to_host_bounce);
}
async fn get_metrics(State(st): State<AppState>, headers: HeaderMap) -> Response {
let metrics_scope = match authorize_metrics(&st.api_auth, &st.metrics_auth, &headers) {
Ok(scope) => scope,
Err(response) => return response,
};
let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
let mut body = if metrics_scope.process_wide() {
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,
"computed_tokens_in": m.prompt_tokens_in.saturating_sub(m.cached_tokens_in),
"admission_session_defers": m.admission_session_defers,
"admission_vram_defers": m.admission_vram_defers,
"step_oom_parks": m.step_oom_parks,
"continuation_pool_hits": m.continuation_pool_hits,
"continuation_pool_evictions": m.continuation_pool_evictions,
"plain_affinity_rewinds": m.plain_affinity_rewinds,
"spec_pool_hits": m.spec_pool_hits,
"spec_pool_misses": m.spec_pool_misses,
"spec_pool_affinity_rewinds": m.spec_pool_affinity_rewinds,
"spec_pool_evictions": m.spec_pool_evictions,
})
} else {
json!({})
};
if metrics_scope.operator() {
if let Some(budgets) = st.tenant_budgets.as_ref() {
let budget_health = budgets.health();
body["budget_source_reload_failed"] = json!(budget_health.source_reload_failed);
body["budget_source_reload_consecutive"] =
json!(budget_health.source_reload_consecutive);
body["budget_source_available"] = json!(budget_health.source_available);
}
body["cache_hit_token_ratio"] = json!(if m.prompt_tokens_in > 0 {
m.cached_tokens_in as f64 / m.prompt_tokens_in as f64 } else { 0.0 });
body["prefix_cache_hits"] = json!(m.prefix_hits);
body["prefix_cache_misses"] = json!(m.prefix_misses);
body["prefix_cache_inserts"] = json!(m.prefix_inserts);
body["prefix_cache_evictions"] = json!(m.prefix_evictions);
body["prefix_cache_hit_tokens"] = json!(m.prefix_hit_tokens);
body["lcp_histogram"] = json!({
"edges": worker::LCP_HIST_EDGES.to_vec(),
"counts": m.lcp_hist.to_vec(),
});
let idle_s = darklane::ValleySignal::new(st.health.clone()).idle_seconds();
body["prefix_cache_entries"] = json!(m.prefix_entries);
body["prefix_cache_bytes"] = json!(m.prefix_bytes);
body["active_sessions"] = json!(m.active_sessions);
body["queued_requests"] = json!(m.queued_requests);
body["continuation_pool_entries"] = json!(m.continuation_pool_entries);
body["spec_pool_entries"] = json!(m.spec_pool_entries);
body["cuda_driver_free_bytes"] = json!(m.cuda_driver_free_bytes);
body["cuda_pool_reserved_bytes"] = json!(m.cuda_pool_reserved_bytes);
body["cuda_pool_used_bytes"] = json!(m.cuda_pool_used_bytes);
body["cuda_pool_cached_bytes"] = json!(m.cuda_pool_cached_bytes);
if !m.constraint_compiler_fail_closed.is_empty() {
body["constraint_compiler_fail_closed"] = serde_json::Value::Object(
m.constraint_compiler_fail_closed.iter().map(|(model, gauge)| {
let value = u8::from(gauge.load(std::sync::atomic::Ordering::Acquire));
(model.clone(), json!(value))
}).collect(),
);
}
body["serve_idle_seconds"] = json!((idle_s * 1000.0).round() / 1000.0);
}
if !m.ns_tokens.is_empty() {
let tenants: serde_json::Map<String, serde_json::Value> = m.ns_tokens.iter()
.filter(|(ns, _)| metrics_scope.includes(ns))
.map(|(ns, [p, c])| (ns.clone(), json!({
"prompt_tokens_in": p,
"cached_tokens_in": c,
"cache_hit_token_ratio": if *p > 0 { *c as f64 / *p as f64 } else { 0.0 },
})))
.collect();
if !tenants.is_empty() {
body["tenants"] = serde_json::Value::Object(tenants);
}
}
let adsd_suspect_total: serde_json::Map<String, serde_json::Value> =
m.adsd_suspect_total.iter()
.filter(|(tenant, _)| metrics_scope.includes(tenant))
.map(|(tenant, total)| (tenant.clone(), json!(total)))
.collect();
if !adsd_suspect_total.is_empty() {
body["adsd_suspect_total"] = serde_json::Value::Object(adsd_suspect_total);
}
if metrics_scope.operator() {
if let Some((bg, mode)) = &st.bg {
body["bg"] = bg.to_json(mode);
}
}
if metrics_scope.operator() {
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);
}
}
insert_spec_acceptance_metrics(&mut body, &metrics_scope, || m.spec_window.clone());
insert_dual_pp_metrics(&mut body, &metrics_scope, DualPpMetricsSnapshot::current);
insert_peer_probe_metrics(&mut body, &metrics_scope, memra_engine::pp::peer_probe_metrics);
Json(body).into_response()
}
#[derive(Debug, Default, Deserialize)]
struct ModelsQuery {
#[serde(default)]
schema: Option<String>,
}
fn models_openai_body(models: &[String]) -> serde_json::Value {
let data: Vec<_> = models
.iter()
.map(|m| json!({ "id": m, "object": "model" }))
.collect();
json!({ "object": "list", "data": data })
}
fn openrouter_supported_parameters(
caps: Option<&ModelCaps>,
max_output_length: Option<u64>,
) -> serde_json::Value {
let mut parameters = serde_json::Map::new();
for name in [
"temperature",
"top_p",
"min_p",
"frequency_penalty",
"presence_penalty",
"repetition_penalty",
"stop",
] {
parameters.insert(name.into(), json!({ "type": "unknown" }));
}
parameters.insert("top_k".into(), json!({ "type": "integer", "min": 0 }));
parameters.insert(
"seed".into(),
json!({ "type": "integer", "min": 0, "max": JSON_SAFE_INTEGER_MAX }),
);
let mut max_tokens = json!({ "type": "integer", "min": 1, "unit": "token" });
if let Some(max) = max_output_length {
max_tokens["max"] = json!(max);
}
parameters.insert("max_tokens".into(), max_tokens);
parameters.insert("json_mode".into(), json!({ "type": "boolean" }));
parameters.insert(
"structured_outputs".into(),
json!({ "type": "boolean" }),
);
if caps.is_some_and(|c| c.tools_branch) {
parameters.insert("tools".into(), json!({ "type": "boolean" }));
parameters.insert(
"tool_choice".into(),
json!({ "type": "enum", "values": ["auto", "none"] }),
);
}
if caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think) {
parameters.insert("reasoning".into(), json!({ "type": "boolean" }));
}
serde_json::Value::Object(parameters)
}
fn model_entry_openrouter(
name: &str,
caps: Option<&ModelCaps>,
metadata: Option<&OpenRouterModelMetadata>,
) -> serde_json::Value {
let empty = OpenRouterModelMetadata::default();
let metadata = metadata.unwrap_or(&empty);
let context_length = caps
.map(|c| c.context_length as u64)
.filter(|&v| v > 0 && v <= JSON_SAFE_INTEGER_MAX);
let tokenizer = caps
.map(|c| c.tokenizer.as_str())
.filter(|tokenizer| !tokenizer.is_empty());
let mut input = serde_json::Map::new();
input.insert("type".into(), json!("text"));
let mut supported_inputs = serde_json::Map::new();
if let Some(value) = context_length {
supported_inputs.insert(
"max_context_length".into(),
json!({ "value": value, "unit": "token" }),
);
}
if let Some(value) = metadata.max_prompt_length {
supported_inputs.insert(
"max_prompt_length".into(),
json!({ "value": value, "unit": "token" }),
);
}
if !supported_inputs.is_empty() {
input.insert(
"supported_inputs".into(),
serde_json::Value::Object(supported_inputs),
);
}
let mut input_pricing = Vec::new();
for (kind, cost) in [
("prompt", metadata.pricing.prompt.as_deref()),
(
"cached_prompt",
metadata.pricing.cached_prompt.as_deref(),
),
("cache_write", metadata.pricing.cache_write.as_deref()),
] {
if let Some(cost) = cost {
input_pricing.push(json!({
"type": kind,
"unit": "token",
"cost_usd": cost,
}));
}
}
if !input_pricing.is_empty() {
input.insert("pricing".into(), serde_json::Value::Array(input_pricing));
}
let mut input_capacity = Vec::new();
for (kind, value) in [
("prompt", metadata.capacity.prompt_tpm),
("cached_prompt", metadata.capacity.cached_prompt_tpm),
] {
if let Some(value) = value {
input_capacity.push(json!({
"type": kind,
"unit": "token",
"per": "minute",
"value": value,
}));
}
}
if !input_capacity.is_empty() {
input.insert(
"capacity".into(),
serde_json::Value::Array(input_capacity),
);
}
let mut output = serde_json::Map::new();
output.insert("type".into(), json!("text"));
output.insert(
"supported_parameters".into(),
openrouter_supported_parameters(caps, metadata.max_output_length),
);
output.insert("streaming".into(), json!(true));
if let Some(value) = metadata.max_output_length {
output.insert(
"max_length".into(),
json!({ "value": value, "unit": "token" }),
);
}
let mut output_pricing = Vec::new();
for (kind, cost) in [
("completion", metadata.pricing.completion.as_deref()),
(
"internal_reasoning",
metadata.pricing.internal_reasoning.as_deref(),
),
] {
if let Some(cost) = cost {
output_pricing.push(json!({
"type": kind,
"unit": "token",
"cost_usd": cost,
}));
}
}
if !output_pricing.is_empty() {
output.insert(
"pricing".into(),
serde_json::Value::Array(output_pricing),
);
}
let mut output_capacity = Vec::new();
if let Some(value) = metadata.capacity.completion_tpm {
output_capacity.push(json!({
"type": "completion",
"unit": "token",
"per": "minute",
"value": value,
}));
}
if let Some(value) = metadata.capacity.concurrency {
output_capacity.push(json!({
"type": "concurrency",
"unit": "request",
"value": value,
}));
}
if !output_capacity.is_empty() {
output.insert(
"capacity".into(),
serde_json::Value::Array(output_capacity),
);
}
let mut entry = serde_json::Map::new();
entry.insert("schema_version".into(), json!(OPENROUTER_SCHEMA_VERSION));
entry.insert("id".into(), json!(name));
entry.insert("name".into(), json!(name));
if let Some(value) = metadata.hugging_face_id.as_deref() {
entry.insert("hugging_face_id".into(), json!(value));
}
if let Some(value) = metadata.created {
entry.insert("created".into(), json!(value));
}
if let Some(value) = metadata.quantization.as_deref() {
entry.insert("quantization".into(), json!(value));
}
if let Some(value) = tokenizer {
entry.insert("tokenizer".into(), json!(value));
}
if let Some(value) = metadata.description.as_deref() {
entry.insert("description".into(), json!(value));
}
entry.insert(
"input_modalities".into(),
serde_json::Value::Array(vec![serde_json::Value::Object(input)]),
);
entry.insert(
"output_modalities".into(),
serde_json::Value::Array(vec![serde_json::Value::Object(output)]),
);
if let Some(cost) = metadata.pricing.request.as_deref() {
entry.insert(
"pricing".into(),
json!([{ "type": "request", "unit": "request", "cost_usd": cost }]),
);
}
if let Some(value) = metadata.capacity.request_rpm {
entry.insert(
"capacity".into(),
json!([{
"type": "request",
"unit": "request",
"per": "minute",
"value": value,
}]),
);
}
if let Some(value) = metadata.is_ready {
entry.insert("is_ready".into(), json!(value));
}
if let Some(value) = metadata.is_free {
entry.insert("is_free".into(), json!(value));
}
if let Some(value) = metadata.discount_to_user {
entry.insert("discount_to_user".into(), json!(value));
}
if let Some(value) = metadata.openrouter_slug.as_deref() {
entry.insert("openrouter".into(), json!({ "slug": value }));
}
if !metadata.datacenters.is_empty() {
entry.insert("datacenters".into(), json!(metadata.datacenters));
}
let mut compliance = serde_json::Map::new();
if let Some(value) = metadata.zdr {
compliance.insert("zdr".into(), json!(value));
}
if let Some(value) = metadata.hipaa {
compliance.insert("hipaa".into(), json!(value));
}
if !compliance.is_empty() {
entry.insert(
"compliance".into(),
serde_json::Value::Object(compliance),
);
}
serde_json::Value::Object(entry)
}
fn models_openrouter_body(st: &AppState) -> serde_json::Value {
let data: Vec<_> = st
.models
.iter()
.map(|model| {
model_entry_openrouter(
model,
st.caps.get(model),
st.openrouter_metadata.get(model),
)
})
.collect();
json!({ "data": data })
}
fn model_entry_openmodels(
name: &str,
caps: Option<&ModelCaps>,
metadata: Option<&OpenRouterModelMetadata>,
) -> Result<serde_json::Value, String> {
let metadata = metadata.ok_or_else(|| {
format!(
"OpenModels feed requires MEMRA_MODEL_METADATA for model {name:?}"
)
})?;
let context_length = caps
.map(|c| c.context_length as u64)
.filter(|&value| value > 0 && value <= JSON_SAFE_INTEGER_MAX)
.ok_or_else(|| format!("OpenModels feed requires context_length for model {name:?}"))?;
let created = metadata
.created
.ok_or_else(|| format!("OpenModels feed requires created for model {name:?}"))?;
let max_output_length = metadata.max_output_length.ok_or_else(|| {
format!("OpenModels feed requires max_output_length for model {name:?}")
})?;
let prompt = metadata
.pricing
.prompt
.as_deref()
.ok_or_else(|| format!("OpenModels feed requires pricing.prompt for model {name:?}"))?;
let completion = metadata.pricing.completion.as_deref().ok_or_else(|| {
format!("OpenModels feed requires pricing.completion for model {name:?}")
})?;
let input_cache_read = metadata.pricing.cached_prompt.as_deref().ok_or_else(|| {
format!("OpenModels feed requires pricing.cached_prompt for model {name:?}")
})?;
let is_ready = metadata
.is_ready
.ok_or_else(|| format!("OpenModels feed requires is_ready for model {name:?}"))?;
let is_free = metadata
.is_free
.ok_or_else(|| format!("OpenModels feed requires is_free for model {name:?}"))?;
let discount_to_user = metadata.discount_to_user.ok_or_else(|| {
format!("OpenModels feed requires discount_to_user for model {name:?}")
})?;
let mut pricing = serde_json::Map::new();
pricing.insert("prompt".into(), json!(prompt));
pricing.insert("completion".into(), json!(completion));
pricing.insert("input_cache_read".into(), json!(input_cache_read));
if let Some(value) = metadata.pricing.request.as_deref() {
pricing.insert("request".into(), json!(value));
}
let mut supported_features = Vec::new();
if caps.is_some_and(|c| c.tools_branch) {
supported_features.push("tool_calling");
}
if caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think) {
supported_features.push("reasoning");
}
let mut entry = serde_json::Map::new();
entry.insert("id".into(), json!(name));
entry.insert("name".into(), json!(name));
entry.insert("created".into(), json!(created));
entry.insert("input_modalities".into(), json!(["text"]));
entry.insert("output_modalities".into(), json!(["text"]));
entry.insert("context_length".into(), json!(context_length));
entry.insert("max_output_length".into(), json!(max_output_length));
entry.insert("currency".into(), json!("USD"));
entry.insert("pricing".into(), serde_json::Value::Object(pricing));
entry.insert("supported_features".into(), json!(supported_features));
entry.insert("is_ready".into(), json!(is_ready));
entry.insert("is_free".into(), json!(is_free));
entry.insert("discount_to_user".into(), json!(discount_to_user));
Ok(serde_json::Value::Object(entry))
}
fn models_openmodels_body(st: &AppState) -> Result<serde_json::Value, String> {
let data: Result<Vec<_>, _> = st
.models
.iter()
.map(|model| {
model_entry_openmodels(
model,
st.caps.get(model),
st.openrouter_metadata.get(model),
)
})
.collect();
Ok(json!({ "data": data? }))
}
async fn list_models(
State(st): State<AppState>,
Query(query): Query<ModelsQuery>,
) -> Response {
match query.schema.as_deref() {
None | Some("openai") => Json(models_openai_body(st.models.as_ref())).into_response(),
Some("openrouter") => Json(models_openrouter_body(&st)).into_response(),
Some("openmodels") => match models_openmodels_body(&st) {
Ok(body) => Json(body).into_response(),
Err(error) => bad_request(&error, Some("schema")),
},
Some(schema) => bad_request(
&format!(
"unsupported models schema {schema:?}; expected openai, openrouter, or openmodels"
),
Some("schema"),
),
}
}
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());
let thinking = caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think);
let mut supported = vec![
"max_tokens", "temperature", "top_p", "stop", "seed", "response_format",
"structured_outputs", "logit_bias", "logprobs", "top_logprobs",
];
if caps.is_some_and(|c| c.tools_branch) {
supported.extend(["tools", "tool_choice"]);
}
if thinking {
supported.extend(["reasoning", "include_reasoning", "reasoning_effort"]);
}
json!({
"id": name,
"name": name,
"object": "model",
"created": created,
"context_length": ctx,
"architecture": {
"modality": "text->text",
"tokenizer": tokenizer,
"instruct_type": instruct,
},
"supported_parameters": supported,
"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>, headers: HeaderMap) -> Response {
let metrics_scope = match authorize_metrics(&st.api_auth, &st.metrics_auth, &headers) {
Ok(scope) => scope,
Err(response) => return response,
};
if !metrics_scope.process_wide() {
return error_response(
StatusCode::FORBIDDEN,
"completion api keys do not authorize process-wide yield metrics; configure \
MEMRA_METRICS_TOKEN",
"authentication_error",
None,
);
}
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],
});
let mut body = json!({
"lanes": {
"interactive": lane(0), "judge": lane(1), "harvest": lane(2),
},
"interactive_step_ms": { "p50": m.step_p50_ms, "p99": m.step_p99_ms },
});
if metrics_scope.operator() {
body["batch_size_last"] = json!(m.batch_size_last);
}
Json(body).into_response()
}
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)) => Err(engine_error_response(&e)),
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)
}
}
}
#[cfg(test)]
fn build_request(req: &CompletionReq, tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>) -> Request {
build_request_with_trace(req, tx, lane, affinity, None)
}
fn build_request_with_trace(
req: &CompletionReq,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane,
affinity: Option<String>,
ttft: Option<Arc<ttft::Trace>>,
) -> 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,
reasoning_effort: None, params,
sampler_cfg,
stop_strings: req.stop.clone().into_vec(),
trace_id: req.trace_id.clone(),
max_prompt_tokens: None,
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar: None, prepared_constraint: None,
constraint_ready: None,
oom_retries: 0, spec_k_replay: None,
prepared_prompt: None,
ttft,
tx,
}
}
struct ChatPlan {
request: Request,
parser: Option<ToolStreamParser>,
}
#[cfg(test)]
fn build_chat_request(req: ChatCompletionReq, caps: Option<&ModelCaps>,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>)
-> Result<ChatPlan, String> {
build_chat_request_with_trace(req, caps, tx, lane, affinity, None)
}
fn build_chat_request_with_trace(
req: ChatCompletionReq,
caps: Option<&ModelCaps>,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane,
affinity: Option<String>,
ttft: Option<Arc<ttft::Trace>>,
)
-> Result<ChatPlan, String> {
let (temperature, top_p) = chat_sampling_values(&req, caps);
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, effort_level) = parse_think(&req.reasoning_effort, &req.reasoning)?;
let reasoning_effort = if caps.map(|c| c.effort_levels).unwrap_or(false) {
effort_level
} else {
None
};
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());
}
let role = if msg.role == "developer" { "system".to_string() } else { msg.role.clone() };
turns.push(TmplTurn { role, 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 if caps.map(|c| c.gemma_think).unwrap_or(false) {
Some(ToolStreamParser::gemma_thought(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,
reasoning_effort,
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(
temperature, req.top_k, 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,
max_prompt_tokens: None,
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar,
prepared_constraint: None,
constraint_ready: None,
oom_retries: 0, spec_k_replay: None,
prepared_prompt: None,
ttft,
tx,
},
parser,
})
}
fn bearer_token(headers: &HeaderMap) -> Option<&str> {
headers.get("authorization")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
}
fn authentication_error(why: auth::AuthDenied) -> Response {
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 authenticate(api_auth: &ApiAuth, headers: &HeaderMap) -> Result<auth::TenantCtx, Response> {
auth::authenticate_with(
api_auth.keyring,
api_auth.single_key.as_deref(),
bearer_token(headers),
).map_err(authentication_error)
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum MetricsScope {
All,
CompletionDomain,
Tenant(String),
}
impl MetricsScope {
fn operator(&self) -> bool {
matches!(self, MetricsScope::All)
}
fn process_wide(&self) -> bool {
matches!(self, MetricsScope::All | MetricsScope::CompletionDomain)
}
fn includes(&self, tenant_row: &str) -> bool {
match self {
MetricsScope::All | MetricsScope::CompletionDomain => true,
MetricsScope::Tenant(tenant) => tenant == tenant_row,
}
}
}
fn authorize_metrics(
api_auth: &ApiAuth,
metrics_auth: &MetricsAuth,
headers: &HeaderMap,
) -> Result<MetricsScope, Response> {
if !metrics_auth.required {
return Ok(MetricsScope::All);
}
let Some(candidate) = bearer_token(headers) else {
return Err(authentication_error(auth::AuthDenied::Unknown));
};
if let Some(token) = metrics_auth.token.as_deref() {
if auth::constant_time_secret_eq(token, candidate) {
return Ok(MetricsScope::All);
}
if api_auth.configured() {
return match auth::authenticate_with(
api_auth.keyring,
api_auth.single_key.as_deref(),
Some(candidate),
) {
Ok(_) => Err(error_response(
StatusCode::FORBIDDEN,
"completion api keys do not authorize metrics while \
MEMRA_METRICS_TOKEN is configured",
"authentication_error",
None,
)),
Err(why) => Err(authentication_error(why)),
};
}
return Err(authentication_error(auth::AuthDenied::Unknown));
}
if api_auth.configured() {
let tenant = authenticate(api_auth, headers)?;
return Ok(if api_auth.keyring.is_some() {
MetricsScope::Tenant(format!("t:{}", tenant.tenant))
} else {
MetricsScope::CompletionDomain
});
}
Err(authentication_error(auth::AuthDenied::Unknown))
}
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(|| {
error_response_coded(StatusCode::BAD_REQUEST,
&format!("unknown x-lane {v:?}; expected one of interactive, judge, harvest"),
"invalid_request_error", Some("x-lane"), Some("invalid_lane"))
})?),
};
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>,
) -> Result<String, &'static str> {
let keyring_configured = auth::global().is_some();
let raw = validate_cache_namespace(cache_salt, keyring_configured)?;
if keyring_configured {
Ok(auth::scope_namespace(&tenant.tenant, &raw))
} else {
Ok(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);
}
fn apply_model_request_limits(
request: &mut Request,
metadata: Option<&OpenRouterModelMetadata>,
) -> Result<(), (String, &'static str)> {
let Some(metadata) = metadata else { return Ok(()) };
let max_prompt = metadata.max_prompt_length
.map(usize::try_from)
.transpose()
.map_err(|_| ("configured model prompt limit does not fit this platform".into(), "model"))?;
let max_output = metadata.max_output_length
.map(usize::try_from)
.transpose()
.map_err(|_| ("configured model output limit does not fit this platform".into(), "model"))?;
request.max_prompt_tokens = max_prompt;
if let Some(max_output) = max_output {
if request.params.max_new == worker::MAX_NEW_CTX_BOUNDED {
request.params.max_new = max_output;
} else if request.params.max_new > max_output {
return Err((
format!(
"max_tokens {} exceeds configured model maximum {max_output}",
request.params.max_new
),
"max_tokens",
));
}
}
if let (Some(max_prompt), Some(max_output), Some(requested_ctx)) =
(max_prompt, max_output, request.params.max_ctx)
{
let operational_ctx = max_prompt
.checked_add(max_output)
.and_then(|value| value.checked_add(8))
.ok_or_else(|| ("configured model context envelope overflowed".into(), "model"))?;
if requested_ctx > operational_ctx {
return Err((
format!(
"max_ctx {requested_ctx} exceeds configured model envelope {operational_ctx}"
),
"max_ctx",
));
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn start_request_receipt(
st: &AppState,
env: &Envelope,
tenant: &auth::TenantCtx,
model: &str,
route: &'static str,
lane: lanes::Lane,
stream: bool,
budget_permit: Option<ledger::BudgetPermit>,
) -> Option<ledger::PendingReceipt> {
st.request_ledger.as_ref().map(|ledger| {
ledger.start_with_budget(
&env.id,
&tenant.tenant,
model,
route,
lane.as_str(),
stream,
budget_permit,
)
})
}
enum BudgetRejection {
Invalid(String),
Insufficient,
RequestInProgress,
Unavailable(String),
}
impl BudgetRejection {
fn into_response(self) -> (Response, &'static str) {
match self {
Self::Invalid(message) => (bad_request(&message, Some("prompt")), "invalid_request"),
Self::Insufficient => (
error_response_coded(
StatusCode::PAYMENT_REQUIRED,
"tenant prepaid balance is insufficient for this request",
"insufficient_balance",
None,
Some("insufficient_balance"),
),
"insufficient_balance",
),
Self::RequestInProgress => (
retry_contract_response(
error_response_coded(
StatusCode::TOO_MANY_REQUESTS,
"a budgeted request for this tenant is already in progress",
"rate_limit_error",
None,
Some("budget_request_in_progress"),
),
Some(1),
),
"budget_request_in_progress",
),
Self::Unavailable(err) => {
eprintln!("[budget] ERROR: admission unavailable: {err}");
(
error_response_coded(
StatusCode::SERVICE_UNAVAILABLE,
"tenant budget accounting is unavailable",
"server_error",
None,
Some("tenant_budget_unavailable"),
),
"tenant_budget_unavailable",
)
}
}
}
}
fn prepare_budget_prompt(
request: &mut Request,
tokenizer: Option<&Tokenizer>,
) -> Result<usize, String> {
if request.prepared_prompt.is_none() {
if let Some(trace) = request.ttft.as_ref() {
trace.mark_tokenize_start();
}
let prompt = if !request.prompt_ids.is_empty() {
request.prompt_ids.clone()
} else if !request.chat_turns.is_empty() {
let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
let plain = request.tools_json.is_empty()
&& request.think == ThinkMode::Default
&& request.reasoning_effort.is_none()
&& request
.chat_turns
.iter()
.all(|turn| turn.role != "tool" && turn.tool_calls.is_empty());
let rendered = if plain {
let messages: Vec<_> = request
.chat_turns
.iter()
.map(|turn| (turn.role.as_str(), turn.content.as_str()))
.collect();
tokenizer.apply_chat_template(&messages, true)
} else {
tokenizer
.apply_chat_template_tools(
&request.chat_turns,
true,
&request.tools_json,
request.think,
request.reasoning_effort.as_deref(),
)
.map_err(|err| format!("chat template: {err}"))?
};
tokenizer.encode(&rendered, true)
} else if request.chat {
let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
let rendered =
tokenizer.apply_chat_template(&[("user", request.prompt_text.as_str())], true);
tokenizer.encode(&rendered, true)
} else {
let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
tokenizer.encode(&request.prompt_text, true)
};
if prompt.is_empty() {
return Err("empty prompt after tokenization".into());
}
if let Some(trace) = request.ttft.as_ref() {
trace.mark_tokenize_end(prompt.len());
}
request.prepared_prompt = Some(prompt);
}
let prompt_tokens = request
.prepared_prompt
.as_ref()
.expect("budget prompt was prepared")
.len();
if let Some(limit) = request.max_prompt_tokens
&& prompt_tokens > limit
{
return Err(format!(
"prompt ({prompt_tokens} tok) exceeds configured model maximum ({limit})"
));
}
Ok(prompt_tokens)
}
fn budget_completion_bound(
request: &Request,
prompt_tokens: usize,
caps: Option<&ModelCaps>,
) -> Result<usize, String> {
let max_new = request.params.max_new;
let ctx_cap = match (request.params.max_ctx, max_new) {
(Some(cap), _) => cap,
(None, worker::MAX_NEW_CTX_BOUNDED) => {
let server_ctx = std::env::var("MEMRA_CTX")
.ok()
.and_then(|value| value.parse().ok())
.unwrap_or(8192usize);
let mut cap = server_ctx;
if prompt_tokens.saturating_add(16) > cap {
cap = prompt_tokens.saturating_add(server_ctx);
}
let model_ctx = caps.map_or(0, |caps| caps.context_length);
if model_ctx > 0 {
cap = cap.min(model_ctx);
}
cap
}
(None, max_new) => prompt_tokens
.checked_add(max_new)
.and_then(|value| value.checked_add(8))
.ok_or_else(|| "request context bound overflowed".to_string())?,
};
if prompt_tokens >= ctx_cap {
return Err(format!(
"prompt ({prompt_tokens} tok) >= context cap ({ctx_cap})"
));
}
Ok(max_new.min(ctx_cap - prompt_tokens))
}
fn admit_tenant_budget(
st: &AppState,
tenant: &auth::TenantCtx,
request: &mut Request,
) -> Result<Option<ledger::BudgetPermit>, BudgetRejection> {
let Some(budgets) = st.tenant_budgets.as_ref() else {
return Ok(None);
};
match budgets.is_limited(&tenant.tenant) {
Ok(false) => return Ok(None),
Ok(true) => {}
Err(ledger::BudgetAdmissionError::Unavailable(err)) => {
return Err(BudgetRejection::Unavailable(err));
}
Err(other) => {
return Err(BudgetRejection::Unavailable(format!(
"unexpected budget enrollment result: {other:?}"
)));
}
}
let tokenizer = st
.budget_tokenizers
.as_ref()
.and_then(|tokenizers| tokenizers.get(&request.model))
.map(Arc::as_ref);
if request.prompt_ids.is_empty() && tokenizer.is_none() {
return Err(BudgetRejection::Unavailable(format!(
"no reservation tokenizer for model {:?}",
request.model
)));
}
let prompt_tokens =
prepare_budget_prompt(request, tokenizer).map_err(BudgetRejection::Invalid)?;
let completion_tokens =
budget_completion_bound(request, prompt_tokens, st.caps.get(&request.model))
.map_err(BudgetRejection::Invalid)?;
let prompt_tokens = u64::try_from(prompt_tokens)
.map_err(|_| BudgetRejection::Unavailable("prompt token count exceeds u64".into()))?;
let completion_tokens = u64::try_from(completion_tokens)
.map_err(|_| BudgetRejection::Unavailable("completion token bound exceeds u64".into()))?;
let ledger = st.request_ledger.as_ref().ok_or_else(|| {
BudgetRejection::Unavailable("tenant budgets require the request ledger".into())
})?;
match ledger.reserve_budget(
&tenant.tenant,
&request.model,
prompt_tokens,
completion_tokens,
) {
Ok(permit) => Ok(permit),
Err(ledger::BudgetAdmissionError::Insufficient) => Err(BudgetRejection::Insufficient),
Err(ledger::BudgetAdmissionError::RequestInProgress) => {
Err(BudgetRejection::RequestInProgress)
}
Err(ledger::BudgetAdmissionError::Unavailable(err)) => {
Err(BudgetRejection::Unavailable(err))
}
}
}
fn request_ledger_error_response() -> Response {
error_response_coded(
StatusCode::INTERNAL_SERVER_ERROR,
"request completion could not be committed to the billing ledger",
"server_error",
None,
Some("request_ledger_unavailable"),
)
}
fn request_ledger_error_body() -> serde_json::Value {
error_body(
"request completion could not be committed to the billing ledger",
"server_error",
None,
Some("request_ledger_unavailable"),
)
}
fn ledger_rejected(
mut receipt: Option<ledger::PendingReceipt>,
response: Response,
error_code: &str,
request_id: &str,
) -> Response {
let status = response.status().as_u16();
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.reject(status, error_code)
{
eprintln!("[ledger] ERROR: request {request_id} rejection receipt failed: {err}");
return with_request_id(request_id, request_ledger_error_response());
}
with_request_id(request_id, response)
}
fn engine_error_code(class: worker::ErrClass) -> &'static str {
use worker::ErrClass as C;
match class {
C::InvalidRequest => "invalid_request",
C::ContextLength => "context_length_exceeded",
C::ModelNotFound => "model_not_found",
C::RateLimit => "rate_limit_exceeded",
C::Overloaded => "overloaded",
C::Engine => "engine_error",
}
}
async fn completions(State(st): State<AppState>, headers: axum::http::HeaderMap,
trace: Option<Extension<TtftRequestTrace>>,
Json(req): Json<CompletionReq>) -> Response {
let env = Envelope::new(false);
let ttft = trace.and_then(|Extension(trace)| trace.0);
if let Some(trace) = ttft.as_ref() {
trace.mark_parsed();
trace.bind_request(&env.id, &req.model);
}
let tenant = match authenticate(&st.api_auth, &headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
let cache_ns = match tenant_namespace(&tenant, &req.cache_salt) {
Ok(ns) => ns,
Err(msg) => return with_request_id(&env.id, bad_request(msg, Some("cache_salt"))),
};
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,
};
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_with_trace(&req, tx, lane, affinity, ttft.clone());
request.cache_ns = cache_ns;
if let Err((message, param)) =
apply_model_request_limits(&mut request, st.openrouter_metadata.get(&model))
{
return with_request_id(&env.id, bad_request(&message, Some(param)));
}
if draining() {
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&req.model,
"/v1/completions",
lane,
req.stream,
None,
);
return ledger_rejected(receipt, drain_response(), "draining", &env.id);
}
let budget_permit = match admit_tenant_budget(&st, &tenant, &mut request) {
Ok(permit) => permit,
Err(rejection) => {
let (response, error_code) = rejection.into_response();
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&req.model,
"/v1/completions",
lane,
req.stream,
None,
);
return ledger_rejected(receipt, response, error_code, &env.id);
}
};
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&req.model,
"/v1/completions",
lane,
req.stream,
budget_permit,
);
let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
Ok(slot) => slot,
Err(resp) => {
return ledger_rejected(receipt, resp, "rate_limit_exceeded", &env.id);
}
};
meter_admit(&env, &tenant, &model, lane);
let stop_strings = request.stop_strings.clone();
worker::PENDING_ADMITS.fetch_add(1, std::sync::atomic::Ordering::Release);
if let Some(trace) = ttft.as_ref() {
trace.mark_submitted();
}
if st.cmd_tx.send(Cmd::Generate(Box::new(request))).is_err() {
worker::PENDING_ADMITS.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
return ledger_rejected(
receipt,
rl.attach(worker_unavailable_response()),
"worker_unavailable",
&env.id,
);
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => {
return ledger_rejected(
receipt,
rl.attach(resp),
"admission_shed",
&env.id,
);
}
};
let resp = if stream {
sse_response_with_receipt(
rx,
model,
false,
None,
env.clone(),
stop_strings,
Some(guard),
receipt,
)
.into_response()
} else {
let resp = blocking_response_with_receipt(
rx,
model,
false,
stop_strings,
None,
env.clone(),
receipt,
)
.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,
trace: Option<Extension<TtftRequestTrace>>,
Json(req): Json<ChatCompletionReq>) -> Response {
let env = Envelope::new(true);
let ttft = trace.and_then(|Extension(trace)| trace.0);
if let Some(trace) = ttft.as_ref() {
trace.mark_parsed();
trace.bind_request(&env.id, &req.model);
}
let tenant = match authenticate(&st.api_auth, &headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
let cache_ns = match tenant_namespace(&tenant, &req.cache_salt) {
Ok(ns) => ns,
Err(msg) => return with_request_id(&env.id, bad_request(msg, Some("cache_salt"))),
};
if req.messages.is_empty() || req.messages.iter().any(|message| {
!matches!(message.role.as_str(),
"system" | "developer" | "user" | "assistant" | "tool")
}) {
return with_request_id(&env.id, bad_request(
"messages must use system/developer/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 (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_with_trace(
req, st.caps.get(&model), tx, lane, affinity, ttft.clone())
{
Ok(plan) => plan,
Err(err) => {
return with_request_id(&env.id, bad_request(&err, None));
}
};
plan.request.cache_ns = cache_ns;
if let Err((message, param)) =
apply_model_request_limits(&mut plan.request, st.openrouter_metadata.get(&model))
{
return with_request_id(&env.id, bad_request(&message, Some(param)));
}
if draining() {
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&model,
"/v1/chat/completions",
lane,
stream,
None,
);
return ledger_rejected(receipt, drain_response(), "draining", &env.id);
}
let budget_permit = match admit_tenant_budget(&st, &tenant, &mut plan.request) {
Ok(permit) => permit,
Err(rejection) => {
let (response, error_code) = rejection.into_response();
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&model,
"/v1/chat/completions",
lane,
stream,
None,
);
return ledger_rejected(receipt, response, error_code, &env.id);
}
};
let receipt = start_request_receipt(
&st,
&env,
&tenant,
&model,
"/v1/chat/completions",
lane,
stream,
budget_permit,
);
let constraint_ready = if plan.request.grammar.is_some() {
let (ready_tx, ready_rx) = tokio::sync::oneshot::channel();
plan.request.constraint_ready = Some(ready_tx);
Some(ready_rx)
} else {
None
};
let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
Ok(slot) => slot,
Err(resp) => {
return ledger_rejected(receipt, resp, "rate_limit_exceeded", &env.id);
}
};
meter_admit(&env, &tenant, &model, lane);
let stop_strings = plan.request.stop_strings.clone();
worker::PENDING_ADMITS.fetch_add(1, std::sync::atomic::Ordering::Release);
if let Some(trace) = ttft.as_ref() {
trace.mark_submitted();
}
if st.cmd_tx.send(Cmd::Generate(Box::new(plan.request))).is_err() {
worker::PENDING_ADMITS.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
return ledger_rejected(
receipt,
rl.attach(worker_unavailable_response()),
"worker_unavailable",
&env.id,
);
}
if let Some(ready) = constraint_ready {
match tokio::time::timeout(constrained::CONSTRAINT_COMPILE_TIMEOUT, ready).await {
Ok(Ok(Ok(()))) => {}
Ok(Ok(Err(err))) => {
return ledger_rejected(
receipt,
rl.attach(engine_error_response(&err)),
engine_error_code(err.class),
&env.id,
);
}
Ok(Err(_)) => {
return ledger_rejected(
receipt,
rl.attach(worker_unavailable_response()),
"worker_unavailable",
&env.id,
);
}
Err(_) => {
return ledger_rejected(
receipt,
rl.attach(engine_error_response(&worker::constraint_timeout_error())),
"constraint_compile_timeout",
&env.id,
);
}
}
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => {
return ledger_rejected(
receipt,
rl.attach(resp),
"admission_shed",
&env.id,
);
}
};
let resp = if stream {
sse_response_with_receipt(
rx,
model,
true,
plan.parser,
env.clone(),
stop_strings,
Some(guard),
receipt,
)
.into_response()
} else {
let resp = blocking_response_with_receipt(
rx,
model,
true,
stop_strings,
plan.parser,
env.clone(),
receipt,
)
.await
.into_response();
drop(guard); resp
};
rl.attach(with_request_id(&env.id, resp))
}
#[cfg(test)]
fn sse_response(rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String, chat: bool,
parser: Option<ToolStreamParser>, env: Envelope,
stop_strings: Vec<String>, guard: Option<InflightGuard>)
-> Sse<impl futures_core::Stream<Item = Result<SseEvent, std::convert::Infallible>>> {
sse_response_with_receipt(
rx,
model,
chat,
parser,
env,
stop_strings,
guard,
None,
)
}
fn sse_response_with_receipt(
mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>,
model: String,
chat: bool,
mut parser: Option<ToolStreamParser>,
env: Envelope,
stop_strings: Vec<String>,
guard: Option<InflightGuard>,
mut receipt: Option<ledger::PendingReceipt>,
) -> 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::PromptUsage { n_prompt, n_cached } => {
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.record_prompt_usage(
n_prompt as u64,
n_cached as u64,
)
{
eprintln!(
"[ledger] ERROR: request {} partial prompt receipt failed: {err}",
env.id
);
let payload = request_ledger_error_body().to_string();
if chat || openai_compat() {
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
yield Ok(SseEvent::default().event("error").data(payload));
}
break;
}
}
Event::Token { id, text } => {
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.record_completion_token()
{
eprintln!(
"[ledger] ERROR: request {} partial completion receipt failed: {err}",
env.id
);
let payload = request_ledger_error_body().to_string();
if chat || openai_compat() {
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
yield Ok(SseEvent::default().event("error").data(payload));
}
break;
}
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::TokenSnapshot(_) => {}
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 let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.complete(
ledger::Usage {
prompt_tokens: n_prompt as u64,
cached_prompt_tokens: n_cached as u64,
completion_tokens: n_tokens as u64,
},
elapsed_s,
)
{
eprintln!(
"[ledger] ERROR: request {} completion receipt failed: {err}",
env.id
);
let payload = request_ledger_error_body().to_string();
if chat || openai_compat() {
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
yield Ok(SseEvent::default().event("error").data(payload));
}
break;
}
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(err) => {
let ledger_error = if let Some(receipt) = receipt.as_mut() {
receipt
.reject(class_http(err.class).0.as_u16(), engine_error_code(err.class))
.err()
} else {
None
};
if let Some(ref ledger_error) = ledger_error {
eprintln!(
"[ledger] ERROR: request {} failure receipt failed: {ledger_error}",
env.id
);
}
let payload = if ledger_error.is_some() {
request_ledger_error_body().to_string()
} else {
engine_error_body(&err).to_string()
};
if chat || openai_compat() {
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
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)
}
}
#[cfg(test)]
async fn blocking_response(rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String,
chat: bool, stop_strings: Vec<String>,
parser: Option<ToolStreamParser>, env: Envelope) -> Response {
blocking_response_with_receipt(
rx,
model,
chat,
stop_strings,
parser,
env,
None,
)
.await
}
async fn blocking_response_with_receipt(
mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>,
model: String,
chat: bool,
stop_strings: Vec<String>,
mut parser: Option<ToolStreamParser>,
env: Envelope,
mut receipt: Option<ledger::PendingReceipt>,
) -> 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::PromptUsage { n_prompt, n_cached } => {
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.record_prompt_usage(
n_prompt as u64,
n_cached as u64,
)
{
eprintln!(
"[ledger] ERROR: request {} partial prompt receipt failed: {err}",
env.id
);
return request_ledger_error_response();
}
}
Event::Token { id, text: delta } => {
if let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.record_completion_token()
{
eprintln!(
"[ledger] ERROR: request {} partial completion receipt failed: {err}",
env.id
);
return request_ledger_error_response();
}
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::TokenSnapshot(ids) => tokens = ids,
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 let Some(receipt) = receipt.as_mut()
&& let Err(err) = receipt.complete(
ledger::Usage {
prompt_tokens: n_prompt as u64,
cached_prompt_tokens: n_cached as u64,
completion_tokens: n_tokens as u64,
},
elapsed_s,
)
{
eprintln!(
"[ledger] ERROR: request {} completion receipt failed: {err}",
env.id
);
return request_ledger_error_response();
}
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(err) => {
if let Some(receipt) = receipt.as_mut()
&& let Err(ledger_err) = receipt.reject(
class_http(err.class).0.as_u16(),
engine_error_code(err.class),
)
{
eprintln!(
"[ledger] ERROR: request {} failure receipt failed: {ledger_err}",
env.id
);
return request_ledger_error_response();
}
return engine_error_response(&err);
}
}
}
let e = worker::EngineError::overloaded(
"worker closed the stream without completing (worker restart in progress)");
if let Some(receipt) = receipt.as_mut()
&& let Err(ledger_err) = receipt.reject(
class_http(e.class).0.as_u16(),
engine_error_code(e.class),
)
{
eprintln!(
"[ledger] ERROR: request {} closed-stream receipt failed: {ledger_err}",
env.id
);
return request_ledger_error_response();
}
engine_error_response(&e)
}
#[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 models_v1_entry_advertises_thinking_support() {
let step_caps = ModelCaps { effort_levels: true, ..tool_caps() };
let entry = model_entry_v1("stepfun/step-3.7-flash", Some(&step_caps), 1);
let supported: Vec<&str> = entry["supported_parameters"]
.as_array().unwrap().iter().map(|v| v.as_str().unwrap()).collect();
for knob in ["reasoning", "include_reasoning", "reasoning_effort"] {
assert!(supported.contains(&knob), "missing {knob} for thinking model");
}
assert!(supported.contains(&"tools"));
let plain = ModelCaps { chat_ok: true, ..Default::default() };
let entry = model_entry_v1("plain", Some(&plain), 1);
let supported: Vec<&str> = entry["supported_parameters"]
.as_array().unwrap().iter().map(|v| v.as_str().unwrap()).collect();
for knob in ["reasoning", "include_reasoning", "reasoning_effort", "tools"] {
assert!(!supported.contains(&knob), "{knob} advertised without capability");
}
let entry = model_entry_v1("unknown", None, 1);
assert!(entry["supported_parameters"].as_array().unwrap().len() >= 6);
}
#[test]
fn chat_request_preserves_turns_and_openai_stop_forms() {
let payload = serde_json::json!({
"model": "plain_quant",
"messages": [
{"role": "system", "content": "rules"},
{"role": "developer", "content": "dev 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()),
("system".into(), "dev 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 native_response_uses_terminal_token_snapshot_for_coalesced_events() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 4, text: "hello".into() }).unwrap();
tx.send(Event::TokenSnapshot(vec![1, 2, 3, 4])).unwrap();
tx.send(Event::Done {
stop_reason: "MaxNew".into(), n_tokens: 4, n_prompt: 2, n_cached: 0,
elapsed_s: 0.5, spec: None,
}).unwrap();
drop(tx);
let response = blocking_response(rx, "plain_quant".into(), false, Vec::new(), 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"], "hello");
assert_eq!(payload["tokens"], serde_json::json!([1, 2, 3, 4]));
assert_eq!(payload["n_tokens"], 4);
}
#[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::Think),
(json!({"reasoning_effort": "none"}), ThinkMode::NoThink),
(json!({"reasoning_effort": "minimal"}), ThinkMode::NoThink),
(json!({"reasoning_effort": "high"}), ThinkMode::Think),
(json!({"reasoning_effort": "medium"}), ThinkMode::Think),
(json!({"reasoning": {"enabled": false}}), ThinkMode::NoThink),
(json!({"reasoning": {"effort": "low"}}), ThinkMode::Think),
(json!({"reasoning": {"enabled": true}}), ThinkMode::Think),
] {
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 reasoning_effort_maps_to_effort_level_on_step35_class_templates() {
let effort_caps = ModelCaps { effort_levels: true, ..tool_caps() };
for (extra, want) in [
(json!({}), None),
(json!({"reasoning_effort": "low"}), Some("low")),
(json!({"reasoning_effort": "medium"}), Some("medium")),
(json!({"reasoning_effort": "high"}), Some("high")),
(json!({"reasoning_effort": "none"}), Some("low")),
(json!({"reasoning_effort": "minimal"}), Some("low")),
(json!({"reasoning": {"effort": "high"}}), Some("high")),
(json!({"reasoning": {"enabled": false}}), Some("low")),
] {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(extra.clone()),
Some(&effort_caps), tx, lanes::Lane::Interactive, None).unwrap();
assert_eq!(plan.request.reasoning_effort.as_deref(), want, "extra={extra}");
}
for extra in [json!({}), json!({"reasoning_effort": "high"}),
json!({"reasoning": {"effort": "low"}})] {
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.reasoning_effort, None, "extra={extra}");
}
}
#[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 cache_salt_validation_rejects_oversized_value() {
let salt = Some("a".repeat(CACHE_SALT_MAX_BYTES + 1));
assert_eq!(
validate_cache_namespace(&salt, false),
Err("cache_salt must be at most 64 bytes")
);
}
#[test]
fn cache_salt_validation_rejects_reserved_open_namespace() {
let salt = Some("t:acme\u{1f}private".to_string());
assert_eq!(
validate_cache_namespace(&salt, false),
Err("cache_salt must not use the reserved t: prefix without a keyring")
);
}
#[test]
fn cache_salt_validation_accepts_normal_value() {
let salt = Some("tenant-A_7.c2VjcmV0LXNjb3Bl+/=".to_string());
assert_eq!(validate_cache_namespace(&salt, false).unwrap(), salt.unwrap());
assert_eq!(validate_cache_namespace(&None, false).unwrap(), "");
let max = Some("a".repeat(CACHE_SALT_MAX_BYTES));
assert_eq!(validate_cache_namespace(&max, false).unwrap(), max.unwrap());
}
#[test]
fn cache_salt_validation_rejects_unsupported_characters() {
let salt = Some("tenant salt".to_string());
assert_eq!(
validate_cache_namespace(&salt, false),
Err("cache_salt contains unsupported characters")
);
}
#[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_token_events_equal_usage_on_every_finish_path() {
for (stop_reason, expected_finish) in [
("Eos", "stop"),
("Callback", "stop"),
("MaxNew", "length"),
("ContextFull", "length"),
] {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 248_046, text: String::new() }).unwrap();
tx.send(Event::Done {
stop_reason: stop_reason.into(), n_tokens: 1, n_prompt: 8, n_cached: 8,
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(|line| serde_json::from_str(line).unwrap())
.collect();
let token_events = chunks.iter().filter(|chunk| {
chunk["choices"][0]["finish_reason"].is_null()
}).count();
let terminal = chunks.last().unwrap();
assert_eq!(token_events, 1, "{stop_reason} SSE token count");
assert_eq!(terminal["usage"]["completion_tokens"], token_events);
assert_eq!(terminal["choices"][0]["finish_reason"], expected_finish);
}
}
#[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(worker::EngineError::engine("boom"))).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!(err["error"]["code"], "engine_error");
assert_eq!(lines.last(), Some(&"[DONE]"));
}
#[test]
fn ttft_sse_marker_ignores_keepalive_comments() {
assert!(!is_sse_data_frame(b": keep-alive\n\n"));
assert!(is_sse_data_frame(b"data: {\"choices\":[]}\n\n"));
assert!(is_sse_data_frame(
b"event: error\ndata: {\"error\":\"failed\"}\n\n"
));
}
#[tokio::test]
async fn error_bodies_use_the_openai_object_shape() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error(worker::EngineError::model_not_found("unknown model \"x\""))).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_eq!(payload["error"]["param"], "model");
assert_eq!(payload["error"]["code"], "model_not_found");
}
fn retry_after(resp: &Response) -> Option<String> {
resp.headers().get(axum::http::header::RETRY_AFTER)
.and_then(|v| v.to_str().ok()).map(str::to_string)
}
#[test]
fn taxonomy_maps_every_class_to_its_status_and_code() {
use worker::{EngineError as E, ErrClass as C};
let cases: Vec<(worker::EngineError, StatusCode, &str, &str)> = vec![
(E::invalid_param("bad json", "response_format"),
StatusCode::BAD_REQUEST, "invalid_request_error", ""),
(E::context_length("prompt (9000 tok) >= context cap (8192)"),
StatusCode::BAD_REQUEST, "invalid_request_error", "context_length_exceeded"),
(E::model_not_found("unknown model \"nope\""),
StatusCode::BAD_REQUEST, "invalid_request_error", "model_not_found"),
(E::rate_limit("lane judge is at capacity, retry"),
StatusCode::TOO_MANY_REQUESTS, "rate_limit_error", "rate_limit_exceeded"),
(E::overloaded("no VRAM for a new session"),
StatusCode::SERVICE_UNAVAILABLE, "server_error", "overloaded"),
(E::engine("graph step failed: launch error"),
StatusCode::INTERNAL_SERVER_ERROR, "server_error", "engine_error"),
];
for (err, want_status, want_type, want_code) in cases {
let (status, etype, code) = class_http(err.class);
assert_eq!(status, want_status, "{:?}", err);
assert_eq!(etype, want_type, "{:?}", err);
if !want_code.is_empty() {
assert_eq!(code, Some(want_code), "{:?}", err);
}
let body = engine_error_body(&err);
assert_eq!(body["error"]["message"], err.message);
assert_eq!(body["error"]["type"], want_type);
}
for c in [C::InvalidRequest, C::ContextLength, C::ModelNotFound,
C::RateLimit, C::Overloaded, C::Engine] {
let (s, t, _) = class_http(c);
assert!(s.is_client_error() || s.is_server_error(), "{c:?} -> {s}");
assert!(!t.is_empty());
}
}
#[test]
fn a_cuda_oom_message_is_capacity_503_not_a_500() {
let e = worker::EngineError::engine(
"step error: DriverError(CUDA_ERROR_OUT_OF_MEMORY, \"out of memory\")");
let resp = engine_error_response(&e);
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(retry_after(&resp).as_deref(), Some("5"));
}
#[test]
fn retry_headers_follow_the_sdk_contract() {
for e in [worker::EngineError::rate_limit("shed"),
worker::EngineError::overloaded("no VRAM")] {
let resp = engine_error_response(&e);
let ra = retry_after(&resp).expect("retryable class must carry Retry-After");
let secs: u64 = ra.parse().expect("Retry-After must be integer delay-seconds");
assert!(secs > 0 && secs <= 60, "Retry-After {secs}s outside the honored window");
let ms = resp.headers().get("retry-after-ms").unwrap().to_str().unwrap();
assert_eq!(ms.parse::<u64>().unwrap(), secs * 1000, "the two headers disagree");
assert!(resp.headers().get("x-should-retry").is_none(),
"a retryable class must not say x-should-retry: false");
}
}
#[tokio::test]
async fn command_send_failure_obeys_the_retry_contract() {
let _l = DRAIN_LOCK.lock().unwrap();
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let mut st = fake_worker_state();
let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
drop(cmd_rx);
st.cmd_tx = cmd_tx;
let completion = completions(
State(st.clone()),
axum::http::HeaderMap::new(),
None,
Json(serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "test"
})).unwrap()),
).await;
let chat = chat_completions(
State(st),
axum::http::HeaderMap::new(),
None,
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "test"}]
})).unwrap()),
).await;
for resp in [completion, chat] {
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(retry_after(&resp).as_deref(), Some("2"));
assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
assert_ne!(resp.headers().get("x-should-retry")
.and_then(|v| v.to_str().ok()), Some("false"));
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["error"]["type"], "server_error");
assert_eq!(payload["error"]["code"], "overloaded");
}
}
#[test]
fn unfixable_client_errors_say_x_should_retry_false() {
for e in [worker::EngineError::model_not_found("unknown model \"x\""),
worker::EngineError::context_length("prompt too long"),
worker::EngineError::invalid_param("bad", "messages")] {
let resp = engine_error_response(&e);
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
assert!(retry_after(&resp).is_none(), "a 400 must not promise a retry window");
}
}
#[tokio::test]
async fn a_closed_worker_channel_is_503_not_500() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Event>();
drop(tx);
let resp = blocking_response(rx, "m".into(), true, Vec::new(), None,
Envelope::new(true)).await;
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(retry_after(&resp).as_deref(), Some("5"));
}
#[tokio::test]
async fn a_dark_lane_shed_is_429_with_an_openai_object_body() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error(worker::EngineError::rate_limit(
"lane judge shed: interactive p99 over budget, retry"))).unwrap();
let resp = peek_shed(lanes::Lane::Judge, rx).await
.err().expect("a shed must not be forwarded into the stream");
assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(retry_after(&resp).as_deref(), Some("2"));
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"].is_object(), "bare-string error body: {payload}");
assert_eq!(payload["error"]["type"], "rate_limit_error");
assert!(payload["error"]["message"].as_str().unwrap().contains("shed"));
}
#[tokio::test]
async fn interactive_never_peeks_so_its_first_token_is_not_held() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error(worker::EngineError::rate_limit("would-be shed"))).unwrap();
assert!(peek_shed(lanes::Lane::Interactive, rx).await.is_ok());
}
#[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 step35_chat_uses_published_sampling_defaults_only_when_omitted() {
let caps = ModelCaps {
chat_temperature_default: Some(0.5),
chat_top_p_default: Some(0.9),
chat_ok: true,
..Default::default()
};
let cfg = |extra: serde_json::Value| {
let mut body = serde_json::json!({
"model": "step35",
"messages": [{"role": "user", "content": "task"}]
});
body.as_object_mut().unwrap().extend(extra.as_object().unwrap().clone());
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, Some(&caps), tx, lanes::Lane::Interactive, None)
.unwrap().request.sampler_cfg
};
let omitted = cfg(serde_json::json!({}));
assert_eq!(omitted.temperature, 0.5);
assert_eq!(omitted.top_p, 0.9);
let explicit_temp = cfg(serde_json::json!({"temperature": 0.7}));
assert_eq!(explicit_temp.temperature, 0.7);
assert_eq!(explicit_temp.top_p, 0.9,
"omitting top_p must retain StepFun's nucleus default");
let explicit = cfg(serde_json::json!({"temperature": 0.0, "top_p": 1.0}));
assert_eq!(explicit.temperature, 0.0, "explicit greedy must remain authoritative");
assert_eq!(explicit.top_p, 1.0, "explicit untruncated sampling must remain authoritative");
}
#[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 {
fake_worker_state_with_steps(1, std::time::Duration::ZERO)
}
fn fake_worker_state_with_steps(steps: usize, step_delay: std::time::Duration) -> AppState {
let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
let health = health::WorkerHealth::new();
let h = health.clone();
std::thread::spawn(move || {
h.mark_ready();
while let Ok(Cmd::Generate(mut req)) = cmd_rx.recv() {
let _ = worker::PENDING_ADMITS.fetch_update(
std::sync::atomic::Ordering::AcqRel,
std::sync::atomic::Ordering::Acquire,
|v| v.checked_sub(1),
);
h.beat_busy();
if let Some(ready) = req.constraint_ready.take() {
let _ = ready.send(Ok(()));
}
let _ = req.tx.send(Event::PromptUsage {
n_prompt: 1,
n_cached: 0,
});
for step in 0..steps {
h.beat_busy();
let text = if steps == 1 { "ok" } else { "x" };
let _ = req.tx.send(Event::Token {
id: step as u32 + 1,
text: text.into(),
});
if !step_delay.is_zero() {
std::thread::sleep(step_delay);
}
}
let _ = req.tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: steps, n_prompt: 1, n_cached: 0,
elapsed_s: 0.01, spec: None,
});
h.set_phase(health::PHASE_IDLE);
}
});
for _ in 0..2000 {
if health.live().is_ok() { break; }
std::thread::sleep(std::time::Duration::from_millis(1));
}
AppState {
cmd_tx,
models: Arc::new(vec!["m".into()]),
caps: Arc::new(HashMap::new()),
openrouter_metadata: Arc::new(HashMap::new()),
request_ledger: None,
tenant_budgets: None,
budget_tokenizers: None,
api_auth: ApiAuth::default(),
metrics_auth: MetricsAuth::default(),
metrics: SharedMetrics::default(),
started: 1,
inflight: Arc::new(Default::default()),
tenant_inflight: Arc::new(Default::default()),
health,
bg: None,
}
}
#[tokio::test]
async fn deep_schema_fails_while_normal_decode_keeps_stepping() {
let _l = DRAIN_LOCK.lock().unwrap();
let st = fake_worker_state_with_steps(
64,
std::time::Duration::from_millis(5),
);
let normal_state = st.clone();
let normal = tokio::spawn(async move {
chat_completions(
State(normal_state),
axum::http::HeaderMap::new(),
None,
Json(serde_json::from_value(serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "keep decoding"}],
})).unwrap()),
).await
});
tokio::time::sleep(std::time::Duration::from_millis(15)).await;
let mut deep = serde_json::json!({"type": "string"});
for _ in 0..(constrained::MAX_SCHEMA_DEPTH / 2 + 1) {
deep = serde_json::json!({"allOf": [deep]});
}
let bad = chat_completions(
State(st.clone()),
axum::http::HeaderMap::new(),
None,
Json(serde_json::from_value(serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "bad schema"}],
"response_format": {
"type": "json_schema",
"json_schema": {"schema": deep},
},
})).unwrap()),
).await;
assert_eq!(bad.status(), StatusCode::BAD_REQUEST);
assert_eq!(bad.headers().get("x-should-retry").unwrap(), "false");
let bytes = axum::body::to_bytes(bad.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("maximum nesting depth"));
assert!(!normal.is_finished(), "bad schema stalled or replaced the normal decode");
let normal_response = normal.await.unwrap();
assert_eq!(normal_response.status(), StatusCode::OK);
let snapshot = st.health.snapshot();
assert!(st.health.live().is_ok(), "normal decode left health stalled");
assert!(snapshot.beat_age_ms < snapshot.stall_threshold_ms);
}
#[tokio::test]
async fn valid_response_format_preflight_preserves_generation() {
let _l = DRAIN_LOCK.lock().unwrap();
let response = chat_completions(
State(fake_worker_state()),
axum::http::HeaderMap::new(),
None,
Json(serde_json::from_value(serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "valid schema"}],
"response_format": {"type": "json_object"},
})).unwrap()),
).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]["message"]["content"], "ok");
}
const METRICS_KEY_ACME: &str = "completion-acme-secret";
const METRICS_KEY_BLUE: &str = "completion-blue-secret";
fn multi_key_metrics_state(metrics_token: Option<&str>) -> AppState {
let spec = format!(
"acme:{},blue:{}",
auth::sha256_hex(METRICS_KEY_ACME),
auth::sha256_hex(METRICS_KEY_BLUE),
);
let keyring = Box::leak(Box::new(auth::KeyStore::from_spec(&spec).unwrap()));
let mut st = fake_worker_state();
st.api_auth.keyring = Some(keyring);
st.metrics_auth = MetricsAuth::new(
true,
st.api_auth.configured(),
metrics_token.map(str::to_string),
);
{
let mut metrics = st.metrics.lock().unwrap();
metrics.admitted = 17;
metrics.prompt_tokens_in = 400;
metrics.cached_tokens_in = 60;
metrics.prefix_hits = 2;
metrics.prefix_misses = 3;
metrics.prefix_inserts = 5;
metrics.prefix_evictions = 7;
metrics.prefix_hit_tokens = 11;
metrics.lcp_hist[4] = 13;
metrics.ns_tokens.insert("t:acme".into(), [100, 40]);
metrics.ns_tokens.insert("t:blue".into(), [300, 20]);
metrics.adsd_suspect_total.insert("t:acme".into(), 1);
metrics.adsd_suspect_total.insert("t:blue".into(), 2);
metrics.prefix_entries = 29;
metrics.prefix_bytes = 31;
metrics.active_sessions = 3;
metrics.queued_requests = 5;
metrics.continuation_pool_entries = 7;
metrics.spec_pool_entries = 11;
metrics.cuda_driver_free_bytes = 13;
metrics.cuda_pool_reserved_bytes = 17;
metrics.cuda_pool_used_bytes = 19;
metrics.cuda_pool_cached_bytes = 23;
metrics.batch_size_last = 37;
metrics.spec.insert("m".into(), memra_engine::spec::SpecTelemetry {
rounds: 2,
drafted: 6,
accepted: 4,
..Default::default()
});
let mut spec_window = memra_engine::spec::SpecTelemetry {
rounds: 4,
drafted: 12,
accepted: 6,
..Default::default()
};
spec_window.pos_drafted[..3].copy_from_slice(&[4, 4, 4]);
spec_window.pos_accepted[..3].copy_from_slice(&[3, 2, 1]);
metrics.spec_window.insert("m".into(), spec_window);
metrics.constraint_compiler_fail_closed.insert(
"m".into(),
Arc::new(std::sync::atomic::AtomicBool::new(true)),
);
}
st
}
async fn metrics_json(st: AppState, bearer: &str) -> serde_json::Value {
let mut headers = HeaderMap::new();
headers.insert("authorization", format!("Bearer {bearer}").parse().unwrap());
let response = get_metrics(State(st), headers).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
serde_json::from_slice(&bytes).unwrap()
}
async fn yield_metrics_json(st: AppState, bearer: &str) -> serde_json::Value {
let mut headers = HeaderMap::new();
headers.insert("authorization", format!("Bearer {bearer}").parse().unwrap());
let response = yield_metrics(State(st), headers).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
serde_json::from_slice(&bytes).unwrap()
}
#[test]
fn exposed_open_bind_is_refused_before_server_start() {
assert!(validate_bind_security("127.0.0.1:8080", false, false).unwrap());
assert!(validate_bind_security("[::1]:8080", false, false).unwrap());
let err = validate_bind_security("0.0.0.0:8000", false, false).unwrap_err();
assert!(err.contains("refusing unauthenticated non-loopback bind"));
assert!(err.contains("MEMRA_API_KEY"));
assert!(err.contains("MEMRA_ALLOW_OPEN_BIND=1"));
assert!(validate_bind_security("[::]:8000", false, false).is_err());
assert!(!validate_bind_security("0.0.0.0:8000", true, false).unwrap());
assert!(!validate_bind_security("0.0.0.0:8000", false, true).unwrap());
}
#[tokio::test]
async fn keyed_metrics_require_and_accept_api_bearer() {
let mut st = fake_worker_state();
st.api_auth.single_key = Some(Arc::from("completion-secret"));
st.metrics_auth = MetricsAuth::new(true, st.api_auth.configured(), None);
let response = get_metrics(State(st.clone()), HeaderMap::new()).await;
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
let response = yield_metrics(State(st.clone()), HeaderMap::new()).await;
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer completion-secret".parse().unwrap());
assert_eq!(
get_metrics(State(st.clone()), headers.clone()).await.status(),
StatusCode::OK,
);
let body = metrics_json(st.clone(), "completion-secret").await;
assert!(
body.get("admitted").is_some(),
"the legacy single-key domain keeps cumulative counters",
);
assert!(
body.get("active_sessions").is_none(),
"a static completion key is not an operator metrics principal",
);
assert_eq!(yield_metrics(State(st), headers).await.status(), StatusCode::OK);
}
#[tokio::test]
async fn keyring_metrics_bearer_sees_only_its_tenant_rows() {
let st = multi_key_metrics_state(None);
let body = metrics_json(st.clone(), METRICS_KEY_ACME).await;
assert_eq!(
body.as_object().unwrap().len(),
2,
"completion metrics must contain only tenant-scoped rows",
);
let tenants = body["tenants"].as_object().unwrap();
assert_eq!(tenants.len(), 1);
assert_eq!(tenants["t:acme"]["prompt_tokens_in"], 100);
assert!(!tenants.contains_key("t:blue"));
let adsd = body["adsd_suspect_total"].as_object().unwrap();
assert_eq!(adsd.len(), 1);
assert_eq!(adsd["t:acme"], 1);
assert!(!adsd.contains_key("t:blue"));
let mut headers = HeaderMap::new();
headers.insert(
"authorization",
format!("Bearer {METRICS_KEY_ACME}").parse().unwrap(),
);
assert_eq!(
yield_metrics(State(st), headers).await.status(),
StatusCode::FORBIDDEN,
"the process-wide yield view requires an operator metrics token",
);
}
#[tokio::test]
async fn tenant_metrics_hide_capacity_and_aggregate_spec() {
let body = metrics_json(multi_key_metrics_state(None), METRICS_KEY_ACME).await;
for operator_only in [
"prefix_cache_entries",
"prefix_cache_bytes",
"active_sessions",
"queued_requests",
"continuation_pool_entries",
"spec_pool_entries",
"cuda_driver_free_bytes",
"cuda_pool_reserved_bytes",
"cuda_pool_used_bytes",
"cuda_pool_cached_bytes",
"constraint_compiler_fail_closed",
"serve_idle_seconds",
"spec",
"spec_tau",
"spec_accept_by_position",
"dual_pp",
"peer_probe_bypassed",
"peer_probe_boundary_copies",
"peer_probe_runtime_reprobes",
"peer_probe_runtime_failures",
"peer_probe_deferred_total",
"peer_probe_integrity_degraded",
"peer_probe_degraded_to_host_bounce",
] {
assert!(
body.get(operator_only).is_none(),
"tenant metrics must not expose operator field {operator_only}",
);
}
}
#[test]
fn populated_spec_acceptance_metrics_are_operator_only() {
for scope in [
MetricsScope::CompletionDomain,
MetricsScope::Tenant("t:acme".into()),
] {
let mut body = json!({});
insert_spec_acceptance_metrics(&mut body, &scope, || {
panic!("tenant scope evaluated the process-wide spec snapshot")
});
assert!(body.get("spec_tau").is_none(), "{scope:?} leaked spec tau");
assert!(body.get("spec_accept_by_position").is_none(),
"{scope:?} leaked the accept histogram");
}
let mut telemetry = memra_engine::spec::SpecTelemetry {
rounds: 4,
drafted: 12,
accepted: 6,
..Default::default()
};
telemetry.pos_drafted[..3].copy_from_slice(&[4, 4, 4]);
telemetry.pos_accepted[..3].copy_from_slice(&[3, 2, 1]);
let mut body = json!({});
insert_spec_acceptance_metrics(&mut body, &MetricsScope::All, || {
HashMap::from([("model-a".to_string(), telemetry)])
});
assert_eq!(body["spec_tau"]["model-a"], 1.5);
let histogram = &body["spec_accept_by_position"]["model-a"];
assert_eq!(histogram["window_seconds"], worker::SPEC_METRICS_WINDOW_S);
assert_eq!(histogram["rounds"], 4);
assert_eq!(histogram["offered"], json!([4, 4, 4]));
assert_eq!(histogram["accepted"], json!([3, 2, 1]));
assert_eq!(histogram["accept_rate"], json!([0.75, 0.5, 0.25]));
}
#[test]
fn populated_dual_pp_metrics_are_operator_only() {
let populated = DualPpMetricsSnapshot {
stage_ns: [1_000_000, 2_000_000, 3_000_000, 4_000_000],
stage_samples: [1, 1, 1, 1],
dropped_timing_samples: 0,
overlaps: 17,
slot_pairs: 19,
slot_uses: [19, 19],
slot_collisions: 0,
};
for scope in [
MetricsScope::CompletionDomain,
MetricsScope::Tenant("t:acme".into()),
] {
let mut body = json!({});
insert_dual_pp_metrics(&mut body, &scope, || populated);
assert!(body.get("dual_pp").is_none(), "{scope:?} leaked dual PP topology");
}
let mut body = json!({});
insert_dual_pp_metrics(&mut body, &MetricsScope::All, || populated);
assert_eq!(body["dual_pp"]["overlaps"], 17);
assert_eq!(body["dual_pp"]["slot_pairs"], 19);
assert_eq!(body["dual_pp"]["slot_uses"], json!([19, 19]));
assert_eq!(body["dual_pp"]["slot_collisions"], 0);
assert_eq!(body["dual_pp"]["cuda_event_spans"]["wave_a_stage0"]["mean_ms"], 1.0);
}
#[test]
fn peer_probe_metrics_are_operator_only() {
let populated = memra_engine::pp::PeerProbeMetrics {
bypassed: 1,
boundary_copies: 8_192,
runtime_probes: 1,
runtime_failures: 0,
deferred_total: 4,
integrity_degraded: true,
degraded_to_host_bounce: true,
};
for scope in [
MetricsScope::CompletionDomain,
MetricsScope::Tenant("t:acme".into()),
] {
let mut body = json!({});
insert_peer_probe_metrics(&mut body, &scope, || populated);
assert!(body.get("peer_probe_bypassed").is_none());
}
let mut body = json!({});
insert_peer_probe_metrics(&mut body, &MetricsScope::All, || populated);
assert_eq!(body["peer_probe_bypassed"], 1);
assert_eq!(body["peer_probe_boundary_copies"], 8_192);
assert_eq!(body["peer_probe_runtime_reprobes"], 1);
assert_eq!(body["peer_probe_runtime_failures"], 0);
assert_eq!(body["peer_probe_deferred_total"], 4);
assert_eq!(body["peer_probe_integrity_degraded"], true);
assert_eq!(body["peer_probe_degraded_to_host_bounce"], true);
}
#[tokio::test]
async fn prefix_aggregate_metrics_are_operator_only_but_tenant_ratio_remains() {
let tenant_body = metrics_json(multi_key_metrics_state(None), METRICS_KEY_ACME).await;
for operator_only in [
"lcp_histogram",
"cache_hit_token_ratio",
"prefix_cache_hits",
"prefix_cache_misses",
"prefix_cache_inserts",
"prefix_cache_evictions",
"prefix_cache_hit_tokens",
] {
assert!(
tenant_body.get(operator_only).is_none(),
"tenant metrics must not expose global prefix field {operator_only}",
);
}
assert_eq!(tenant_body["tenants"].as_object().unwrap().len(), 1);
assert_eq!(tenant_body["tenants"]["t:acme"]["prompt_tokens_in"], 100);
assert_eq!(tenant_body["tenants"]["t:acme"]["cached_tokens_in"], 40);
assert_eq!(tenant_body["tenants"]["t:acme"]["cache_hit_token_ratio"], 0.4);
let operator_body = metrics_json(
multi_key_metrics_state(Some("scrape-secret")),
"scrape-secret",
).await;
assert_eq!(operator_body["prefix_cache_hits"], 2);
assert_eq!(operator_body["prefix_cache_misses"], 3);
assert_eq!(operator_body["prefix_cache_inserts"], 5);
assert_eq!(operator_body["prefix_cache_evictions"], 7);
assert_eq!(operator_body["prefix_cache_hit_tokens"], 11);
assert_eq!(operator_body["cache_hit_token_ratio"], 0.15);
assert_eq!(operator_body["lcp_histogram"]["counts"][4], 13);
}
#[tokio::test]
async fn configured_metrics_token_is_exclusive_and_sees_all_tenants() {
let st = multi_key_metrics_state(Some("scrape-secret"));
let mut completion_headers = HeaderMap::new();
completion_headers.insert(
"authorization",
format!("Bearer {METRICS_KEY_ACME}").parse().unwrap(),
);
assert_eq!(
get_metrics(State(st.clone()), completion_headers.clone()).await.status(),
StatusCode::FORBIDDEN,
);
assert_eq!(
yield_metrics(State(st.clone()), completion_headers).await.status(),
StatusCode::FORBIDDEN,
);
let body = metrics_json(st.clone(), "scrape-secret").await;
let tenants = body["tenants"].as_object().unwrap();
assert_eq!(tenants.len(), 2);
assert!(tenants.contains_key("t:acme"));
assert!(tenants.contains_key("t:blue"));
assert_eq!(body["adsd_suspect_total"]["t:acme"], 1);
assert_eq!(body["adsd_suspect_total"]["t:blue"], 2);
assert_eq!(body["active_sessions"], 3);
assert_eq!(body["queued_requests"], 5);
assert_eq!(body["prefix_cache_bytes"], 31);
assert_eq!(body["cuda_driver_free_bytes"], 13);
assert_eq!(body["constraint_compiler_fail_closed"]["m"], 1);
assert_eq!(body["spec"]["m"]["drafted"], 6);
assert_eq!(body["spec_tau"]["m"], 1.5);
assert_eq!(body["spec_accept_by_position"]["m"]["accepted"], json!([3, 2, 1]));
let yield_body = yield_metrics_json(st, "scrape-secret").await;
assert_eq!(yield_body["batch_size_last"], 37);
}
#[tokio::test]
async fn metrics_token_protects_public_override_without_api_keys() {
let mut st = fake_worker_state();
st.metrics_auth = MetricsAuth::new(false, false, Some("scrape-secret".into()));
assert_eq!(
get_metrics(State(st.clone()), HeaderMap::new()).await.status(),
StatusCode::UNAUTHORIZED,
);
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer scrape-secret".parse().unwrap());
assert_eq!(
get_metrics(State(st.clone()), headers.clone()).await.status(),
StatusCode::OK,
);
assert_eq!(yield_metrics(State(st), headers).await.status(), StatusCode::OK);
}
#[tokio::test]
async fn no_key_loopback_metrics_remain_open_for_development() {
let mut st = fake_worker_state();
st.metrics_auth = MetricsAuth::new(true, false, None);
let response = get_metrics(State(st.clone()), HeaderMap::new()).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert!(
body.get("active_sessions").is_some(),
"no-key loopback development keeps full operator visibility",
);
assert_eq!(
yield_metrics(State(st), HeaderMap::new()).await.status(),
StatusCode::OK,
);
}
#[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::try_acquire(
counts.clone(), lanes::Lane::Interactive, tenants.clone(), "acme", None)
.unwrap();
let (g2, n2, t2) = InflightGuard::try_acquire(
counts.clone(), lanes::Lane::Interactive, tenants.clone(), "acme", None)
.unwrap();
assert_eq!((n1, n2), (1, 2));
assert_eq!((t1, t2), (1, 2));
let (gj, nj, tj) = InflightGuard::try_acquire(
counts.clone(), lanes::Lane::Judge, tenants.clone(), "blue", None)
.unwrap();
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_concurrency_cap_is_atomic_across_arrivals() {
let counts: InflightCounts = Arc::new(Default::default());
let tenants: TenantGauge = Arc::new(Default::default());
let start = Arc::new(std::sync::Barrier::new(3));
let attempted = Arc::new(std::sync::Barrier::new(3));
let mut joins = Vec::new();
for _ in 0..2 {
let counts = counts.clone();
let tenants = tenants.clone();
let start = start.clone();
let attempted = attempted.clone();
joins.push(std::thread::spawn(move || {
start.wait();
let result = InflightGuard::try_acquire(
counts, lanes::Lane::Interactive, tenants, "preview_001", Some(1));
let won = result.is_ok();
attempted.wait(); drop(result);
won
}));
}
start.wait();
attempted.wait();
let wins = joins.into_iter()
.map(|join| join.join().unwrap())
.filter(|won| *won)
.count();
assert_eq!(wins, 1, "exactly one simultaneous request may pass cap=1");
assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 0);
assert!(tenants.lock().unwrap().is_empty());
}
#[tokio::test]
async fn tenant_concurrency_cap_rejects_before_worker_admission() {
let st = fake_worker_state();
let tenant = auth::TenantCtx {
tenant: "preview_001".into(),
lane_class: auth::LaneClass::Interactive,
rate_limit: Some(1),
};
let first_env = Envelope::new(true);
let (guard, first_rl) = match acquire_request_slot(
&st, lanes::Lane::Interactive, &tenant, &first_env)
{
Ok(slot) => slot,
Err(_) => panic!("the first request must acquire the tenant slot"),
};
assert_eq!((first_rl.limit, first_rl.remaining), (1, 0));
let second_env = Envelope::new(true);
let response = match acquire_request_slot(
&st, lanes::Lane::Interactive, &tenant, &second_env)
{
Err(response) => response,
Ok(_) => panic!("the second request must be rejected at the tenant cap"),
};
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(response.headers()["retry-after"], "2");
assert_eq!(response.headers()["retry-after-ms"], "2000");
assert_eq!(response.headers()["x-ratelimit-limit"], "1");
assert_eq!(response.headers()["x-ratelimit-remaining"], "0");
assert_eq!(response.headers()["x-request-id"], second_env.id);
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 1,
"rejected request must not consume a lane slot");
assert_eq!(
st.tenant_inflight.lock().unwrap().get("preview_001").copied(), Some(1),
"rejected request must not increment the tenant gauge");
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"]["type"], "rate_limit_error");
assert_eq!(payload["error"]["code"], "rate_limit_exceeded");
assert!(payload["error"]["message"].as_str().unwrap()
.contains("concurrent request limit"));
drop(guard);
let _ = InflightGuard::try_acquire(
st.inflight.clone(), lanes::Lane::Interactive,
st.tenant_inflight.clone(), "preview_001", Some(1))
.expect("slot must reopen after the in-flight request completes");
}
#[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);
}
#[tokio::test]
async fn handler_layer_refusals_are_openai_objects_with_x_should_retry() {
let hdr = |v: &str| {
let mut h = axum::http::HeaderMap::new();
h.insert("x-lane", axum::http::HeaderValue::from_str(v).unwrap());
h
};
let body = |resp: Response| async move {
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap()
};
let resp = lane_for_tenant(&hdr("turbo"), &auth::TenantCtx::default_tenant()).unwrap_err();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
let payload = body(resp).await;
assert!(payload["error"].is_object(), "bare-string error body: {payload}");
assert_eq!(payload["error"]["type"], "invalid_request_error");
assert_eq!(payload["error"]["param"], "x-lane");
assert_eq!(payload["error"]["code"], "invalid_lane");
let batch = auth::TenantCtx {
tenant: "bulk".into(), lane_class: auth::LaneClass::Batch, rate_limit: None,
};
let resp = lane_for_tenant(&hdr("interactive"), &batch).unwrap_err();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
let payload = body(resp).await;
assert_eq!(payload["error"]["type"], "authentication_error");
assert_eq!(payload["error"]["param"], "x-lane");
}
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(), None,
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(), None,
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 handlers_sync_worker_truth_usage_and_cost_before_terminal_response() {
let _l = DRAIN_LOCK.lock().unwrap();
let dir = std::env::temp_dir().join(format!("memra-handler-ledger-{}", gen_hex128()));
std::fs::create_dir(&dir).unwrap();
let path = dir.join("requests.jsonl");
let metadata = OpenRouterModelMetadata {
pricing: OpenRouterPricing {
prompt: Some("0.000000289".into()),
cached_prompt: Some("0.0000000289".into()),
completion: Some("0.000002".into()),
..Default::default()
},
..Default::default()
};
let mut st = fake_worker_state();
st.request_ledger = Some(ledger::Ledger::for_test(&path, "m", &metadata));
let nonstream = chat_completions(
State(st.clone()),
HeaderMap::new(),
None,
Json(
serde_json::from_value(json!({
"model": "m",
"messages": [{"role": "user", "content": "t"}],
}))
.unwrap(),
),
)
.await;
assert_eq!(nonstream.status(), StatusCode::OK);
let nonstream_id = nonstream.headers()["x-request-id"].to_str().unwrap().to_string();
let stream = completions(
State(st),
HeaderMap::new(),
None,
Json(
serde_json::from_value(json!({
"model": "m",
"prompt": "t",
"stream": true,
}))
.unwrap(),
),
)
.await;
assert_eq!(stream.status(), StatusCode::OK);
let stream_id = stream.headers()["x-request-id"].to_str().unwrap().to_string();
let _ = axum::body::to_bytes(stream.into_body(), usize::MAX).await.unwrap();
let rows: Vec<serde_json::Value> = std::fs::read_to_string(&path)
.unwrap()
.lines()
.map(|line| serde_json::from_str(line).unwrap())
.collect();
assert_eq!(rows.len(), 2);
assert_eq!(rows[0]["request_id"], nonstream_id);
assert_eq!(rows[1]["request_id"], stream_id);
for row in rows {
assert_eq!(row["outcome"], "completed");
assert_eq!(row["http_status"], 200);
assert_eq!(row["usage"]["prompt_tokens"], 1);
assert_eq!(row["usage"]["cached_prompt_tokens"], 0);
assert_eq!(row["usage"]["ordinary_prompt_tokens"], 1);
assert_eq!(row["usage"]["completion_tokens"], 1);
assert_eq!(row["cost_usd"]["total"], "0.0000022890");
}
std::fs::remove_dir_all(dir).unwrap();
}
#[tokio::test]
async fn completion_admission_returns_402_at_zero_then_debits_a_credited_balance() {
let _l = DRAIN_LOCK.lock().unwrap();
let dir = std::env::temp_dir().join(format!(
"memra-handler-budget-{}",
gen_hex128(),
));
std::fs::create_dir(&dir).unwrap();
let request_path = dir.join("requests.jsonl");
let budget_path = dir.join("budgets.toml");
std::fs::write(
&budget_path,
"[[budgets]]\ntenant = \"default\"\ncurrency = \"USD\"\nbalance_micro = 0\n",
)
.unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt as _;
std::fs::set_permissions(
&budget_path,
std::fs::Permissions::from_mode(0o640),
)
.unwrap();
}
let metadata = OpenRouterModelMetadata {
pricing: OpenRouterPricing {
prompt: Some("0.000000289".into()),
cached_prompt: Some("0.0000000289".into()),
completion: Some("0.000002".into()),
..Default::default()
},
..Default::default()
};
let ledger = ledger::Ledger::for_test_with_budgets(
&request_path,
&budget_path,
"m",
&metadata,
);
let budgets = ledger.budgets().unwrap();
let mut st = fake_worker_state();
st.request_ledger = Some(ledger);
st.tenant_budgets = Some(budgets.clone());
let metrics = get_metrics(State(st.clone()), HeaderMap::new()).await;
assert_eq!(metrics.status(), StatusCode::OK);
let metrics_body = axum::body::to_bytes(metrics.into_body(), usize::MAX)
.await
.unwrap();
let metrics_body: serde_json::Value = serde_json::from_slice(&metrics_body).unwrap();
assert_eq!(metrics_body["budget_source_reload_failed"], 0);
assert_eq!(metrics_body["budget_source_reload_consecutive"], 0);
assert_eq!(metrics_body["budget_source_available"], true);
let denied = completions(
State(st.clone()),
HeaderMap::new(),
None,
Json(
serde_json::from_value(json!({
"model": "m",
"prompt_ids": [1],
"max_tokens": 1,
}))
.unwrap(),
),
)
.await;
assert_eq!(denied.status(), StatusCode::PAYMENT_REQUIRED);
let denied_body = axum::body::to_bytes(denied.into_body(), usize::MAX)
.await
.unwrap();
let denied_body: serde_json::Value = serde_json::from_slice(&denied_body).unwrap();
assert_eq!(denied_body["error"]["type"], "insufficient_balance");
assert_eq!(denied_body["error"]["code"], "insufficient_balance");
budgets.credit("default", 10, "handler-credit-1").unwrap();
let admitted = completions(
State(st.clone()),
HeaderMap::new(),
None,
Json(
serde_json::from_value(json!({
"model": "m",
"prompt_ids": [1],
"max_tokens": 1,
}))
.unwrap(),
),
)
.await;
assert_eq!(admitted.status(), StatusCode::OK);
assert_eq!(
budgets.balance("default").unwrap().unwrap().balance_micro,
7
);
let rows: Vec<serde_json::Value> = std::fs::read_to_string(&request_path)
.unwrap()
.lines()
.map(|line| serde_json::from_str(line).unwrap())
.collect();
assert_eq!(rows.len(), 2);
assert_eq!(rows[0]["http_status"], 402);
assert!(rows[0]["budget"].is_null());
assert_eq!(rows[1]["http_status"], 200);
assert_eq!(rows[1]["cost_usd"]["total"], "0.0000022890");
assert_eq!(rows[1]["budget"]["debit_micro"], 3);
std::fs::remove_dir_all(dir).unwrap();
}
#[tokio::test]
async fn streaming_client_disconnect_records_partial_usage_and_cost() {
let _l = DRAIN_LOCK.lock().unwrap();
let dir = std::env::temp_dir().join(format!(
"memra-handler-ledger-disconnect-{}",
gen_hex128(),
));
std::fs::create_dir(&dir).unwrap();
let path = dir.join("requests.jsonl");
let metadata = OpenRouterModelMetadata {
pricing: OpenRouterPricing {
prompt: Some("0.000000289".into()),
cached_prompt: Some("0.0000000289".into()),
completion: Some("0.000002".into()),
..Default::default()
},
..Default::default()
};
let mut st = fake_worker_state_with_steps(
4,
std::time::Duration::from_millis(100),
);
st.request_ledger = Some(ledger::Ledger::for_test(&path, "m", &metadata));
let response = completions(
State(st),
HeaderMap::new(),
None,
Json(
serde_json::from_value(json!({
"model": "m",
"prompt": "disconnect after one delta",
"stream": true,
}))
.unwrap(),
),
)
.await;
assert_eq!(response.status(), StatusCode::OK);
let request_id = response.headers()["x-request-id"].to_str().unwrap().to_string();
let mut body = Box::pin(response.into_body().into_data_stream());
let first = std::future::poll_fn(|cx| body.as_mut().poll_next(cx))
.await
.expect("stream ended before first delta")
.expect("stream body failed");
assert!(is_sse_data_frame(&first), "first frame was not SSE data: {first:?}");
drop(body);
let row: serde_json::Value = serde_json::from_str(
std::fs::read_to_string(&path).unwrap().trim(),
)
.unwrap();
eprintln!("[disconnect-cell] receipt={row}");
assert_eq!(row["request_id"], request_id);
assert_eq!(row["outcome"], "abandoned");
assert_eq!(row["http_status"], 499);
assert_eq!(row["usage"]["prompt_tokens"], 1);
assert_eq!(row["usage"]["cached_prompt_tokens"], 0);
assert_eq!(row["usage"]["ordinary_prompt_tokens"], 1);
assert_eq!(row["usage"]["completion_tokens"], 1);
assert_eq!(row["usage"]["total_tokens"], 2);
assert_eq!(row["cost_usd"]["total"], "0.0000022890");
std::fs::remove_dir_all(dir).unwrap();
}
#[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(), None,
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
let ra = resp.headers()["retry-after"].to_str().unwrap().to_string();
let ra_s: u64 = ra.parse().expect("Retry-After must be integer delay-seconds");
assert!(ra_s > 0 && ra_s <= 60, "Retry-After {ra_s}s is outside the honored window");
let ra_ms: u64 = resp.headers()["retry-after-ms"].to_str().unwrap().parse().unwrap();
assert_eq!(ra_ms, ra_s * 1000, "the two retry headers must agree");
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"));
assert_eq!(payload["error"]["type"], "server_error");
assert_eq!(payload["error"]["code"], "draining");
let resp = completions(State(st.clone()), axum::http::HeaderMap::new(), None,
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_live(State(st.clone())).await.into_response();
assert_eq!(resp.status(), StatusCode::OK, "a drain must not look like a liveness fault");
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");
let resp = health_ready(State(st.clone())).await.into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
let retry_s = drain_deadline_s().clamp(1, 60);
let retry_s_text = retry_s.to_string();
let retry_ms_text = (retry_s * 1000).to_string();
assert_eq!(retry_after(&resp).as_deref(), Some(retry_s_text.as_str()));
assert_eq!(resp.headers().get("retry-after-ms").unwrap(),
retry_ms_text.as_str());
assert_ne!(resp.headers().get("x-should-retry")
.and_then(|v| v.to_str().ok()), Some("false"));
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"], "not_ready");
assert!(payload["detail"].as_str().unwrap().contains("draining"));
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let resp = chat_completions(State(st.clone()), axum::http::HeaderMap::new(), None,
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn health_is_green_only_while_the_worker_is_alive() {
let _l = DRAIN_LOCK.lock().unwrap();
let st = fake_worker_state();
let resp = health_live(State(st.clone())).await.into_response();
assert_eq!(resp.status(), StatusCode::OK);
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"], "ok");
assert_eq!(payload["worker"]["phase"], "idle");
assert!(payload["worker"]["stall_threshold_ms"].as_u64().unwrap() > 0);
let ready = health_ready(State(st.clone())).await.into_response();
assert_eq!(ready.status(), StatusCode::OK);
st.health.mark_dead("worker thread panicked: test-injected");
let resp = health_live(State(st.clone())).await.into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE,
"a dead worker MUST NOT report a healthy liveness");
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"], "unhealthy");
assert!(payload["detail"].as_str().unwrap().contains("test-injected"),
"cause not surfaced: {payload}");
let ready = health_ready(State(st.clone())).await.into_response();
assert_eq!(ready.status(), StatusCode::SERVICE_UNAVAILABLE, "dead is also not ready");
st.health.mark_ready();
assert_eq!(health_live(State(st.clone())).await.into_response().status(),
StatusCode::OK, "mark_ready must clear the latch (a successful respawn)");
}
#[tokio::test]
async fn readyz_peer_probe_integrity_is_present_and_advisory() {
let _l = DRAIN_LOCK.lock().unwrap();
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let st = fake_worker_state();
let ready = health_ready(State(st.clone())).await.into_response();
assert_eq!(ready.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(ready.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["peer_probe_integrity"], "ok");
st.health.note_peer_probe_deferral(2, false);
let deferred = health_ready(State(st.clone())).await.into_response();
assert_eq!(deferred.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(deferred.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["peer_probe_integrity"], "deferred_2");
st.health.note_peer_probe_deferral(4, true);
let degraded = health_ready(State(st.clone())).await.into_response();
assert_eq!(degraded.status(), StatusCode::OK,
"peer degradation is advisory while plain serving remains healthy");
let bytes = axum::body::to_bytes(degraded.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["peer_probe_integrity"], "degraded");
st.health.mark_dead("test-injected worker failure");
let unready = health_ready(State(st)).await.into_response();
assert_eq!(unready.status(), StatusCode::SERVICE_UNAVAILABLE);
let bytes = axum::body::to_bytes(unready.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["peer_probe_integrity"], "degraded",
"the advisory field must also survive an unrelated readiness failure");
}
#[tokio::test]
async fn liveness_failure_obeys_the_retry_contract() {
let _l = DRAIN_LOCK.lock().unwrap();
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let st = fake_worker_state();
st.health.mark_dead("worker thread panicked: retry-contract-test");
let resp = health_live(State(st)).await.into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(retry_after(&resp).as_deref(), Some("2"));
assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
assert_ne!(resp.headers().get("x-should-retry")
.and_then(|v| v.to_str().ok()), Some("false"));
}
#[tokio::test]
async fn readiness_failure_obeys_the_retry_contract() {
let _l = DRAIN_LOCK.lock().unwrap();
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let st = fake_worker_state();
st.health.mark_dead("worker thread panicked: retry-contract-test");
let resp = health_ready(State(st)).await.into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(retry_after(&resp).as_deref(), Some("2"));
assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
assert_ne!(resp.headers().get("x-should-retry")
.and_then(|v| v.to_str().ok()), Some("false"));
}
#[tokio::test]
async fn a_wedged_gpu_flips_health_even_though_the_worker_thread_is_fine() {
let _l = DRAIN_LOCK.lock().unwrap();
DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
let st = fake_worker_state();
assert_eq!(health_live(State(st.clone())).await.into_response().status(), StatusCode::OK);
st.health.mark_gpu_fault("nvidia-smi probe exceeded 10s deadline (GSP hang class)");
let resp = health_live(State(st.clone())).await.into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
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["detail"].as_str().unwrap().contains("probe exceeded"));
st.health.mark_ready();
assert_eq!(health_live(State(st.clone())).await.into_response().status(),
StatusCode::SERVICE_UNAVAILABLE,
"a GPU fault must not be cleared by an in-process respawn");
}
#[test]
fn v1_models_entry_keeps_catalog_shape_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()),
effort_levels: false,
gemma_think: false,
chat_temperature_default: None,
chat_top_p_default: None,
};
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());
}
#[test]
fn models_openai_default_body_stays_byte_identical() {
let body = models_openai_body(&["main".into(), "judge".into()]);
let bytes = serde_json::to_vec(&body).unwrap();
assert_eq!(
bytes,
br#"{"object":"list","data":[{"id":"main","object":"model"},{"id":"judge","object":"model"}]}"#
);
}
#[test]
fn openrouter_models_entry_serializes_complete_metadata() {
let metadata = OpenRouterMetadataFile::from_toml(
r#"
[models.main]
hugging_face_id = "Qwen/Qwen3.6-27B"
created = 1786032000
quantization = "nvfp4"
description = "Qwen3.6 27B served by memra."
max_prompt_length = 245760
max_output_length = 16384
is_ready = true
is_free = false
discount_to_user = 0.1
openrouter_slug = "qwen/qwen3.6-27b"
datacenters = [{ country_code = "US", region = "us-east-1" }]
zdr = true
hipaa = false
[models.main.pricing]
prompt = "0.000000234"
cached_prompt = "0.0000000585"
cache_write = "0.000000234"
completion = "0.000001872"
internal_reasoning = "0.000001872"
request = "0.01"
[models.main.capacity]
prompt_tpm = 1000000
cached_prompt_tpm = 2000000
completion_tpm = 500000
request_rpm = 1000
concurrency = 64
"#,
)
.unwrap();
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()),
..Default::default()
};
let entry = model_entry_openrouter("main", Some(&caps), metadata.get("main"));
assert_eq!(entry["schema_version"], "2.4");
assert_eq!(entry["id"], "main");
assert_eq!(entry["name"], "main");
assert_eq!(entry["hugging_face_id"], "Qwen/Qwen3.6-27B");
assert_eq!(entry["created"], 1786032000u64);
assert_eq!(entry["quantization"], "nvfp4");
assert_eq!(entry["tokenizer"], "qwen2");
assert_eq!(entry["description"], "Qwen3.6 27B served by memra.");
assert!(
entry.get("object").is_none(),
"OpenRouter schema 2.4 rejects unknown OpenAI fields"
);
let input = &entry["input_modalities"][0];
assert_eq!(input["type"], "text");
assert_eq!(
input["supported_inputs"]["max_context_length"]["value"],
262144
);
assert_eq!(
input["supported_inputs"]["max_prompt_length"]["value"],
245760
);
let input_prices = input["pricing"].as_array().unwrap();
let input_price = |kind: &str| {
input_prices
.iter()
.find(|price| price["type"] == kind)
.unwrap()
};
assert_eq!(input_price("prompt")["cost_usd"], "0.000000234");
assert_eq!(
input_price("cached_prompt")["cost_usd"],
"0.0000000585"
);
assert_eq!(input_price("cache_write")["cost_usd"], "0.000000234");
assert_eq!(input["capacity"][0]["value"], 1000000);
assert_eq!(input["capacity"][1]["value"], 2000000);
let output = &entry["output_modalities"][0];
assert_eq!(output["type"], "text");
assert_eq!(output["max_length"]["value"], 16384);
assert_eq!(output["streaming"], true);
assert_eq!(output["supported_parameters"]["tools"]["type"], "boolean");
assert_eq!(
output["supported_parameters"]["structured_outputs"]["type"],
"boolean"
);
assert_eq!(
output["supported_parameters"]["reasoning"]["type"],
"boolean"
);
assert_eq!(output["pricing"][0]["type"], "completion");
assert_eq!(output["pricing"][0]["cost_usd"], "0.000001872");
assert_eq!(output["pricing"][1]["type"], "internal_reasoning");
assert_eq!(output["capacity"][0]["value"], 500000);
assert_eq!(output["capacity"][1]["type"], "concurrency");
assert_eq!(output["capacity"][1]["value"], 64);
assert_eq!(entry["pricing"][0]["type"], "request");
assert_eq!(entry["pricing"][0]["cost_usd"], "0.01");
assert_eq!(entry["capacity"][0]["value"], 1000);
assert_eq!(entry["is_ready"], true);
assert_eq!(entry["is_free"], false);
assert_eq!(entry["discount_to_user"], 0.1);
assert_eq!(entry["openrouter"]["slug"], "qwen/qwen3.6-27b");
assert_eq!(entry["datacenters"][0]["country_code"], "US");
assert_eq!(entry["compliance"]["zdr"], true);
assert_eq!(entry["compliance"]["hipaa"], false);
}
#[test]
fn gateway_registry_generates_the_deployed_pair_shape() {
let metadata = OpenRouterMetadataFile::from_toml(include_str!(
"../../../deploy/gateway/q27-models.toml"
))
.unwrap();
let caps = ModelCaps {
tools_branch: true,
qwen_think: true,
think_switch: true,
chat_ok: true,
context_length: 8192,
tokenizer: "qwen2".into(),
instruct_type: Some("chatml".into()),
..Default::default()
};
let entry = model_entry_openrouter(
"qwen/qwen3.6-27b",
Some(&caps),
metadata.get("qwen/qwen3.6-27b"),
);
assert_eq!(entry["schema_version"], "2.4");
assert_eq!(entry["created"], 1777255064u64);
assert_eq!(entry["quantization"], "nvfp4");
assert_eq!(entry["is_ready"], true);
assert_eq!(
entry["input_modalities"][0]["supported_inputs"]["max_context_length"]["value"],
8192
);
assert_eq!(
entry["input_modalities"][0]["supported_inputs"]["max_prompt_length"]["value"],
7680
);
let prices = entry["input_modalities"][0]["pricing"]
.as_array()
.unwrap();
let price = |kind: &str| {
prices
.iter()
.find(|candidate| candidate["type"] == kind)
.unwrap()
};
assert_eq!(price("prompt")["cost_usd"], "0.000000280");
assert_eq!(price("cached_prompt")["cost_usd"], "0.000000070");
assert_eq!(
entry["output_modalities"][0]["pricing"][0]["cost_usd"],
"0.000002690"
);
assert_eq!(entry["output_modalities"][0]["capacity"][1]["value"], 4);
assert_eq!(entry["datacenters"][0]["country_code"], "CA");
let q35_entry = model_entry_openrouter(
"qwen/qwen3.6-35b-a3b",
Some(&caps),
metadata.get("qwen/qwen3.6-35b-a3b"),
);
assert_eq!(q35_entry["created"], 1777260255u64);
assert_eq!(q35_entry["quantization"], "int4");
assert_eq!(q35_entry["is_ready"], true);
assert_eq!(
q35_entry["input_modalities"][0]["pricing"][0]["cost_usd"],
"0.000000120"
);
assert_eq!(
q35_entry["input_modalities"][0]["pricing"][1]["cost_usd"],
"0.000000030"
);
assert_eq!(
q35_entry["output_modalities"][0]["pricing"][0]["cost_usd"],
"0.000001030"
);
assert_eq!(
q35_entry["output_modalities"][0]["capacity"][0]["value"],
4
);
for model in ["qwen/qwen3.6-27b", "qwen/qwen3.6-35b-a3b"] {
let openmodels = model_entry_openmodels(
model,
Some(&caps),
metadata.get(model),
)
.unwrap();
assert_eq!(openmodels["currency"], "USD");
assert_eq!(openmodels["is_ready"], true);
assert_eq!(openmodels["is_free"], false);
assert_eq!(openmodels["discount_to_user"], 0.0);
}
}
#[test]
fn q27_gateway_registry_limits_are_live_request_limits() {
let metadata_file = OpenRouterMetadataFile::from_toml(include_str!(
"../../../deploy/gateway/q27-models.toml"
))
.unwrap();
let metadata = metadata_file.get("qwen/qwen3.6-27b").unwrap();
let build = |value: serde_json::Value| {
let req: CompletionReq = serde_json::from_value(value).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_request(&req, tx, lanes::Lane::Interactive, None)
};
let mut omitted = build(json!({
"model": "qwen/qwen3.6-27b",
"prompt_ids": [1, 2, 3]
}));
apply_model_request_limits(&mut omitted, Some(metadata)).unwrap();
assert_eq!(omitted.params.max_new, 512);
assert_eq!(omitted.max_prompt_tokens, Some(7_680));
let mut too_much_output = build(json!({
"model": "qwen/qwen3.6-27b",
"prompt_ids": [1],
"max_tokens": 513
}));
let (message, param) =
apply_model_request_limits(&mut too_much_output, Some(metadata)).unwrap_err();
assert_eq!(param, "max_tokens");
assert!(message.contains("513"));
let mut oversized_allocation = build(json!({
"model": "qwen/qwen3.6-27b",
"prompt_ids": [1],
"max_tokens": 1,
"max_ctx": 8_201
}));
let (_, param) =
apply_model_request_limits(&mut oversized_allocation, Some(metadata)).unwrap_err();
assert_eq!(param, "max_ctx");
}
#[test]
fn openrouter_models_entry_omits_undeclared_optional_fields() {
let entry = model_entry_openrouter("minimal", None, None);
let object = entry.as_object().unwrap();
for field in [
"hugging_face_id",
"created",
"quantization",
"tokenizer",
"description",
"pricing",
"capacity",
"is_ready",
"is_free",
"discount_to_user",
"openrouter",
"datacenters",
"compliance",
] {
assert!(
!object.contains_key(field),
"optional field {field} must be absent, not null"
);
}
assert_eq!(entry["schema_version"], "2.4");
assert_eq!(entry["input_modalities"][0]["type"], "text");
assert!(
entry["input_modalities"][0]
.get("supported_inputs")
.is_none()
);
assert!(entry["input_modalities"][0].get("pricing").is_none());
assert!(entry["input_modalities"][0].get("capacity").is_none());
assert_eq!(entry["output_modalities"][0]["type"], "text");
assert_eq!(entry["output_modalities"][0]["streaming"], true);
assert!(
entry["output_modalities"][0]["supported_parameters"]
.is_object()
);
assert!(entry["output_modalities"][0].get("max_length").is_none());
assert!(entry["output_modalities"][0].get("pricing").is_none());
assert!(entry["output_modalities"][0].get("capacity").is_none());
}
#[test]
fn openmodels_entry_serializes_standard_provider_shape() {
let metadata = OpenRouterMetadataFile::from_toml(
r#"
[models."qwen/qwen3.6-27b"]
created = 1786032000
max_output_length = 16384
is_ready = true
is_free = false
discount_to_user = 0.05
[models."qwen/qwen3.6-27b".pricing]
prompt = "0.000000291"
cached_prompt = "0.000000291"
completion = "0.000002763"
request = "0"
"#,
)
.unwrap();
let caps = ModelCaps {
tools_branch: true,
qwen_think: true,
chat_ok: true,
context_length: 262144,
..Default::default()
};
let entry = model_entry_openmodels(
"qwen/qwen3.6-27b",
Some(&caps),
metadata.get("qwen/qwen3.6-27b"),
)
.unwrap();
assert_eq!(entry["id"], "qwen/qwen3.6-27b");
assert_eq!(entry["name"], "qwen/qwen3.6-27b");
assert_eq!(entry["created"], 1786032000u64);
assert_eq!(entry["input_modalities"], json!(["text"]));
assert_eq!(entry["output_modalities"], json!(["text"]));
assert_eq!(entry["context_length"], 262144u64);
assert_eq!(entry["max_output_length"], 16384u64);
assert_eq!(entry["currency"], "USD");
assert_eq!(entry["pricing"]["prompt"], "0.000000291");
assert_eq!(entry["pricing"]["completion"], "0.000002763");
assert_eq!(entry["pricing"]["input_cache_read"], "0.000000291");
assert_eq!(entry["pricing"]["request"], "0");
assert_eq!(
entry["supported_features"],
json!(["tool_calling", "reasoning"])
);
assert_eq!(entry["is_ready"], true);
assert_eq!(entry["is_free"], false);
assert_eq!(entry["discount_to_user"], 0.05);
assert!(entry.get("schema_version").is_none());
assert!(entry.get("quantization").is_none());
}
#[test]
fn openmodels_entry_rejects_missing_operator_metadata() {
let caps = ModelCaps {
context_length: 262144,
..Default::default()
};
let error = model_entry_openmodels("qwen/qwen3.6-27b", Some(&caps), None).unwrap_err();
assert_eq!(
error,
"OpenModels feed requires MEMRA_MODEL_METADATA for model \"qwen/qwen3.6-27b\""
);
}
#[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");
}
}