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 toolcall;
mod worker;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::mpsc::Sender;
use axum::{
Json, Router,
extract::{Query, State},
http::StatusCode,
response::{sse::{Event as SseEvent, Sse}, IntoResponse, Response},
routing::{get, post},
};
use serde::{Deserialize, Serialize};
use serde_json::json;
use memra_engine::decode::GenParams;
use memra_engine::sampler::SamplerConfig;
use memra_tokenizer::chat::{ThinkMode, ToolCall as TmplToolCall, Turn as TmplTurn};
use toolcall::{ParsedToolCall, Piece, ToolStreamParser};
use worker::{Cmd, Event, ModelCaps, Request, SharedMetrics};
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>>,
metrics: SharedMetrics,
started: u64,
inflight: InflightCounts,
tenant_inflight: TenantGauge,
health: health::SharedHealth,
bg: Option<(Arc<darklane::BgJobState>, &'static str)>,
}
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 = "default_temperature")]
temperature: f32,
#[serde(default = "one")]
top_p: f32,
#[serde(default)]
top_k: usize,
#[serde(default)]
min_p: f32,
#[serde(default)]
frequency_penalty: f32,
#[serde(default)]
presence_penalty: f32,
#[serde(default = "one")]
repetition_penalty: f32,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
stop: StopSequences,
#[serde(default)]
stream: bool,
#[serde(default)]
max_ctx: Option<usize>,
#[serde(default)]
response_format: Option<serde_json::Value>,
#[serde(default)]
logit_bias: Option<serde_json::Value>,
#[serde(default)]
logprobs: Option<serde_json::Value>,
#[serde(default)]
top_logprobs: Option<usize>,
#[serde(default)]
n: Option<usize>,
#[serde(default)]
tools: Vec<serde_json::Value>,
#[serde(default)]
tool_choice: Option<serde_json::Value>,
#[serde(default)]
reasoning_effort: Option<String>,
#[serde(default)]
reasoning: Option<serde_json::Value>,
#[serde(default)]
include_reasoning: Option<bool>,
#[serde(default)]
cache_salt: Option<String>,
#[serde(default)]
session_id: Option<String>,
#[serde(default)]
user: Option<String>,
}
fn one() -> f32 { 1.0 }
fn default_temperature() -> f32 { 1.0 }
#[derive(Serialize)]
struct CompletionResp {
model: String,
text: String,
tokens: Vec<u32>,
stop_reason: String,
n_tokens: usize,
prompt_tokens: usize,
cached_tokens: usize,
elapsed_s: f64,
}
fn usage_json(n_prompt: usize, n_tokens: usize, n_cached: usize, elapsed_s: f64,
spec: Option<worker::SpecUsage>) -> serde_json::Value {
let mut u = json!({
"prompt_tokens": n_prompt,
"completion_tokens": n_tokens,
"total_tokens": n_prompt + n_tokens,
"prompt_tokens_details": { "cached_tokens": n_cached },
"elapsed_s": elapsed_s,
});
if let Some(sp) = spec {
u["spec"] = json!({
"rounds": sp.rounds,
"drafted": sp.drafted,
"accepted": sp.accepted,
"acceptance_rate": if sp.drafted > 0 {
sp.accepted as f64 / sp.drafted as f64 } else { 0.0 },
});
}
u
}
const SYSTEM_FINGERPRINT: &str = concat!("memra-", env!("MEMRA_BUILD_SHA"));
fn gen_hex128() -> String {
use std::hash::{BuildHasher, Hasher};
static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let n = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let t = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let mut h1 = std::collections::hash_map::RandomState::new().build_hasher();
h1.write_u64(n);
h1.write_u64(t);
let mut h2 = std::collections::hash_map::RandomState::new().build_hasher();
h2.write_u64(t.rotate_left(17));
h2.write_u64(n);
format!("{:016x}{:016x}", h1.finish(), h2.finish())
}
#[derive(Clone)]
struct Envelope {
id: String,
created: u64,
}
impl Envelope {
fn new(chat: bool) -> Self {
Envelope {
id: format!("{}-{}", if chat { "chatcmpl" } else { "cmpl" }, gen_hex128()),
created: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0),
}
}
fn stamp(&self, mut v: serde_json::Value) -> serde_json::Value {
v["id"] = json!(self.id);
v["created"] = json!(self.created);
v["system_fingerprint"] = json!(SYSTEM_FINGERPRINT);
v
}
}
fn with_request_id(id: &str, mut resp: Response) -> Response {
if let Ok(v) = axum::http::HeaderValue::from_str(id) {
resp.headers_mut()
.insert(axum::http::HeaderName::from_static("x-request-id"), v);
}
resp
}
fn openai_compat() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| {
match std::env::var("MEMRA_COMPAT").as_deref() {
Ok("openai") => true,
Ok(_) => false,
Err(_) => std::env::var("MEMRA_API_KEY").is_ok(),
}
})
}
fn cache_namespace(cache_salt: &Option<String>) -> String {
cache_salt.clone().unwrap_or_default()
}
fn affinity_key(
session_id: &Option<String>,
user: &Option<String>,
headers: &axum::http::HeaderMap,
) -> Option<String> {
let clean = |s: &str| -> Option<String> {
let t = s.trim();
if t.is_empty() { None } else { Some(t.to_string()) }
};
session_id.as_deref().and_then(clean)
.or_else(|| user.as_deref().and_then(clean))
.or_else(|| headers.get("x-session-id")
.and_then(|v| v.to_str().ok())
.and_then(clean))
}
fn error_body(message: &str, etype: &str, param: Option<&str>, code: Option<&str>)
-> serde_json::Value {
json!({ "error": {
"message": message,
"type": etype,
"param": param,
"code": code,
} })
}
fn error_response(status: StatusCode, message: &str, etype: &str, param: Option<&str>)
-> Response {
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 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);
}
};
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) = 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),
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 addr = std::env::var("MEMRA_ADDR").unwrap_or_else(|_| "127.0.0.1:8080".into());
let listener = tokio::net::TcpListener::bind(&addr).await?;
eprintln!("[server] listening on http://{addr}");
health::sd_notify("READY=1\nSTATUS=serving");
let inflight = inflight_handle;
axum::serve(listener, app)
.with_graceful_shutdown(async move {
let mut sigterm = match tokio::signal::unix::signal(
tokio::signal::unix::SignalKind::terminate()) {
Ok(s) => s,
Err(err) => {
eprintln!("[server] WARN: no SIGTERM handler ({err}); drain disabled");
std::future::pending::<()>().await;
unreachable!()
}
};
sigterm.recv().await;
DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
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?;
if let Some(h) = bg_handle {
h.shutdown();
}
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 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
}
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(health_payload(&st, "ready", None))).into_response(),
Err(why) => retry_contract_response(
(StatusCode::SERVICE_UNAVAILABLE,
Json(health_payload(&st, "not_ready", Some(&why)))).into_response(),
Some(if is_draining {
drain_deadline_s()
} else {
worker::WORKER_RESPAWN_BACKOFF_BASE_S
}),
),
}
}
async fn get_metrics(State(st): State<AppState>) -> impl IntoResponse {
let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
let idle_s = darklane::ValleySignal::new(st.health.clone()).idle_seconds();
let mut body = json!({
"admitted": m.admitted,
"completed": m.completed,
"tokens_out": m.tokens_out,
"step_p50_ms": m.step_p50_ms,
"step_p99_ms": m.step_p99_ms,
"prompt_tokens_in": m.prompt_tokens_in,
"cached_tokens_in": m.cached_tokens_in,
"computed_tokens_in": m.prompt_tokens_in.saturating_sub(m.cached_tokens_in),
"cache_hit_token_ratio": if m.prompt_tokens_in > 0 {
m.cached_tokens_in as f64 / m.prompt_tokens_in as f64 } else { 0.0 },
"prefix_cache_hits": m.prefix_hits,
"prefix_cache_entries": m.prefix_entries,
"prefix_cache_bytes": m.prefix_bytes,
"prefix_cache_misses": m.prefix_misses,
"prefix_cache_inserts": m.prefix_inserts,
"prefix_cache_evictions": m.prefix_evictions,
"prefix_cache_hit_tokens": m.prefix_hit_tokens,
"lcp_histogram": {
"edges": worker::LCP_HIST_EDGES.to_vec(),
"counts": m.lcp_hist.to_vec(),
},
"serve_idle_seconds": (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()
.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();
body["tenants"] = serde_json::Value::Object(tenants);
}
if let Some((bg, mode)) = &st.bg {
body["bg"] = bg.to_json(mode);
}
let spec: serde_json::Map<String, serde_json::Value> = m.spec.iter().map(|(model, t)| {
let n_pos = t.pos_drafted.iter().rposition(|&d| d > 0).map_or(0, |p| p + 1);
(model.clone(), json!({
"rounds": t.rounds,
"drafted": t.drafted,
"accepted": t.accepted,
"acceptance_rate": if t.drafted > 0 {
t.accepted as f64 / t.drafted as f64 } else { 0.0 },
"tokens_per_round": if t.rounds > 0 {
(t.accepted + t.rounds) as f64 / t.rounds as f64 } else { 0.0 },
"pos_drafted": t.pos_drafted[..n_pos].to_vec(),
"pos_accepted": t.pos_accepted[..n_pos].to_vec(),
"accept_rate_per_pos": (0..n_pos).map(|j| if t.pos_drafted[j] > 0 {
t.pos_accepted[j] as f64 / t.pos_drafted[j] as f64 } else { 0.0 })
.collect::<Vec<f64>>(),
}))
}).collect();
if !spec.is_empty() {
body["spec"] = serde_json::Value::Object(spec);
}
Json(body)
}
#[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 })
}
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(schema) => bad_request(
&format!("unsupported models schema {schema:?}; expected openai or openrouter"),
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());
json!({
"id": name,
"name": name,
"object": "model",
"created": created,
"context_length": ctx,
"architecture": {
"modality": "text->text",
"tokenizer": tokenizer,
"instruct_type": instruct,
},
"pricing": {
"prompt": "0",
"completion": "0",
"request": "0",
"image": "0",
},
"top_provider": {
"context_length": ctx,
"max_completion_tokens": serde_json::Value::Null,
},
})
}
async fn list_models_v1(State(st): State<AppState>) -> impl IntoResponse {
let data: Vec<_> = st.models.iter()
.map(|m| model_entry_v1(m, st.caps.get(m), st.started))
.collect();
Json(json!({ "object": "list", "data": data }))
}
async fn yield_metrics(State(st): State<AppState>) -> impl IntoResponse {
let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
let lane = |i: usize| json!({
"admitted": m.lane_admitted[i], "shed": m.lane_shed[i],
"completed": m.lane_completed[i], "tokens_out": m.lane_tokens[i],
});
Json(json!({
"lanes": {
"interactive": lane(0), "judge": lane(1), "harvest": lane(2),
},
"interactive_step_ms": { "p50": m.step_p50_ms, "p99": m.step_p99_ms },
"batch_size_last": m.batch_size_last,
}))
}
async fn peek_shed(
lane: lanes::Lane,
mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>,
) -> Result<tokio::sync::mpsc::UnboundedReceiver<Event>, Response> {
if lane == lanes::Lane::Interactive {
return Ok(rx);
}
match rx.recv().await {
Some(Event::Error(e)) => 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)
}
}
}
fn build_request(req: &CompletionReq, tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>) -> Request {
let params = GenParams {
max_new: req.max_tokens.unwrap_or(worker::MAX_NEW_CTX_BOUNDED),
max_ctx: req.max_ctx,
eos: Vec::new(), };
let sampler_cfg = sampler_config(
req.temperature, req.top_k, req.top_p, req.min_p,
req.frequency_penalty, req.presence_penalty, req.repetition_penalty, req.seed);
Request {
model: req.model.clone(),
prompt_ids: req.prompt_ids.clone(),
prompt_text: req.prompt.clone(),
chat: req.chat,
chat_turns: Vec::new(),
tools_json: Vec::new(),
think: ThinkMode::Default,
reasoning_effort: None, params,
sampler_cfg,
stop_strings: req.stop.clone().into_vec(),
trace_id: req.trace_id.clone(),
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar: None, oom_retries: 0, tx,
}
}
struct ChatPlan {
request: Request,
parser: Option<ToolStreamParser>,
}
fn build_chat_request(req: ChatCompletionReq, caps: Option<&ModelCaps>,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
lane: lanes::Lane, affinity: Option<String>)
-> Result<ChatPlan, String> {
let tool_choice = parse_tool_choice(&req.tool_choice)?;
if let Some(c) = caps {
if !c.chat_ok {
return Err(format!(
"model {:?} has no chat template (checkpoint carries neither \
tokenizer_config.json chat_template nor chat_template.jinja) — \
/v1/chat/completions unavailable; use /v1/completions with a raw prompt",
req.model));
}
}
let (mut think, 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());
}
turns.push(TmplTurn { role: msg.role.clone(), content, tool_calls });
}
let has_tool_features = !tools_json.is_empty()
|| turns.iter().any(|t| t.role == "tool" || !t.tool_calls.is_empty());
if has_tool_features && !caps.map(|c| c.tools_branch).unwrap_or(false) {
return Err(format!("model {:?} chat template has no tools branch", req.model));
}
let think_open = caps.map(|c| c.qwen_think
&& !(think == ThinkMode::NoThink && c.think_switch)).unwrap_or(false);
let include_reasoning = req.include_reasoning.unwrap_or(true)
&& req.reasoning.as_ref()
.and_then(|r| r.get("exclude")).and_then(|v| v.as_bool()) != Some(true);
let parser = if !tools_json.is_empty() {
Some(ToolStreamParser::new(schemas, think_open)
.with_include_reasoning(include_reasoning))
} else if think_open {
Some(ToolStreamParser::reasoning_only(include_reasoning))
} else 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(
req.temperature, req.top_k, req.top_p, req.min_p,
req.frequency_penalty, req.presence_penalty, req.repetition_penalty,
req.seed),
stop_strings: req.stop.into_vec(),
trace_id: None,
cache_ns: cache_namespace(&req.cache_salt),
affinity,
lane,
grammar,
oom_retries: 0, tx,
},
parser,
})
}
fn authenticate(headers: &axum::http::HeaderMap) -> Result<auth::TenantCtx, Response> {
static SINGLE: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
let single = SINGLE.get_or_init(|| std::env::var("MEMRA_API_KEY").ok());
let bearer = headers.get("authorization")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "));
auth::authenticate_with(auth::global(), single.as_deref(), bearer).map_err(|why| {
match why {
auth::AuthDenied::Unknown => error_response(
StatusCode::UNAUTHORIZED, "invalid api key", "authentication_error", None),
auth::AuthDenied::Disabled => error_response(
StatusCode::FORBIDDEN, "api key is disabled", "authentication_error", None),
}
})
}
fn lane_for_tenant(headers: &axum::http::HeaderMap, tenant: &auth::TenantCtx)
-> Result<lanes::Lane, Response> {
let requested = match headers.get("x-lane").map(|v| v.to_str().unwrap_or("?")) {
None => None,
Some(v) => Some(lanes::Lane::parse(v).ok_or_else(|| {
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>) -> String {
let raw = cache_namespace(cache_salt);
if auth::global().is_some() {
auth::scope_namespace(&tenant.tenant, &raw)
} else {
raw
}
}
fn meter_admit(env: &Envelope, tenant: &auth::TenantCtx, model: &str, lane: lanes::Lane) {
eprintln!("[meter] admit id={} tenant={} lane={} model={:?}",
env.id, tenant.tenant, lane.as_str(), model);
}
async fn completions(State(st): State<AppState>, headers: axum::http::HeaderMap,
Json(req): Json<CompletionReq>) -> Response {
let env = Envelope::new(false);
let tenant = match authenticate(&headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
if let Err((msg, param)) = reject_unsupported(&[
("logit_bias", req.logit_bias.is_some(),
" (device-side sampling has no bias hook yet)"),
("logprobs", req.logprobs.is_some(), ""),
("n", req.n.is_some_and(|n| n != 1), " for n != 1 (single choice only)"),
("best_of", req.best_of.is_some_and(|n| n != 1), " (single choice only)"),
]) {
return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
}
let lane = match lane_for_tenant(&headers, &tenant) {
Ok(l) => l,
Err(resp) => return resp,
};
if draining() {
return with_request_id(&env.id, drain_response());
}
let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
Ok(slot) => slot,
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(&req, tx, lane, affinity);
request.cache_ns = tenant_namespace(&tenant, &req.cache_salt);
meter_admit(&env, &tenant, &model, lane);
let stop_strings = request.stop_strings.clone();
worker::PENDING_ADMITS.fetch_add(1, std::sync::atomic::Ordering::Release);
if st.cmd_tx.send(Cmd::Generate(Box::new(request))).is_err() {
worker::PENDING_ADMITS.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
return rl.attach(with_request_id(&env.id, worker_unavailable_response()));
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => return rl.attach(with_request_id(&env.id, resp)),
};
let resp = if stream {
sse_response(rx, model, false, None, env.clone(), stop_strings, Some(guard))
.into_response()
} else {
let resp = blocking_response(rx, model, false, stop_strings, None, env.clone()).await
.into_response();
drop(guard); resp
};
rl.attach(with_request_id(&env.id, resp))
}
async fn chat_completions(State(st): State<AppState>, headers: axum::http::HeaderMap,
Json(req): Json<ChatCompletionReq>) -> Response {
let env = Envelope::new(true);
let tenant = match authenticate(&headers) {
Ok(t) => t,
Err(resp) => return with_request_id(&env.id, resp),
};
if req.messages.is_empty() || req.messages.iter().any(|message| {
!matches!(message.role.as_str(), "system" | "user" | "assistant" | "tool")
}) {
return with_request_id(&env.id, bad_request(
"messages must use system/user/assistant/tool roles", Some("messages")));
}
if let Err((msg, param)) = reject_unsupported(&[
("logit_bias", req.logit_bias.is_some(),
" (device-side sampling has no bias hook yet)"),
("logprobs", req.logprobs.as_ref().is_some_and(|v| v.as_bool() != Some(false)), ""),
("top_logprobs", req.top_logprobs.is_some(), ""),
("n", req.n.is_some_and(|n| n != 1), " for n != 1 (single choice only)"),
]) {
return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
}
let lane = match lane_for_tenant(&headers, &tenant) {
Ok(l) => l,
Err(resp) => return resp,
};
let model = req.model.clone();
let stream = req.stream;
let cache_salt = req.cache_salt.clone();
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Event>();
let affinity = affinity_key(&req.session_id, &req.user, &headers);
let mut plan = match build_chat_request(req, st.caps.get(&model), tx, lane, affinity) {
Ok(plan) => plan,
Err(err) => {
return with_request_id(&env.id, bad_request(&err, None));
}
};
plan.request.cache_ns = tenant_namespace(&tenant, &cache_salt);
if draining() {
return with_request_id(&env.id, drain_response());
}
let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
Ok(slot) => slot,
Err(resp) => return resp,
};
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 st.cmd_tx.send(Cmd::Generate(Box::new(plan.request))).is_err() {
worker::PENDING_ADMITS.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
return rl.attach(with_request_id(&env.id, worker_unavailable_response()));
}
let rx = match peek_shed(lane, rx).await {
Ok(rx) => rx,
Err(resp) => return rl.attach(with_request_id(&env.id, resp)),
};
let resp = if stream {
sse_response(rx, model, true, plan.parser, env.clone(), stop_strings, Some(guard))
.into_response()
} else {
let resp = blocking_response(rx, model, true, stop_strings, plan.parser, env.clone())
.await.into_response();
drop(guard); resp
};
rl.attach(with_request_id(&env.id, resp))
}
fn sse_response(mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String, chat: bool,
mut parser: Option<ToolStreamParser>, env: Envelope,
stop_strings: Vec<String>, guard: Option<InflightGuard>)
-> Sse<impl futures_core::Stream<Item = Result<SseEvent, std::convert::Infallible>>> {
let mut scrub = (!stop_strings.is_empty() && (chat || openai_compat()))
.then(|| StopScrubber::new(stop_strings));
let stream = async_stream::stream! {
let _guard = guard;
let mut call_index: usize = 0;
let mut role_sent = false;
macro_rules! chat_chunk {
($delta:expr, $finish:expr) => {{
let mut delta = $delta;
if chat && !role_sent {
role_sent = true;
delta["role"] = json!("assistant");
}
env.stamp(json!({ "object": "chat.completion.chunk", "model": model,
"choices": [{ "index": 0, "delta": delta,
"finish_reason": $finish }] }))
.to_string()
}};
}
macro_rules! piece_chunks {
($piece:expr) => {{
let mut payloads: Vec<String> = Vec::new();
match $piece {
Piece::Content(text) => {
let text = match scrub.as_mut() {
Some(sc) => sc.push(&text),
None => text,
};
if !text.is_empty() {
payloads.push(chat_chunk!(json!({ "content": text }),
serde_json::Value::Null));
}
}
Piece::Reasoning(text) => payloads.push(
chat_chunk!(json!({ "reasoning": text }), serde_json::Value::Null)),
Piece::Call(call) => {
payloads.push(chat_chunk!(json!({ "tool_calls": [{
"index": call_index, "id": call.id, "type": "function",
"function": { "name": call.name, "arguments": "" } }] }),
serde_json::Value::Null));
payloads.push(chat_chunk!(json!({ "tool_calls": [{
"index": call_index,
"function": { "arguments": call.arguments } }] }),
serde_json::Value::Null));
call_index += 1;
}
}
payloads
}};
}
while let Some(ev) = rx.recv().await {
match ev {
Event::Token { id, text } => {
if let Some(p) = parser.as_mut() {
for piece in p.push(&text) {
for payload in piece_chunks!(piece) {
yield Ok(SseEvent::default().data(payload));
}
}
continue;
}
let text = match scrub.as_mut() {
Some(sc) => sc.push(&text),
None => text,
};
if text.is_empty() && scrub.is_some() {
continue; }
let payload = if chat {
chat_chunk!(json!({ "content": text }), serde_json::Value::Null)
} else if openai_compat() {
env.stamp(json!({ "object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": text, "finish_reason": null }] }))
.to_string()
} else {
json!({ "model": model, "id": id, "text": text }).to_string()
};
yield Ok(SseEvent::default().data(payload));
}
Event::Done { stop_reason, n_tokens, n_prompt, n_cached, elapsed_s, spec } => {
let mut finish = stop_reason_to_finish(&stop_reason);
if let Some(p) = parser.as_mut() {
for piece in p.finish() {
for payload in piece_chunks!(piece) {
yield Ok(SseEvent::default().data(payload));
}
}
if p.n_calls() > 0 { finish = "tool_calls"; }
}
if let Some(sc) = scrub.as_mut() {
let tail = sc.finish();
if !tail.is_empty() {
let payload = if chat {
chat_chunk!(json!({ "content": tail }),
serde_json::Value::Null)
} else {
env.stamp(json!({ "object": "text_completion",
"model": model,
"choices": [{ "index": 0, "text": tail,
"finish_reason": null }] })).to_string()
};
yield Ok(SseEvent::default().data(payload));
}
}
if chat || openai_compat() {
let usage = usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec);
let fin = if chat {
let mut v = env.stamp(json!({
"object": "chat.completion.chunk", "model": model,
"choices": [{ "index": 0, "delta": {},
"finish_reason": finish }],
"usage": usage }));
if !role_sent {
v["choices"][0]["delta"]["role"] = json!("assistant");
}
v
} else {
env.stamp(json!({ "object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": "",
"finish_reason": finish }],
"usage": usage }))
}.to_string();
yield Ok(SseEvent::default().data(fin));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
let payload = json!({
"stop_reason": stop_reason, "n_tokens": n_tokens,
"prompt_tokens": n_prompt, "cached_tokens": n_cached,
"elapsed_s": elapsed_s
}).to_string();
yield Ok(SseEvent::default().event("done").data(payload));
}
break;
}
Event::Error(err) => {
if chat || openai_compat() {
let payload = engine_error_body(&err).to_string();
yield Ok(SseEvent::default().data(payload));
yield Ok(SseEvent::default().data("[DONE]".to_string()));
} else {
let payload = engine_error_body(&err).to_string();
yield Ok(SseEvent::default().event("error").data(payload));
}
break;
}
}
}
};
Sse::new(stream).keep_alive(
axum::response::sse::KeepAlive::new()
.interval(std::time::Duration::from_secs(5)),
)
}
fn truncate_at_stop(text: &mut String, stop_strings: &[String]) {
if let Some(offset) = stop_strings.iter().filter_map(|stop| text.find(stop)).min() {
text.truncate(offset);
}
}
fn partial_stop_suffix(s: &str, tag: &str) -> usize {
let mut best = 0;
for (k, _) in tag.char_indices().skip(1) {
if k <= s.len() && s.ends_with(&tag[..k]) {
best = k;
}
}
best
}
struct StopScrubber {
stops: Vec<String>,
buf: String,
done: bool,
}
impl StopScrubber {
fn new(stops: Vec<String>) -> Self {
Self { stops, buf: String::new(), done: false }
}
fn push(&mut self, text: &str) -> String {
if self.done {
return String::new();
}
self.buf.push_str(text);
if let Some(i) = self.stops.iter().filter_map(|s| self.buf.find(s.as_str())).min() {
self.done = true;
let out = self.buf[..i].to_string();
self.buf.clear();
return out;
}
let keep = self.stops.iter()
.map(|s| partial_stop_suffix(&self.buf, s)).max().unwrap_or(0);
let emit_to = self.buf.len() - keep;
let out = self.buf[..emit_to].to_string();
self.buf.drain(..emit_to);
out
}
fn finish(&mut self) -> String {
if self.done {
self.buf.clear();
return String::new();
}
std::mem::take(&mut self.buf)
}
}
async fn blocking_response(mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>, model: String,
chat: bool, stop_strings: Vec<String>,
mut parser: Option<ToolStreamParser>, env: Envelope) -> Response {
let mut text = String::new();
let mut reasoning = String::new();
let mut tokens: Vec<u32> = Vec::new();
let mut calls: Vec<ParsedToolCall> = Vec::new();
let consume = |pieces: Vec<Piece>, text: &mut String, reasoning: &mut String,
calls: &mut Vec<ParsedToolCall>| {
for piece in pieces {
match piece {
Piece::Content(t) => text.push_str(&t),
Piece::Reasoning(t) => reasoning.push_str(&t),
Piece::Call(c) => calls.push(c),
}
}
};
while let Some(ev) = rx.recv().await {
match ev {
Event::Token { id, text: delta } => {
tokens.push(id);
match parser.as_mut() {
Some(p) => consume(p.push(&delta), &mut text, &mut reasoning, &mut calls),
None => text.push_str(&delta),
}
}
Event::Done { stop_reason, n_tokens, n_prompt, n_cached, elapsed_s, spec } => {
if let Some(p) = parser.as_mut() {
consume(p.finish(), &mut text, &mut reasoning, &mut calls);
}
truncate_at_stop(&mut text, &stop_strings);
let finish = if calls.is_empty() { stop_reason_to_finish(&stop_reason) }
else { "tool_calls" };
if chat {
let content = if !calls.is_empty() && text.is_empty() {
serde_json::Value::Null
} else {
serde_json::Value::String(text)
};
let mut message = json!({ "role": "assistant", "content": content });
if !reasoning.is_empty() {
message["reasoning"] = json!(reasoning);
message["reasoning_details"] = json!([{
"type": "reasoning.text", "text": reasoning }]);
}
if !calls.is_empty() {
message["tool_calls"] = serde_json::Value::Array(
calls.iter().map(tool_call_json).collect());
}
return Json(env.stamp(json!({
"object": "chat.completion", "model": model,
"choices": [{ "index": 0,
"message": message,
"finish_reason": finish }],
"usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
}))).into_response();
}
if openai_compat() {
return Json(env.stamp(json!({
"object": "text_completion", "model": model,
"choices": [{ "index": 0, "text": text,
"finish_reason": finish }],
"usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
}))).into_response();
}
return Json(CompletionResp {
model, text, tokens, stop_reason, n_tokens,
prompt_tokens: n_prompt, cached_tokens: n_cached, elapsed_s,
}).into_response();
}
Event::Error(err) => {
return engine_error_response(&err);
}
}
}
let e = worker::EngineError::overloaded(
"worker closed the stream without completing (worker restart in progress)");
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 chat_request_preserves_turns_and_openai_stop_forms() {
let payload = serde_json::json!({
"model": "plain_quant",
"messages": [
{"role": "system", "content": "rules"},
{"role": "user", "content": "task"},
{"role": "assistant", "content": "work"}
],
"max_tokens": 64,
"temperature": 0.0,
"stop": "<stop>"
});
let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
let request = plan.request;
assert!(plan.parser.is_none(), "no tools -> no parser (isolation contract)");
assert!(request.tools_json.is_empty());
assert_eq!(request.think, ThinkMode::Default);
assert_eq!(request.model, "plain_quant");
assert_eq!(request.params.max_new, 64);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}]
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
assert_eq!(plan.request.params.max_new, worker::MAX_NEW_CTX_BOUNDED);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"max_completion_tokens": 7
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap().request.params.max_new, 7);
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, None).params.max_new, worker::MAX_NEW_CTX_BOUNDED);
let turns: Vec<(String, String)> = request.chat_turns.iter()
.map(|t| (t.role.clone(), t.content.clone())).collect();
assert_eq!(turns, vec![
("system".into(), "rules".into()),
("user".into(), "task".into()),
("assistant".into(), "work".into()),
]);
assert!(request.chat_turns.iter().all(|t| t.tool_calls.is_empty()));
assert_eq!(request.stop_strings, vec!["<stop>"]);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"stop": ["a", "b"]
})).unwrap();
assert_eq!(req.stop.into_vec(), vec!["a", "b"]);
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
"stop": null
})).unwrap();
assert!(req.stop.into_vec().is_empty());
}
#[tokio::test]
async fn chat_response_has_openai_message_shape() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "hello".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 42, n_cached: 30, elapsed_s: 0.5,
spec: None,
}).unwrap();
drop(tx);
let response = blocking_response(rx, "plain_quant".into(), true, Vec::new(), None,
Envelope::new(true)).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(payload["object"], "chat.completion");
assert!(payload["id"].as_str().unwrap().starts_with("chatcmpl-"));
assert!(payload["created"].as_u64().unwrap() > 1_700_000_000);
assert!(payload["system_fingerprint"].as_str().unwrap().starts_with("memra-"));
assert_eq!(payload["choices"][0]["message"]["role"], "assistant");
assert_eq!(payload["choices"][0]["message"]["content"], "hello");
assert_eq!(payload["choices"][0]["finish_reason"], "stop");
assert_eq!(payload["usage"]["prompt_tokens"], 42);
assert_eq!(payload["usage"]["completion_tokens"], 1);
assert_eq!(payload["usage"]["total_tokens"], 43);
assert_eq!(payload["usage"]["prompt_tokens_details"]["cached_tokens"], 30);
assert!(payload["usage"].get("spec").is_none());
}
#[tokio::test]
async fn chat_usage_carries_spec_acceptance_summary() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "hello".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 42, n_cached: 0, elapsed_s: 0.5,
spec: Some(worker::SpecUsage { rounds: 10, drafted: 30, accepted: 21 }),
}).unwrap();
drop(tx);
let response = blocking_response(rx, "plain_quant".into(), true, Vec::new(), None,
Envelope::new(true)).await;
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
let sp = &payload["usage"]["spec"];
assert_eq!(sp["rounds"], 10);
assert_eq!(sp["drafted"], 30);
assert_eq!(sp["accepted"], 21);
assert!((sp["acceptance_rate"].as_f64().unwrap() - 0.7).abs() < 1e-9);
assert_eq!(payload["usage"]["total_tokens"], 43);
}
fn weather_request(extra: serde_json::Value) -> ChatCompletionReq {
let mut payload = serde_json::json!({
"model": "m",
"messages": [{"role": "user", "content": "Weather in Paris?"}],
"tools": [{"type": "function", "function": {
"name": "get_weather",
"description": "Get current weather",
"parameters": {"type": "object",
"properties": {"city": {"type": "string"},
"days": {"type": "integer"}},
"required": ["city"]}}}],
});
if let Some(obj) = extra.as_object() {
for (k, v) in obj { payload[k] = v.clone(); }
}
serde_json::from_value(payload).unwrap()
}
#[test]
fn tools_request_renders_client_key_order_and_arms_parser() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(json!({})), Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
assert!(plan.parser.is_some());
assert_eq!(plan.request.tools_json.len(), 1);
assert_eq!(plan.request.tools_json[0],
"{\"type\": \"function\", \"function\": {\"name\": \"get_weather\", \
\"description\": \"Get current weather\", \"parameters\": {\"type\": \"object\", \
\"properties\": {\"city\": {\"type\": \"string\"}, \"days\": {\"type\": \
\"integer\"}}, \"required\": [\"city\"]}}}");
}
#[test]
fn tool_choice_none_strips_tools_and_parser() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let plan = build_chat_request(weather_request(json!({"tool_choice": "none"})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).unwrap();
let mut p = plan.parser.expect("think-open chat arms the reasoning splitter");
let pieces = p.push("x</think>\n\n<tool_call> stays prose");
assert_eq!(pieces, vec![
Piece::Reasoning("x".into()),
Piece::Content("<tool_call> stays prose".into()),
]);
assert!(plan.request.tools_json.is_empty());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({"tool_choice": "required"})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).is_err());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({"tool_choice":
{"type": "function", "function": {"name": "get_weather"}}})),
Some(&tool_caps()), tx, lanes::Lane::Interactive, None).is_err());
}
#[test]
fn model_plan_accepts_st_dir_and_rejects_bogus_dir() {
let root = std::env::temp_dir().join(format!("memra_plan_test_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&root);
let st = root.join("st_single");
std::fs::create_dir_all(&st).unwrap();
std::fs::write(st.join("config.json"), "{}").unwrap();
std::fs::write(st.join("model.safetensors"), b"x").unwrap();
assert!(validate_model_path(st.to_str().unwrap()).is_ok());
let sh = root.join("st_sharded");
std::fs::create_dir_all(&sh).unwrap();
std::fs::write(sh.join("config.json"), "{}").unwrap();
std::fs::write(sh.join("model.safetensors.index.json"), "{}").unwrap();
assert!(validate_model_path(sh.to_str().unwrap()).is_ok());
let rp = root.join("repack");
std::fs::create_dir_all(&rp).unwrap();
std::fs::write(rp.join("manifest.json"), "{}").unwrap();
assert!(validate_model_path(rp.to_str().unwrap()).is_ok());
let bogus = root.join("bogus");
std::fs::create_dir_all(&bogus).unwrap();
let err = validate_model_path(bogus.to_str().unwrap()).unwrap_err();
assert!(err.contains("model.safetensors"), "error should say what is missing: {err}");
assert!(err.contains("manifest.json"), "error should mention the repack form: {err}");
let nc = root.join("no_config");
std::fs::create_dir_all(&nc).unwrap();
std::fs::write(nc.join("model.safetensors"), b"x").unwrap();
let err = validate_model_path(nc.to_str().unwrap()).unwrap_err();
assert!(err.contains("config.json"), "error should name config.json: {err}");
let err = validate_model_path(root.join("nowhere").to_str().unwrap()).unwrap_err();
assert!(err.contains("does not exist"), "{err}");
let f = root.join("model.gguf");
std::fs::write(&f, b"g").unwrap();
assert!(validate_model_path(f.to_str().unwrap()).is_ok());
let _ = std::fs::remove_dir_all(&root);
}
#[test]
fn chat_on_templateless_dir_checkpoint_is_rejected_with_clear_message() {
let caps = ModelCaps {
tools_branch: false, qwen_think: false, think_switch: false, chat_ok: false,
..Default::default() };
let payload = serde_json::json!({
"model": "st_model",
"messages": [{"role": "user", "content": "hello"}],
});
let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let err = match build_chat_request(req, Some(&caps), tx, lanes::Lane::Interactive, None) {
Err(e) => e,
Ok(_) => panic!("templateless dir checkpoint must reject chat"),
};
assert!(err.contains("no chat template"), "message should name the cause: {err}");
assert!(err.contains("/v1/completions"), "message should point at the raw-prompt escape hatch: {err}");
}
#[test]
fn tools_on_model_without_tools_branch_is_rejected() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let caps = ModelCaps { chat_ok: true, ..Default::default() };
assert!(build_chat_request(weather_request(json!({})), Some(&caps), tx, lanes::Lane::Interactive, None).is_err());
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_chat_request(weather_request(json!({})), None, tx, lanes::Lane::Interactive, None).is_err());
}
#[test]
fn reasoning_effort_maps_to_think_switch() {
for (extra, want) in [
(json!({}), ThinkMode::Default),
(json!({"reasoning_effort": "low"}), ThinkMode::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 affinity_key_honors_both_client_conventions_in_priority_order() {
use axum::http::HeaderMap;
let hdr = |v: &str| {
let mut h = HeaderMap::new();
h.insert("x-session-id", v.parse().unwrap());
h
};
let empty = HeaderMap::new();
let s = |v: &str| Some(v.to_string());
assert_eq!(affinity_key(&s("explicit"), &None, &empty), s("explicit"));
assert_eq!(affinity_key(&None, &s("openai-user"), &empty), s("openai-user"));
assert_eq!(affinity_key(&None, &None, &hdr("hdr-id")), s("hdr-id"));
assert_eq!(affinity_key(&s("a"), &s("b"), &hdr("c")), s("a"));
assert_eq!(affinity_key(&None, &s("b"), &hdr("c")), s("b"));
assert_eq!(affinity_key(&s(" "), &s(""), &hdr(" ")), None);
assert_eq!(affinity_key(&s(""), &s("real"), &empty), s("real"));
assert_eq!(affinity_key(&s(" padded "), &None, &empty), s("padded"));
assert_eq!(affinity_key(&None, &None, &empty), None);
}
#[test]
fn affinity_key_plumbs_to_the_worker_request_on_both_bodies() {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task", "session_id": "conv-1"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new());
assert_eq!(build_request(&req, tx, lanes::Lane::Interactive, key).affinity.as_deref(),
Some("conv-1"));
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "task"}],
"user": "conv-2"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new());
assert_eq!(build_chat_request(req, None, tx, lanes::Lane::Interactive, key)
.unwrap().request.affinity.as_deref(), Some("conv-2"));
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "task"
})).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(build_request(&req, tx, lanes::Lane::Interactive, None).affinity.is_none());
}
async fn sse_data_lines(resp: Response) -> Vec<String> {
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
String::from_utf8(bytes.to_vec()).unwrap()
.lines()
.filter_map(|l| l.strip_prefix("data: ").map(str::to_string))
.collect()
}
#[tokio::test]
async fn stream_chunks_carry_envelope_and_first_delta_role() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "he".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "llo".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 2, n_prompt: 10, n_cached: 0, elapsed_s: 0.1,
spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true), Vec::new(), None)
.into_response();
let lines = sse_data_lines(resp).await;
assert_eq!(lines.last().map(String::as_str), Some("[DONE]"));
let chunks: Vec<serde_json::Value> = lines[..lines.len() - 1].iter()
.map(|l| serde_json::from_str(l).unwrap()).collect();
let id = chunks[0]["id"].as_str().unwrap().to_string();
assert!(id.starts_with("chatcmpl-"));
for c in &chunks {
assert_eq!(c["id"], id.as_str());
assert!(c["created"].as_u64().unwrap() > 1_700_000_000);
assert!(c["system_fingerprint"].as_str().unwrap().starts_with("memra-"));
assert_eq!(c["object"], "chat.completion.chunk");
}
assert_eq!(chunks[0]["choices"][0]["delta"]["role"], "assistant");
assert_eq!(chunks[0]["choices"][0]["delta"]["content"], "he");
assert!(chunks[1]["choices"][0]["delta"].get("role").is_none());
let fin = chunks.last().unwrap();
assert_eq!(fin["choices"][0]["finish_reason"], "stop");
assert_eq!(fin["usage"]["prompt_tokens"], 10);
}
#[tokio::test]
async fn stream_excludes_stop_text_like_non_stream_does() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "answer\nPro".into() }).unwrap();
tx.send(Event::Token { id: 2, text: "blem: leaked prompt".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Callback".into(), n_tokens: 2, n_prompt: 8, n_cached: 0,
elapsed_s: 0.1, spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true),
vec!["Problem:".into()], None).into_response();
let lines = sse_data_lines(resp).await;
let content: String = lines.iter()
.filter(|l| *l != "[DONE]")
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.filter_map(|c| c["choices"][0]["delta"]["content"].as_str()
.map(str::to_string))
.collect();
assert_eq!(content, "answer\n");
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Token { id: 1, text: "ends in Pro".into() }).unwrap();
tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 8, n_cached: 0, elapsed_s: 0.1,
spec: None,
}).unwrap();
drop(tx);
let resp = sse_response(rx, "m".into(), true, None, Envelope::new(true),
vec!["Problem:".into()], None).into_response();
let lines = sse_data_lines(resp).await;
let content: String = lines.iter()
.filter(|l| *l != "[DONE]")
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.filter_map(|c| c["choices"][0]["delta"]["content"].as_str()
.map(str::to_string))
.collect();
assert_eq!(content, "ends in Pro");
}
#[tokio::test]
async fn stream_worker_error_is_a_data_chunk_not_a_named_event() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tx.send(Event::Error(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]"));
}
#[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(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "test"
})).unwrap()),
).await;
let chat = chat_completions(
State(st),
axum::http::HeaderMap::new(),
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 omitted_seed_is_fresh_entropy_not_a_pinned_zero() {
let comp_seed = |body: serde_json::Value| {
let req: CompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg.seed
};
let chat_seed = |body: serde_json::Value| {
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
.unwrap().request.sampler_cfg.seed
};
let a = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
let b = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
let c = chat_seed(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]}));
assert_ne!(a, 0, "omitted seed must not be the pinned 0 that caused the loop");
assert_ne!(b, 0);
assert_ne!(c, 0);
assert_ne!(a, b, "two seed-omitting requests must get DIFFERENT streams");
assert_ne!(a, c);
assert_eq!(comp_seed(serde_json::json!({
"model": "m", "prompt": "t", "seed": 0})), 0,
"explicit seed 0 must stay 0 — the determinism gates depend on it");
assert_eq!(comp_seed(serde_json::json!({
"model": "m", "prompt": "t", "seed": 12345})), 12345);
assert_eq!(chat_seed(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"seed": 777})), 777);
assert_eq!(comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42})),
comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42})));
let seeds: std::collections::HashSet<u64> = (0..256).map(|_| fresh_seed()).collect();
assert_eq!(seeds.len(), 256, "fresh_seed must not collide across rapid calls");
assert!(!seeds.contains(&0));
}
#[test]
fn response_format_builds_grammar_only_when_present() {
let mk = |rf: Option<serde_json::Value>| {
let mut body = serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]});
if let Some(rf) = rf { body["response_format"] = rf; }
let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
};
assert!(mk(None).unwrap().request.grammar.is_none());
assert!(mk(Some(serde_json::json!({"type": "text"})))
.unwrap().request.grammar.is_none());
assert!(matches!(mk(Some(serde_json::json!({"type": "json_object"})))
.unwrap().request.grammar,
Some(constrained::GrammarSpec::JsonObject)));
assert!(matches!(mk(Some(serde_json::json!({"type": "json_schema",
"json_schema": {"schema": {"type": "object"}}})))
.unwrap().request.grammar,
Some(constrained::GrammarSpec::JsonSchema(_))));
assert!(mk(Some(serde_json::json!({"type": "yaml"}))).is_err());
}
#[test]
fn unsupported_semantic_params_are_named_rejections() {
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"response_format": {"type": "json_object"}
})).unwrap();
assert!(req.response_format.is_some());
let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}],
"response_format": {"type": "text"}, "logprobs": false, "n": 1,
"user": "u-1", "stream_options": {"include_usage": true}
})).unwrap();
assert_eq!(req.response_format.as_ref().unwrap()["type"], "text");
assert_eq!(req.logprobs.as_ref().unwrap().as_bool(), Some(false));
assert_eq!(req.n, Some(1));
assert!(reject_unsupported(&[("logit_bias", false, "")]).is_ok());
let (msg, param) = reject_unsupported(&[("logit_bias", true, " (why)")]).unwrap_err();
assert_eq!(param, "logit_bias");
assert_eq!(msg, "logit_bias is not supported (why)");
}
#[test]
fn completions_accept_openai_stop_forms() {
for (value, expected) in [
(serde_json::json!("Problem:"), vec!["Problem:"]),
(serde_json::json!(["Question:", "Problem:"]), vec!["Question:", "Problem:"]),
(serde_json::Value::Null, Vec::<&str>::new()),
] {
let req: CompletionReq = serde_json::from_value(serde_json::json!({
"model": "plain_quant", "prompt": "task", "stop": value
})).unwrap();
assert_eq!(req.stop.into_vec(), expected);
}
}
fn fake_worker_state() -> AppState {
let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
let health = health::WorkerHealth::new();
let h = health.clone();
std::thread::spawn(move || {
h.mark_ready();
while let Ok(Cmd::Generate(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();
let _ = req.tx.send(Event::Token { id: 1, text: "ok".into() });
let _ = req.tx.send(Event::Done {
stop_reason: "Eos".into(), n_tokens: 1, n_prompt: 1, n_cached: 0,
elapsed_s: 0.01, spec: None,
});
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()),
metrics: SharedMetrics::default(),
started: 1,
inflight: Arc::new(Default::default()),
tenant_inflight: Arc::new(Default::default()),
health,
bg: None,
}
}
#[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(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
let h = resp.headers();
let limit: usize = h["x-ratelimit-limit"].to_str().unwrap().parse().unwrap();
let remaining: usize = h["x-ratelimit-remaining"].to_str().unwrap().parse().unwrap();
assert_eq!(remaining, limit - 1);
assert_eq!(h["x-ratelimit-reset"], "0");
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 0,
"slot must free at completion");
let resp = completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "prompt": "t", "stream": true
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().contains_key("x-ratelimit-limit"));
assert!(resp.headers().contains_key("x-ratelimit-remaining"));
assert!(resp.headers().contains_key("x-ratelimit-reset"));
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 1,
"stream in flight holds the slot");
let _ = axum::body::to_bytes(resp.into_body(), usize::MAX).await.unwrap();
assert_eq!(st.inflight[0].load(std::sync::atomic::Ordering::SeqCst), 0,
"slot must free when the stream completes");
}
#[tokio::test]
async fn draining_rejects_new_requests_with_503_and_retry_after() {
let _l = DRAIN_LOCK.lock().unwrap();
let st = fake_worker_state();
DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
let resp = chat_completions(State(st.clone()), axum::http::HeaderMap::new(),
Json(serde_json::from_value(serde_json::json!({
"model": "m", "messages": [{"role": "user", "content": "t"}]
})).unwrap())).await;
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
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(),
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(),
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 liveness_failure_obeys_the_retry_contract() {
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 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,
};
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 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());
}
#[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");
}
}