use std::sync::Arc;
use std::time::Instant;
use axum::body::Body;
use axum::extract::{Request, State};
use axum::http::HeaderMap;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Extension, Json, Router};
use bytes::Bytes;
use firstpass_core::features::{hour_bucket, token_bucket};
use firstpass_core::hashchain::sha256_hex;
use firstpass_core::{
Attempt, DeferredVerdict, Dialect, FEATURE_VERSION, Features, FinalOutcome, GENESIS_HASH, Mode,
PolicyRef, RequestInfo, Score, ServedFrom, TaskKind, Trace, Verdict,
};
use serde::Deserialize;
use serde_json::Value;
use std::future::Future;
use std::time::Duration;
use tokio::sync::mpsc::error::TrySendError;
use uuid::Uuid;
use crate::config::ProxyConfig;
use crate::error::ProxyError;
use crate::gate::{GateHealthRegistry, resolve_gates};
use crate::provider::{Auth, ChatMessage, ModelRequest, ModelResponse, ProviderRegistry};
use crate::router::{EnforceCtx, EngineOutcome, route_enforce};
use crate::store;
use crate::tenant_auth::{TenantId, auth_middleware};
use crate::upstream::{
forward_anthropic, forward_anthropic_streaming, forward_openai, forward_openai_streaming,
};
use firstpass_core::Route;
#[derive(Clone)]
pub struct AppState {
pub config: Arc<ProxyConfig>,
pub http: reqwest::Client,
pub providers: ProviderRegistry,
pub gate_health: Arc<GateHealthRegistry>,
pub traces: store::TraceSender,
pub adaptive: Option<Arc<std::sync::Mutex<firstpass_core::conformal::AdaptiveConformal>>>,
pub bandit: Option<Arc<std::sync::Mutex<crate::bandit::StartRungBandit>>>,
pub tenant_rate_limiter: Option<Arc<governor::DefaultKeyedRateLimiter<String>>>,
pub spill: Option<store::SpillHandle>,
}
#[must_use]
pub fn build_tenant_rate_limiter(
config: &ProxyConfig,
) -> Option<Arc<governor::DefaultKeyedRateLimiter<String>>> {
let per_sec = config.tenant_rate_per_sec?;
Some(Arc::new(governor::RateLimiter::keyed(
governor::Quota::per_second(per_sec),
)))
}
pub async fn tenant_rate_limit_middleware(
State(state): State<AppState>,
Extension(tenant): Extension<TenantId>,
req: Request,
next: Next,
) -> Response {
if let Some(limiter) = &state.tenant_rate_limiter
&& limiter.check_key(&tenant.0).is_err()
{
return ProxyError::RateLimited.into_response();
}
next.run(req).await
}
impl std::fmt::Debug for AppState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AppState")
.field("config", &self.config)
.finish_non_exhaustive()
}
}
fn offer_trace(traces: &store::TraceSender, spill: Option<&store::SpillHandle>, trace: Trace) {
record_trace_metrics(&trace);
match traces.try_send(trace) {
Ok(()) => {}
Err(TrySendError::Full(t)) => {
if let Some(handle) = spill {
match store::append_to_spill(handle, &t) {
Ok(()) => {
metrics::counter!("firstpass_receipts_spilled_total").increment(1);
}
Err(e) => {
tracing::error!(%e, "durable mode: spill write failed; trace lost");
metrics::counter!("firstpass_traces_dropped_total").increment(1);
}
}
} else {
tracing::warn!("trace channel full; dropping trace (writer behind under load)");
metrics::counter!("firstpass_traces_dropped_total").increment(1);
}
}
Err(TrySendError::Closed(_)) => {
tracing::warn!("trace writer is gone; dropping trace");
}
}
}
fn record_trace_metrics(trace: &Trace) {
if trace.mode == Mode::Enforce {
metrics::histogram!("firstpass_enforce_latency_ms")
.record(trace.final_.total_latency_ms as f64);
if trace.final_.escalations > 0 {
metrics::counter!("firstpass_escalations_total")
.increment(u64::from(trace.final_.escalations));
}
}
let served_from = match trace.final_.served_from {
ServedFrom::Attempt => "attempt",
ServedFrom::BestAttempt => "best_attempt",
ServedFrom::Error => "error",
};
metrics::counter!("firstpass_served_total", "served_from" => served_from).increment(1);
if trace.final_.served_from == ServedFrom::Error {
metrics::counter!("firstpass_upstream_failures_total").increment(1);
}
metrics::gauge!("firstpass_cost_usd_total").increment(trace.final_.total_cost_usd);
metrics::gauge!("firstpass_gate_cost_usd_total").increment(trace.final_.gate_cost_usd);
metrics::gauge!("firstpass_baseline_usd_total")
.increment(trace.final_.counterfactual_baseline_usd);
metrics::gauge!("firstpass_savings_usd_total").increment(trace.final_.savings_usd);
if let Some(rung) = trace.final_.served_rung {
let model = trace
.attempts
.iter()
.find(|a| a.rung == rung)
.map(|a| a.model.clone())
.unwrap_or_else(|| "unknown".to_owned());
metrics::counter!(
"firstpass_served_rung_total",
"rung" => rung.to_string(),
"model" => model
)
.increment(1);
}
}
const MAX_BODY_BYTES: usize = 16 * 1024 * 1024;
pub(crate) fn u01(seed: u128) -> f64 {
let lo = splitmix64_finalise(seed as u64);
let hi = splitmix64_finalise((seed >> 64) as u64);
((lo ^ hi) >> 11) as f64 * (1.0_f64 / (1u64 << 53) as f64)
}
fn splitmix64_finalise(mut z: u64) -> u64 {
z = z.wrapping_add(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
#[must_use]
pub(crate) fn epsilon_propensity(chosen: u32, greedy: u32, epsilon: f64, k: usize) -> f64 {
let greedy_term = if chosen == greedy { 1.0 - epsilon } else { 0.0 };
greedy_term + epsilon / k as f64
}
pub fn app(state: AppState) -> Result<Router, ProxyError> {
crate::metrics::install()?;
let max_concurrency = state.config.max_concurrency;
let business = Router::new()
.route("/v1/messages", post(messages))
.route("/v1/chat/completions", post(chat_completions))
.route("/v1/feedback", post(feedback))
.route("/v1/capabilities", get(capabilities))
.layer(axum::middleware::from_fn_with_state(
state.clone(),
tenant_rate_limit_middleware,
))
.layer(axum::middleware::from_fn_with_state(
state.clone(),
auth_middleware,
));
Ok(Router::new()
.merge(business)
.route("/healthz", get(healthz))
.route("/metrics", get(crate::metrics::handler))
.layer(axum::extract::DefaultBodyLimit::max(MAX_BODY_BYTES))
.layer(tower::limit::GlobalConcurrencyLimitLayer::new(
max_concurrency,
))
.with_state(state))
}
async fn healthz() -> impl IntoResponse {
Json(serde_json::json!({ "status": "ok" }))
}
async fn capabilities(State(state): State<AppState>) -> impl IntoResponse {
let (ladder, gates) = state
.config
.routing
.as_ref()
.and_then(|c| c.routes.iter().find(|r| r.mode == Mode::Enforce))
.map(|r| (r.ladder.clone(), r.gates.clone()))
.unwrap_or_default();
Json(serde_json::json!({
"service": "firstpass",
"version": env!("CARGO_PKG_VERSION"),
"feature_version": FEATURE_VERSION,
"modes": ["observe", "enforce"],
"wire_apis": ["anthropic.messages", "openai.chat_completions"],
"ladder": ladder,
"gates": gates,
"feedback_api": "POST /v1/feedback",
"offboarding": "unset ANTHROPIC_BASE_URL (or OPENAI_BASE_URL for OpenAI clients)",
}))
}
#[derive(Debug, Deserialize)]
struct FeedbackRequest {
trace_id: String,
gate_id: String,
verdict: String,
#[serde(default)]
score: Option<f64>,
reporter: String,
}
async fn feedback(
State(state): State<AppState>,
Extension(TenantId(tenant)): Extension<TenantId>,
body: Bytes,
) -> Response {
let req: FeedbackRequest = match serde_json::from_slice(&body) {
Ok(r) => r,
Err(e) => {
return ProxyError::BadRequest(format!("invalid feedback body: {e}")).into_response();
}
};
let verdict = match req.verdict.as_str() {
"pass" => Verdict::Pass,
"fail" => Verdict::Fail,
"abstain" => Verdict::Abstain,
other => {
return ProxyError::BadRequest(format!("unknown verdict {other:?}")).into_response();
}
};
let score = match req.score {
Some(s) => match Score::new(s) {
Ok(sc) => Some(sc),
Err(_) => {
return ProxyError::BadRequest(format!("score {s} out of range [0,1]"))
.into_response();
}
},
None => None,
};
let db = state.config.db_path.clone();
let (db_check, tenant_check, tid_check) = (db.clone(), tenant.clone(), req.trace_id.clone());
match tokio::task::spawn_blocking(move || {
store::trace_exists(&db_check, &tenant_check, &tid_check)
})
.await
{
Ok(Ok(true)) => {}
Ok(Ok(false)) => {
return ProxyError::NotFound(format!("unknown trace_id {:?}", req.trace_id))
.into_response();
}
Ok(Err(e)) => {
tracing::error!(%e, "feedback: trace_exists check failed");
return ProxyError::Internal(e.to_string()).into_response();
}
Err(e) => {
tracing::error!(%e, "feedback: trace_exists task panicked");
return ProxyError::Internal(e.to_string()).into_response();
}
}
let feedback_signal = match verdict {
Verdict::Pass => Some(true),
Verdict::Fail => Some(false),
Verdict::Abstain => None,
};
let dv = DeferredVerdict {
gate_id: req.gate_id,
verdict,
score,
reported_at: jiff::Timestamp::now(),
reporter: req.reporter,
};
let trace_id = req.trace_id.clone();
match tokio::task::spawn_blocking(move || store::append_deferred(&db, &req.trace_id, &dv)).await
{
Ok(Ok(())) => {
if let (Some(a), Some(correct)) = (state.adaptive.as_ref(), feedback_signal)
&& let Ok(mut g) = a.lock()
{
g.observe_served(correct);
metrics::gauge!("firstpass_serve_threshold").set(g.threshold());
metrics::gauge!("firstpass_realized_served_failure")
.set(g.realized_served_failure());
}
(
axum::http::StatusCode::ACCEPTED,
Json(serde_json::json!({ "status": "recorded", "trace_id": trace_id })),
)
.into_response()
}
Ok(Err(e)) => {
tracing::error!(%e, "feedback: append_deferred failed");
ProxyError::Internal(e.to_string()).into_response()
}
Err(e) => {
tracing::error!(%e, "feedback: append_deferred task panicked");
ProxyError::Internal(e.to_string()).into_response()
}
}
}
const SESSION_HEADER: &str = "x-firstpass-session";
const AGENT_HEADER: &str = "x-firstpass-agent";
const SUBAGENT_HEADER: &str = "x-firstpass-subagent";
async fn messages(
State(state): State<AppState>,
Extension(TenantId(tenant)): Extension<TenantId>,
headers: HeaderMap,
body: Bytes,
) -> Response {
let session_header = header_str(&headers, SESSION_HEADER);
if let Some(routing) = state.config.routing.as_ref() {
let features = extract_features(&headers, &body);
if let Some(route) = routing
.route_for(&features)
.filter(|r| r.mode == Mode::Enforce && !r.ladder.is_empty())
{
let route = route.clone();
if enforce_can_handle(
&features,
&body,
routing.escalation.enforce_structured,
&route.ladder,
&state.providers,
Dialect::Anthropic,
) {
return handle_enforce(
&state,
&headers,
&body,
features,
&route,
session_header,
tenant,
)
.await;
}
tracing::info!(
"enforce route matched but structured request can't be routed faithfully (flag/ladder); serving via observe passthrough"
);
}
}
observe_passthrough(state, headers, body, session_header, tenant).await
}
fn enforce_can_handle(
features: &Features,
body: &[u8],
enforce_structured: bool,
ladder: &[String],
providers: &crate::provider::ProviderRegistry,
inbound: Dialect,
) -> bool {
let structured = features.tool_count > 0
|| features.has_images
|| match inbound {
Dialect::Anthropic => messages_have_tool_blocks(body),
Dialect::Openai => openai_messages_have_tool_calls(body),
Dialect::Gemini => false,
};
if !structured {
return true;
}
if !enforce_structured {
return false;
}
let all_verbatim = ladder.iter().all(|rung| {
let provider_id = rung.split('/').next().unwrap_or_default();
providers
.get(provider_id)
.is_some_and(|p| p.carries_structured_verbatim(inbound))
});
if all_verbatim {
return true;
}
if inbound == Dialect::Openai && !openai_has_http_images(body) {
let all_anthropic = ladder.iter().all(|rung| {
let pid = rung.split('/').next().unwrap_or_default();
providers
.get(pid)
.is_some_and(|p| p.carries_structured_verbatim(Dialect::Anthropic))
});
if all_anthropic {
return true;
}
}
false
}
fn messages_have_tool_blocks(body: &[u8]) -> bool {
serde_json::from_slice::<Value>(body)
.ok()
.and_then(|json| {
json.get("messages")
.and_then(Value::as_array)
.map(|messages| messages.iter().any(message_has_tool_block))
})
.unwrap_or(false)
}
fn message_has_tool_block(message: &Value) -> bool {
message
.get("content")
.and_then(Value::as_array)
.is_some_and(|blocks| {
blocks.iter().any(|block| {
matches!(
block.get("type").and_then(Value::as_str),
Some("tool_use" | "tool_result")
)
})
})
}
fn is_stream_request(body: &[u8]) -> bool {
serde_json::from_slice::<Value>(body)
.ok()
.and_then(|json| json.get("stream").and_then(Value::as_bool))
.unwrap_or(false)
}
fn header_str(headers: &HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::to_owned)
}
fn extract_features(headers: &HeaderMap, body: &[u8]) -> Features {
let (_model, tool_count, has_images) = request_features(body);
let mut f = Features::new(TaskKind::Other);
f.agent = header_str(headers, AGENT_HEADER);
f.subagent = header_str(headers, SUBAGENT_HEADER);
f.tool_count = tool_count;
f.has_images = has_images;
f.prompt_token_bucket = token_bucket(body.len() as u64);
f.hour_bucket = hour_bucket(jiff::Timestamp::now());
f
}
async fn handle_enforce(
state: &AppState,
headers: &HeaderMap,
body: &Bytes,
features: Features,
route: &Route,
session_header: Option<String>,
tenant: String,
) -> Response {
if is_stream_request(body) {
let (tx, rx) = tokio::sync::oneshot::channel::<Result<Value, ProxyError>>();
let (state_c, headers_c, body_c, route_c) =
(state.clone(), headers.clone(), body.clone(), route.clone());
tokio::spawn(async move {
let out = enforce_pipeline(
&state_c,
&headers_c,
&body_c,
features,
&route_c,
session_header,
tenant,
)
.await;
let _ = tx.send(out);
});
return sse_keepalive_response(rx, anthropic_sse_from_message);
}
match enforce_pipeline(
state,
headers,
body,
features,
route,
session_header,
tenant,
)
.await
{
Ok(message) => (axum::http::StatusCode::OK, Json(message)).into_response(),
Err(e) => e.into_response(),
}
}
#[allow(clippy::too_many_arguments)] async fn enforce_pipeline_inner(
state: &AppState,
body: &Bytes,
base_request: ModelRequest,
auth: Auth,
features: Features,
route: &Route,
session_header: Option<String>,
tenant: String,
api: &str,
) -> Result<ModelResponse, ProxyError> {
let gate_defs = state
.config
.routing
.as_ref()
.map_or(&[][..], |cfg| &cfg.gate_defs);
let gates = resolve_gates(
&route.gates,
gate_defs,
&state.providers,
&auth,
&state.config.prices,
);
let session_id = session_header.unwrap_or_else(|| Uuid::now_v7().to_string());
let (budget, max_rungs, speculation, serve_threshold) = match state.config.routing.as_ref() {
Some(cfg) => (
cfg.budget.per_request_usd,
cfg.escalation.max_rungs_per_request,
cfg.escalation.speculation,
cfg.escalation.serve_threshold,
),
None => (None, 3, 0, None),
};
let serve_threshold = state
.adaptive
.as_ref()
.and_then(|a| a.lock().ok().map(|g| g.threshold()))
.or(serve_threshold);
let bandit_ctx = crate::bandit::ContextBucket::from_features(&features);
let (greedy_rung, base_policy_id, ts_propensity) = {
let (chosen, ts_p) = state
.bandit
.as_ref()
.and_then(|b| b.lock().ok())
.map(|mut b| {
b.choose_start_with_propensity(&bandit_ctx, &route.ladder, &state.config.prices)
})
.unwrap_or((0, None));
let policy = if ts_p.is_some() {
"bandit@v2-ts".to_owned()
} else if chosen > 0 {
"bandit@v1".to_owned()
} else {
"static-ladder@v0".to_owned()
};
(chosen, policy, ts_p)
};
let exploration_epsilon = state
.config
.routing
.as_ref()
.and_then(|cfg| cfg.escalation.exploration.as_ref())
.map(|e| e.epsilon);
let (start_rung, policy_id, explore_flag, propensity) = if ts_propensity.is_some() {
if exploration_epsilon.is_some() {
tracing::warn!(
"bandit.algorithm = thompson already logs propensities; \
[escalation.exploration] epsilon is ignored"
);
}
(greedy_rung, base_policy_id, false, ts_propensity)
} else if let Some(epsilon) = exploration_epsilon {
let k = route.ladder.len().max(1);
let u = u01(Uuid::now_v7().as_u128());
let (chosen, eps_branch) = if u < epsilon {
let idx = ((u / epsilon) * k as f64) as u32;
(idx.min(k as u32 - 1), true)
} else {
(greedy_rung, false)
};
let p = epsilon_propensity(chosen, greedy_rung, epsilon, k);
(chosen, format!("{base_policy_id}+eps"), eps_branch, Some(p))
} else {
(greedy_rung, base_policy_id, false, None)
};
let speculation = match state
.config
.routing
.as_ref()
.and_then(|cfg| cfg.escalation.speculation_band)
{
Some([lo, hi]) if speculation > 0 => {
let estimate = state
.bandit
.as_ref()
.and_then(|b| b.lock().ok())
.and_then(|b| b.pass_estimate(&bandit_ctx, start_rung));
match estimate {
Some(p) if p < lo || p > hi => {
metrics::counter!("firstpass_speculation_skipped_total").increment(1);
0
}
_ => speculation,
}
}
_ => speculation,
};
if state.bandit.is_some() {
metrics::counter!(
"firstpass_bandit_start_rung",
"rung" => start_rung.to_string()
)
.increment(1);
}
let ctx = EnforceCtx {
ladder: &route.ladder,
gates: &gates,
health: &state.gate_health,
base_request: &base_request,
providers: &state.providers,
auth: &auth,
prices: &state.config.prices,
budget_per_request_usd: budget,
max_rungs,
speculation,
serve_threshold,
features,
start_rung,
tenant_id: tenant,
session_id,
prompt_hash: prompt_hash(&state.config.prompt_salt, body),
api: api.to_owned(),
policy_id,
};
let (outcome, mut trace) = route_enforce(ctx).await;
trace.policy.explore = explore_flag;
trace.policy.propensity = propensity;
if let Some(bandit) = state.bandit.as_ref()
&& let Ok(mut b) = bandit.lock()
{
for attempt in &trace.attempts {
b.observe(&bandit_ctx, attempt.rung, attempt.verdict);
}
}
offer_trace(&state.traces, state.spill.as_ref(), trace);
match outcome {
EngineOutcome::Served(resp) => Ok(resp),
EngineOutcome::Failed(msg) => Err(ProxyError::Engine(msg)),
}
}
async fn enforce_pipeline(
state: &AppState,
headers: &HeaderMap,
body: &Bytes,
features: Features,
route: &Route,
session_header: Option<String>,
tenant: String,
) -> Result<Value, ProxyError> {
let Some(base_request) = parse_model_request(body) else {
return Err(ProxyError::BadRequest(
"request body is not a valid Anthropic Messages request".to_owned(),
));
};
let auth = Auth::from_headers(headers);
let resp = enforce_pipeline_inner(
state,
body,
base_request,
auth,
features,
route,
session_header,
tenant,
"anthropic.messages",
)
.await?;
Ok(anthropic_response_json(&resp))
}
async fn enforce_pipeline_openai(
state: &AppState,
headers: &HeaderMap,
body: &Bytes,
features: Features,
route: &Route,
session_header: Option<String>,
tenant: String,
) -> Result<Value, ProxyError> {
let providers = &state.providers;
let all_openai = route.ladder.iter().all(|rung| {
let pid = rung.split('/').next().unwrap_or_default();
providers
.get(pid)
.is_some_and(|p| p.carries_structured_verbatim(Dialect::Openai))
});
let Some(base_request) = parse_openai_request(body, all_openai) else {
return Err(ProxyError::BadRequest(
"request body is not a valid OpenAI Chat Completions request".to_owned(),
));
};
let auth = Auth::from_headers(headers);
let resp = enforce_pipeline_inner(
state,
body,
base_request,
auth,
features,
route,
session_header,
tenant,
"openai.chat_completions",
)
.await?;
Ok(openai_response_json(&resp))
}
const SSE_KEEPALIVE_EVERY: Duration = Duration::from_secs(5);
fn sse_keepalive_response(
rx: tokio::sync::oneshot::Receiver<Result<Value, ProxyError>>,
format_message: fn(&Value) -> String,
) -> Response {
let mut ticks = tokio::time::interval(SSE_KEEPALIVE_EVERY);
ticks.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ticks.reset(); let stream = KeepaliveStream {
rx: Some(rx),
ticks,
format_message,
};
(
axum::http::StatusCode::OK,
[(
axum::http::header::CONTENT_TYPE,
"text/event-stream; charset=utf-8",
)],
axum::body::Body::from_stream(stream),
)
.into_response()
}
struct KeepaliveStream {
rx: Option<tokio::sync::oneshot::Receiver<Result<Value, ProxyError>>>,
ticks: tokio::time::Interval,
format_message: fn(&Value) -> String,
}
impl futures_core::Stream for KeepaliveStream {
type Item = Result<Bytes, std::convert::Infallible>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
use std::task::Poll;
let Some(rx) = self.rx.as_mut() else {
return Poll::Ready(None); };
if let Poll::Ready(out) = std::pin::Pin::new(rx).poll(cx) {
let fmt = self.format_message;
let frame = match out {
Ok(Ok(message)) => fmt(&message),
Ok(Err(e)) => sse_error_event(&e),
Err(_) => sse_error_event(&ProxyError::Internal(
"enforce pipeline task dropped".to_owned(),
)),
};
self.rx = None;
return Poll::Ready(Some(Ok(Bytes::from(frame))));
}
if self.ticks.poll_tick(cx).is_ready() {
return Poll::Ready(Some(Ok(Bytes::from_static(b": firstpass routing\n\n"))));
}
Poll::Pending
}
}
fn sse_error_event(e: &ProxyError) -> String {
let mut out = String::new();
sse_event(
&mut out,
"error",
&serde_json::json!({
"type": "error",
"error": { "type": "api_error", "message": e.client_message() }
}),
);
out
}
fn parse_model_request(body: &[u8]) -> Option<ModelRequest> {
let json: Value = serde_json::from_slice(body).ok()?;
let raw = json.clone();
let messages_json = json.get("messages")?.as_array()?;
let messages = messages_json
.iter()
.map(|m| ChatMessage {
role: m
.get("role")
.and_then(Value::as_str)
.unwrap_or("user")
.to_owned(),
content: m
.get("content")
.cloned()
.unwrap_or_else(|| Value::String(String::new())),
})
.collect();
let system = json
.get("system")
.and_then(Value::as_str)
.map(str::to_owned);
let max_tokens = json
.get("max_tokens")
.and_then(Value::as_u64)
.and_then(|n| u32::try_from(n).ok())
.unwrap_or(1024);
let tools = json.get("tools").cloned().unwrap_or(Value::Null);
Some(ModelRequest {
model: json
.get("model")
.and_then(Value::as_str)
.unwrap_or_default()
.to_owned(),
system,
messages,
max_tokens,
tools,
raw,
})
}
fn anthropic_response_json(resp: &ModelResponse) -> Value {
let content = resp
.raw
.get("content")
.filter(|c| c.is_array())
.cloned()
.unwrap_or_else(|| serde_json::json!([{ "type": "text", "text": resp.text }]));
serde_json::json!({
"id": format!("msg_{}", Uuid::now_v7()),
"type": "message",
"role": "assistant",
"model": resp.model,
"content": content,
"usage": { "input_tokens": resp.in_tokens, "output_tokens": resp.out_tokens },
})
}
fn sse_event(out: &mut String, event: &str, data: &Value) {
out.push_str("event: ");
out.push_str(event);
out.push_str("\ndata: ");
out.push_str(&data.to_string());
out.push_str("\n\n");
}
fn anthropic_sse_from_message(message: &Value) -> String {
let mut out = String::new();
let mut start_msg = message.clone();
start_msg["content"] = Value::Array(Vec::new());
sse_event(
&mut out,
"message_start",
&serde_json::json!({ "type": "message_start", "message": start_msg }),
);
let empty = Vec::new();
let blocks = message
.get("content")
.and_then(Value::as_array)
.unwrap_or(&empty);
for (i, block) in blocks.iter().enumerate() {
match block.get("type").and_then(Value::as_str) {
Some("tool_use") => {
let mut shell = block.clone();
shell["input"] = serde_json::json!({});
sse_event(
&mut out,
"content_block_start",
&serde_json::json!({ "type": "content_block_start", "index": i, "content_block": shell }),
);
let input_json = block
.get("input")
.map_or_else(|| "{}".to_owned(), std::string::ToString::to_string);
sse_event(
&mut out,
"content_block_delta",
&serde_json::json!({ "type": "content_block_delta", "index": i,
"delta": { "type": "input_json_delta", "partial_json": input_json } }),
);
}
_ => {
let text = block.get("text").and_then(Value::as_str).unwrap_or("");
sse_event(
&mut out,
"content_block_start",
&serde_json::json!({ "type": "content_block_start", "index": i,
"content_block": { "type": "text", "text": "" } }),
);
sse_event(
&mut out,
"content_block_delta",
&serde_json::json!({ "type": "content_block_delta", "index": i,
"delta": { "type": "text_delta", "text": text } }),
);
}
}
sse_event(
&mut out,
"content_block_stop",
&serde_json::json!({ "type": "content_block_stop", "index": i }),
);
}
let out_tokens = message
.pointer("/usage/output_tokens")
.cloned()
.unwrap_or_else(|| Value::from(0));
sse_event(
&mut out,
"message_delta",
&serde_json::json!({ "type": "message_delta", "delta": { "stop_reason": "end_turn" },
"usage": { "output_tokens": out_tokens } }),
);
sse_event(
&mut out,
"message_stop",
&serde_json::json!({ "type": "message_stop" }),
);
out
}
async fn observe_passthrough(
state: AppState,
headers: HeaderMap,
body: Bytes,
session_header: Option<String>,
tenant: String,
) -> Response {
if is_stream_request(&body) {
return observe_stream(state, headers, body, session_header, tenant).await;
}
let start = Instant::now();
let result = forward_anthropic(
&state.http,
&state.config.upstream_anthropic,
&headers,
body.clone(),
)
.await;
let latency_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
match result {
Ok((status, resp_headers, resp_body)) => {
spawn_trace(
&state,
body,
Some(resp_body.clone()),
latency_ms,
session_header,
tenant,
);
(status, resp_headers, resp_body).into_response()
}
Err(err) => {
spawn_trace(&state, body, None, latency_ms, session_header, tenant);
err.into_response()
}
}
}
async fn observe_stream(
state: AppState,
headers: HeaderMap,
body: Bytes,
session_header: Option<String>,
tenant: String,
) -> Response {
let start = Instant::now();
let result = forward_anthropic_streaming(
&state.http,
&state.config.upstream_anthropic,
&headers,
body.clone(),
)
.await;
let latency_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
match result {
Ok((status, resp_headers, response)) => {
spawn_stream_trace(&state, body, latency_ms, session_header, tenant);
let stream_body = Body::from_stream(response.bytes_stream());
(status, resp_headers, stream_body).into_response()
}
Err(err) => {
spawn_trace(&state, body, None, latency_ms, session_header, tenant);
err.into_response()
}
}
}
fn spawn_stream_trace(
state: &AppState,
req_body: Bytes,
latency_ms: u64,
session_header: Option<String>,
tenant: String,
) {
let config = state.config.clone();
let traces = state.traces.clone();
let spill = state.spill.clone();
tokio::spawn(async move {
let mut trace =
build_stream_trace(&config, &req_body, latency_ms, session_header.as_deref());
trace.tenant_id = tenant;
offer_trace(&traces, spill.as_ref(), trace);
});
}
fn spawn_trace(
state: &AppState,
req_body: Bytes,
resp_body: Option<Bytes>,
latency_ms: u64,
session_header: Option<String>,
tenant: String,
) {
let config = state.config.clone();
let traces = state.traces.clone();
let spill = state.spill.clone();
tokio::spawn(async move {
let mut trace = match resp_body {
Some(resp) => build_trace(
&config,
&req_body,
&resp,
latency_ms,
session_header.as_deref(),
),
None => build_error_trace(&config, &req_body, latency_ms, session_header.as_deref()),
};
trace.tenant_id = tenant;
offer_trace(&traces, spill.as_ref(), trace);
});
}
fn session_id(session_header: Option<&str>, trace_id: Uuid) -> String {
session_header
.map(str::to_owned)
.unwrap_or_else(|| trace_id.to_string())
}
fn prompt_hash(salt: &str, body: &[u8]) -> String {
let mut salted = Vec::with_capacity(salt.len() + body.len());
salted.extend_from_slice(salt.as_bytes());
salted.extend_from_slice(body);
sha256_hex(&salted)
}
fn request_features(body: &[u8]) -> (Option<String>, u32, bool) {
let Ok(json) = serde_json::from_slice::<Value>(body) else {
return (None, 0, false);
};
let model = json.get("model").and_then(Value::as_str).map(str::to_owned);
let tool_count = json
.get("tools")
.and_then(Value::as_array)
.map_or(0, |tools| u32::try_from(tools.len()).unwrap_or(u32::MAX));
let has_images = json
.get("messages")
.and_then(Value::as_array)
.is_some_and(|messages| messages.iter().any(message_has_image));
(model, tool_count, has_images)
}
fn message_has_image(message: &Value) -> bool {
message
.get("content")
.and_then(Value::as_array)
.is_some_and(|blocks| {
blocks
.iter()
.any(|block| block.get("type").and_then(Value::as_str) == Some("image"))
})
}
fn response_usage(body: &[u8]) -> (Option<String>, u64, u64) {
let Ok(json) = serde_json::from_slice::<Value>(body) else {
return (None, 0, 0);
};
let model = json.get("model").and_then(Value::as_str).map(str::to_owned);
let in_tokens = json
.pointer("/usage/input_tokens")
.and_then(Value::as_u64)
.unwrap_or(0);
let out_tokens = json
.pointer("/usage/output_tokens")
.and_then(Value::as_u64)
.unwrap_or(0);
(model, in_tokens, out_tokens)
}
fn build_trace(
config: &ProxyConfig,
req_body: &Bytes,
resp_body: &Bytes,
latency_ms: u64,
session_header: Option<&str>,
) -> Trace {
let (req_model, tool_count, has_images) = request_features(req_body);
let (resp_model, in_tokens, out_tokens) = response_usage(resp_body);
let model = resp_model
.or(req_model)
.unwrap_or_else(|| "unknown".to_owned());
let cost_usd = config
.prices
.cost_usd(&format!("anthropic/{model}"), in_tokens, out_tokens)
.unwrap_or(0.0);
let attempt = Attempt {
rung: 0,
model,
provider: "anthropic".to_owned(),
in_tokens,
out_tokens,
cost_usd,
latency_ms,
gates: Vec::new(),
verdict: Verdict::Pass,
};
let mut trace = base_trace(config, req_body, latency_ms, session_header);
trace.request.features.prompt_token_bucket = token_bucket(in_tokens);
trace.request.features.tool_count = tool_count;
trace.request.features.has_images = has_images;
trace.attempts.push(attempt);
trace.final_ = FinalOutcome {
served_rung: Some(0),
served_from: ServedFrom::Attempt,
total_cost_usd: cost_usd,
gate_cost_usd: 0.0,
total_latency_ms: latency_ms,
escalations: 0,
counterfactual_baseline_usd: cost_usd,
savings_usd: 0.0,
};
trace.recompute_savings();
trace
}
fn build_stream_trace(
config: &ProxyConfig,
req_body: &Bytes,
latency_ms: u64,
session_header: Option<&str>,
) -> Trace {
let (req_model, tool_count, has_images) = request_features(req_body);
let model = req_model.unwrap_or_else(|| "unknown".to_owned());
let attempt = Attempt {
rung: 0,
model,
provider: "anthropic".to_owned(),
in_tokens: 0,
out_tokens: 0,
cost_usd: 0.0,
latency_ms,
gates: Vec::new(),
verdict: Verdict::Pass,
};
let mut trace = base_trace(config, req_body, latency_ms, session_header);
trace.request.features.tool_count = tool_count;
trace.request.features.has_images = has_images;
trace.attempts.push(attempt);
trace.final_ = FinalOutcome {
served_rung: Some(0),
served_from: ServedFrom::Attempt,
total_cost_usd: 0.0,
gate_cost_usd: 0.0,
total_latency_ms: latency_ms,
escalations: 0,
counterfactual_baseline_usd: 0.0,
savings_usd: 0.0,
};
trace.recompute_savings();
trace
}
fn build_error_trace(
config: &ProxyConfig,
req_body: &Bytes,
latency_ms: u64,
session_header: Option<&str>,
) -> Trace {
let (_, tool_count, has_images) = request_features(req_body);
let mut trace = base_trace(config, req_body, latency_ms, session_header);
trace.request.features.tool_count = tool_count;
trace.request.features.has_images = has_images;
trace.final_ = FinalOutcome {
served_rung: None,
served_from: ServedFrom::Error,
total_cost_usd: 0.0,
gate_cost_usd: 0.0,
total_latency_ms: latency_ms,
escalations: 0,
counterfactual_baseline_usd: 0.0,
savings_usd: 0.0,
};
trace.recompute_savings();
trace
}
fn base_trace(
config: &ProxyConfig,
req_body: &Bytes,
latency_ms: u64,
session_header: Option<&str>,
) -> Trace {
let trace_id = Uuid::now_v7();
let mut features = Features::new(TaskKind::Other);
features.hour_bucket = hour_bucket(jiff::Timestamp::now());
Trace {
trace_id,
prev_hash: GENESIS_HASH.to_owned(),
tenant_id: config.tenant_id.clone(),
session_id: session_id(session_header, trace_id),
ts: jiff::Timestamp::now(),
mode: Mode::Observe,
policy: PolicyRef {
id: "observe-passthrough@v0".to_owned(),
explore: false,
propensity: None,
},
request: RequestInfo {
api: "anthropic.messages".to_owned(),
prompt_hash: prompt_hash(&config.prompt_salt, req_body),
features,
},
attempts: Vec::new(),
deferred: Vec::new(),
final_: FinalOutcome {
served_rung: None,
served_from: ServedFrom::Error,
total_cost_usd: 0.0,
gate_cost_usd: 0.0,
total_latency_ms: latency_ms,
escalations: 0,
counterfactual_baseline_usd: 0.0,
savings_usd: 0.0,
},
}
}
fn openai_messages_have_tool_calls(body: &[u8]) -> bool {
serde_json::from_slice::<Value>(body)
.ok()
.and_then(|json| {
json.get("messages").and_then(Value::as_array).map(|msgs| {
msgs.iter().any(|m| {
m.get("tool_calls").is_some()
|| m.get("role").and_then(Value::as_str) == Some("tool")
})
})
})
.unwrap_or(false)
}
fn openai_has_http_images(body: &[u8]) -> bool {
serde_json::from_slice::<Value>(body)
.ok()
.and_then(|json| {
json.get("messages").and_then(Value::as_array).map(|msgs| {
msgs.iter().any(|m| {
m.get("content")
.and_then(Value::as_array)
.is_some_and(|parts| {
parts.iter().any(|p| {
p.get("type").and_then(Value::as_str) == Some("image_url")
&& p.pointer("/image_url/url")
.and_then(Value::as_str)
.is_some_and(|u| {
u.starts_with("http://") || u.starts_with("https://")
})
})
})
})
})
})
.unwrap_or(false)
}
fn openai_messages_have_images(body: &[u8]) -> bool {
serde_json::from_slice::<Value>(body)
.ok()
.and_then(|json| {
json.get("messages").and_then(Value::as_array).map(|msgs| {
msgs.iter().any(|m| {
m.get("content")
.and_then(Value::as_array)
.is_some_and(|parts| {
parts
.iter()
.any(|p| p.get("type").and_then(Value::as_str) == Some("image_url"))
})
})
})
})
.unwrap_or(false)
}
fn extract_openai_features(headers: &HeaderMap, body: &[u8]) -> Features {
let Ok(json) = serde_json::from_slice::<Value>(body) else {
let mut f = Features::new(TaskKind::Other);
f.hour_bucket = hour_bucket(jiff::Timestamp::now());
return f;
};
let tool_count = json
.get("tools")
.and_then(Value::as_array)
.map_or(0, |tools| u32::try_from(tools.len()).unwrap_or(u32::MAX));
let has_images = openai_messages_have_images(body);
let mut f = Features::new(TaskKind::Other);
f.agent = header_str(headers, AGENT_HEADER);
f.subagent = header_str(headers, SUBAGENT_HEADER);
f.tool_count = tool_count;
f.has_images = has_images;
f.prompt_token_bucket = token_bucket(body.len() as u64);
f.hour_bucket = hour_bucket(jiff::Timestamp::now());
f
}
fn parse_data_url(url: &str) -> Option<(&str, &str)> {
let rest = url.strip_prefix("data:")?;
let (meta, data) = rest.split_once(',')?;
let media_type = meta.strip_suffix(";base64")?;
Some((media_type, data))
}
fn translate_openai_user_content(content: &Value) -> Option<Value> {
match content {
Value::String(_) => Some(content.clone()),
Value::Array(parts) => {
let mut blocks: Vec<Value> = Vec::with_capacity(parts.len());
for part in parts {
match part.get("type").and_then(Value::as_str) {
Some("text") => {
let text = part.get("text").and_then(Value::as_str).unwrap_or("");
blocks.push(serde_json::json!({ "type": "text", "text": text }));
}
Some("image_url") => {
let url = part.pointer("/image_url/url").and_then(Value::as_str)?;
if url.starts_with("http://") || url.starts_with("https://") {
return None; }
let (media_type, data) = parse_data_url(url)?;
blocks.push(serde_json::json!({
"type": "image",
"source": { "type": "base64", "media_type": media_type, "data": data }
}));
}
_ => {} }
}
Some(Value::Array(blocks))
}
_ => Some(Value::String(String::new())),
}
}
fn translate_openai_tools(tools: &Value) -> Value {
let Some(arr) = tools.as_array() else {
return Value::Null;
};
let converted: Vec<Value> = arr
.iter()
.map(|tool| {
let func = tool.get("function").unwrap_or(&Value::Null);
let mut out = serde_json::json!({
"name": func.get("name").cloned().unwrap_or(Value::String(String::new())),
"input_schema": func.get("parameters").cloned()
.unwrap_or_else(|| serde_json::json!({ "type": "object" })),
});
if let Some(desc) = func.get("description") {
out["description"] = desc.clone();
}
out
})
.collect();
Value::Array(converted)
}
fn translate_openai_tool_choice(tc: &Value) -> Value {
match tc {
Value::String(s) => match s.as_str() {
"auto" => serde_json::json!({ "type": "auto" }),
"required" => serde_json::json!({ "type": "any" }),
_ => serde_json::json!({ "type": "auto" }),
},
Value::Object(_) => {
if tc.get("type").and_then(Value::as_str) == Some("function") {
let name = tc.pointer("/function/name").cloned().unwrap_or(Value::Null);
serde_json::json!({ "type": "tool", "name": name })
} else {
serde_json::json!({ "type": "auto" })
}
}
_ => serde_json::json!({ "type": "auto" }),
}
}
pub fn parse_openai_request(body: &[u8], carry_raw: bool) -> Option<ModelRequest> {
let json: Value = serde_json::from_slice(body).ok()?;
let raw = if carry_raw { json.clone() } else { Value::Null };
let messages_json = json.get("messages")?.as_array()?;
let mut system: Option<String> = None;
let mut messages: Vec<ChatMessage> = Vec::with_capacity(messages_json.len());
let mut tools = Value::Null;
let mut tool_choice_override: Option<Value> = None;
for msg in messages_json {
let role = msg.get("role").and_then(Value::as_str).unwrap_or("user");
match role {
"system" => {
if let Some(s) = msg.get("content").and_then(Value::as_str) {
system = Some(s.to_owned());
}
}
"user" => {
let content_val = msg.get("content").unwrap_or(&Value::Null);
let translated = translate_openai_user_content(content_val)?;
messages.push(ChatMessage {
role: "user".to_owned(),
content: translated,
});
}
"assistant" => {
if let Some(tc_arr) = msg.get("tool_calls").and_then(Value::as_array) {
let mut blocks: Vec<Value> = Vec::new();
if let Some(text) = msg.get("content").and_then(Value::as_str)
&& !text.is_empty()
{
blocks.push(serde_json::json!({ "type": "text", "text": text }));
}
for tc in tc_arr {
let id = tc.get("id").and_then(Value::as_str).unwrap_or("");
let func = tc.get("function").unwrap_or(&Value::Null);
let name = func.get("name").and_then(Value::as_str).unwrap_or("");
let args_str = func
.get("arguments")
.and_then(Value::as_str)
.unwrap_or("{}");
let input: Value = serde_json::from_str(args_str)
.unwrap_or_else(|_| serde_json::json!({}));
blocks.push(serde_json::json!({
"type": "tool_use",
"id": id,
"name": name,
"input": input,
}));
}
messages.push(ChatMessage {
role: "assistant".to_owned(),
content: Value::Array(blocks),
});
} else {
let content = msg
.get("content")
.cloned()
.unwrap_or_else(|| Value::String(String::new()));
messages.push(ChatMessage {
role: "assistant".to_owned(),
content,
});
}
}
"tool" => {
let tool_call_id = msg
.get("tool_call_id")
.and_then(Value::as_str)
.unwrap_or("");
let content = msg
.get("content")
.cloned()
.unwrap_or_else(|| Value::String(String::new()));
let result_block = serde_json::json!({
"type": "tool_result",
"tool_use_id": tool_call_id,
"content": content,
});
messages.push(ChatMessage {
role: "user".to_owned(),
content: Value::Array(vec![result_block]),
});
}
_ => {} }
}
if !carry_raw {
if let Some(t) = json.get("tools") {
tools = translate_openai_tools(t);
}
if let Some(tc) = json.get("tool_choice") {
tool_choice_override = Some(translate_openai_tool_choice(tc));
}
} else {
tools = json.get("tools").cloned().unwrap_or(Value::Null);
}
let max_tokens = json
.get("max_tokens")
.and_then(Value::as_u64)
.and_then(|n| u32::try_from(n).ok())
.unwrap_or(1024);
let model = json
.get("model")
.and_then(Value::as_str)
.unwrap_or_default()
.to_owned();
let _ = tool_choice_override;
Some(ModelRequest {
model,
system,
messages,
max_tokens,
tools,
raw,
})
}
fn extract_openai_content_and_tools(raw: &Value, text: &str) -> (Value, Option<Value>) {
if let Some(blocks) = raw.get("content").and_then(Value::as_array) {
let mut text_parts: Vec<&str> = Vec::new();
let mut tool_calls: Vec<Value> = Vec::new();
for block in blocks {
match block.get("type").and_then(Value::as_str) {
Some("text") => {
if let Some(t) = block.get("text").and_then(Value::as_str) {
text_parts.push(t);
}
}
Some("tool_use") => {
let id = block.get("id").and_then(Value::as_str).unwrap_or("");
let name = block.get("name").and_then(Value::as_str).unwrap_or("");
let input_str = block
.get("input")
.map_or_else(|| "{}".to_owned(), std::string::ToString::to_string);
tool_calls.push(serde_json::json!({
"id": id,
"type": "function",
"function": { "name": name, "arguments": input_str },
}));
}
_ => {}
}
}
let content_text = if tool_calls.is_empty() || !text_parts.is_empty() {
Value::String(text_parts.join(""))
} else {
Value::Null };
let tc = if tool_calls.is_empty() {
None
} else {
Some(Value::Array(tool_calls))
};
return (content_text, tc);
}
if let Some(msg) = raw.pointer("/choices/0/message") {
let content = msg
.get("content")
.cloned()
.unwrap_or(Value::String(text.to_owned()));
let tc = msg.get("tool_calls").cloned();
return (content, tc);
}
(Value::String(text.to_owned()), None)
}
fn openai_response_json(resp: &ModelResponse) -> Value {
let (content_text, tool_calls) = extract_openai_content_and_tools(&resp.raw, &resp.text);
let finish_reason = if tool_calls.is_some() {
"tool_calls"
} else {
"stop"
};
let mut message = serde_json::json!({
"role": "assistant",
"content": content_text,
});
if let Some(tc) = tool_calls {
message["tool_calls"] = tc;
}
serde_json::json!({
"id": format!("chatcmpl-{}", Uuid::now_v7()),
"object": "chat.completion",
"created": jiff::Timestamp::now().as_second(),
"model": resp.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": finish_reason,
}],
"usage": {
"prompt_tokens": resp.in_tokens,
"completion_tokens": resp.out_tokens,
"total_tokens": resp.in_tokens + resp.out_tokens,
}
})
}
fn openai_sse_from_message(message: &Value) -> String {
let id = message
.get("id")
.and_then(Value::as_str)
.unwrap_or("chatcmpl-unknown")
.to_owned();
let created = message.get("created").cloned().unwrap_or(Value::from(0));
let model = message
.get("model")
.and_then(Value::as_str)
.unwrap_or("unknown");
let choices = message.get("choices").and_then(Value::as_array);
let msg = choices
.and_then(|c| c.first())
.and_then(|c| c.get("message"));
let content = msg
.and_then(|m| m.get("content"))
.cloned()
.unwrap_or(Value::Null);
let tool_calls = msg.and_then(|m| m.get("tool_calls")).cloned();
let finish_reason = choices
.and_then(|c| c.first())
.and_then(|c| c.get("finish_reason"))
.cloned()
.unwrap_or_else(|| Value::String("stop".to_owned()));
let mut out = String::new();
let chunk = |delta: Value| {
serde_json::json!({
"id": id,
"object": "chat.completion.chunk",
"created": created,
"model": model,
"choices": [{ "index": 0, "delta": delta, "finish_reason": Value::Null }]
})
};
let role_chunk = chunk(serde_json::json!({ "role": "assistant", "content": "" }));
out.push_str("data: ");
out.push_str(&role_chunk.to_string());
out.push_str("\n\n");
if let Value::String(text) = &content
&& !text.is_empty()
{
let content_chunk = chunk(serde_json::json!({ "content": text }));
out.push_str("data: ");
out.push_str(&content_chunk.to_string());
out.push_str("\n\n");
}
if let Some(tc) = tool_calls {
let tc_chunk = chunk(serde_json::json!({ "tool_calls": tc }));
out.push_str("data: ");
out.push_str(&tc_chunk.to_string());
out.push_str("\n\n");
}
let finish_chunk = serde_json::json!({
"id": id,
"object": "chat.completion.chunk",
"created": created,
"model": model,
"choices": [{ "index": 0, "delta": {}, "finish_reason": finish_reason }]
});
out.push_str("data: ");
out.push_str(&finish_chunk.to_string());
out.push_str("\n\n");
out.push_str("data: [DONE]\n\n");
out
}
async fn handle_enforce_openai(
state: &AppState,
headers: &HeaderMap,
body: &Bytes,
features: Features,
route: &Route,
session_header: Option<String>,
tenant: String,
) -> Response {
if is_stream_request(body) {
let (tx, rx) = tokio::sync::oneshot::channel::<Result<Value, ProxyError>>();
let (state_c, headers_c, body_c, route_c) =
(state.clone(), headers.clone(), body.clone(), route.clone());
tokio::spawn(async move {
let out = enforce_pipeline_openai(
&state_c,
&headers_c,
&body_c,
features,
&route_c,
session_header,
tenant,
)
.await;
let _ = tx.send(out);
});
return sse_keepalive_response(rx, openai_sse_from_message);
}
match enforce_pipeline_openai(
state,
headers,
body,
features,
route,
session_header,
tenant,
)
.await
{
Ok(message) => (axum::http::StatusCode::OK, Json(message)).into_response(),
Err(e) => e.into_response(),
}
}
async fn observe_passthrough_openai(
state: AppState,
headers: HeaderMap,
body: Bytes,
session_header: Option<String>,
tenant: String,
) -> Response {
if is_stream_request(&body) {
return observe_stream_openai(state, headers, body, session_header, tenant).await;
}
let start = Instant::now();
let result = forward_openai(
&state.http,
&state.config.upstream_openai,
&headers,
body.clone(),
)
.await;
let latency_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
match result {
Ok((status, resp_headers, resp_body)) => {
spawn_trace(
&state,
body,
Some(resp_body.clone()),
latency_ms,
session_header,
tenant,
);
(status, resp_headers, resp_body).into_response()
}
Err(err) => {
spawn_trace(&state, body, None, latency_ms, session_header, tenant);
err.into_response()
}
}
}
async fn observe_stream_openai(
state: AppState,
headers: HeaderMap,
body: Bytes,
session_header: Option<String>,
tenant: String,
) -> Response {
let start = Instant::now();
let result = forward_openai_streaming(
&state.http,
&state.config.upstream_openai,
&headers,
body.clone(),
)
.await;
let latency_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
match result {
Ok((status, resp_headers, response)) => {
spawn_stream_trace(&state, body, latency_ms, session_header, tenant);
let stream_body = Body::from_stream(response.bytes_stream());
(status, resp_headers, stream_body).into_response()
}
Err(err) => {
spawn_trace(&state, body, None, latency_ms, session_header, tenant);
err.into_response()
}
}
}
async fn chat_completions(
State(state): State<AppState>,
Extension(TenantId(tenant)): Extension<TenantId>,
headers: HeaderMap,
body: Bytes,
) -> Response {
let session_header = header_str(&headers, SESSION_HEADER);
if let Some(routing) = state.config.routing.as_ref() {
let features = extract_openai_features(&headers, &body);
if let Some(route) = routing
.route_for(&features)
.filter(|r| r.mode == Mode::Enforce && !r.ladder.is_empty())
{
let route = route.clone();
if enforce_can_handle(
&features,
&body,
routing.escalation.enforce_structured,
&route.ladder,
&state.providers,
Dialect::Openai,
) {
return handle_enforce_openai(
&state,
&headers,
&body,
features,
&route,
session_header,
tenant,
)
.await;
}
tracing::info!(
"enforce route matched but OpenAI structured request can't be routed faithfully (flag/ladder); serving via observe passthrough"
);
}
}
observe_passthrough_openai(state, headers, body, session_header, tenant).await
}
#[cfg(test)]
mod tests {
use bytes::Bytes;
use super::*;
fn test_config() -> ProxyConfig {
ProxyConfig::from_lookup(|_| None).unwrap()
}
#[test]
fn build_trace_maps_request_and_response_fields() {
let config = test_config();
let req = Bytes::from_static(
br#"{"model":"claude-haiku-4-5","tools":[{"name":"a"}],"messages":[]}"#,
);
let resp = Bytes::from_static(
br#"{"model":"claude-haiku-4-5","usage":{"input_tokens":1200,"output_tokens":300}}"#,
);
let trace = build_trace(&config, &req, &resp, 42, Some("sess-1"));
assert_eq!(trace.request.api, "anthropic.messages");
assert_eq!(trace.session_id, "sess-1");
assert_eq!(trace.attempts.len(), 1);
let attempt = &trace.attempts[0];
assert_eq!(attempt.model, "claude-haiku-4-5");
assert_eq!(attempt.provider, "anthropic");
assert_eq!(attempt.in_tokens, 1200);
assert_eq!(attempt.out_tokens, 300);
assert!(attempt.cost_usd > 0.0);
assert_eq!(trace.request.features.tool_count, 1);
assert!(!trace.request.features.has_images);
assert_eq!(trace.final_.served_rung, Some(0));
}
#[test]
fn build_trace_falls_back_to_trace_id_session_when_header_absent() {
let config = test_config();
let req = Bytes::from_static(b"{}");
let resp = Bytes::from_static(b"{}");
let trace = build_trace(&config, &req, &resp, 1, None);
assert_eq!(trace.session_id, trace.trace_id.to_string());
}
#[test]
fn build_error_trace_has_no_attempts_and_served_from_error() {
let config = test_config();
let req = Bytes::from_static(br#"{"model":"claude-haiku-4-5"}"#);
let trace = build_error_trace(&config, &req, 7, None);
assert!(trace.attempts.is_empty());
assert_eq!(trace.final_.served_from, ServedFrom::Error);
assert_eq!(trace.final_.served_rung, None);
}
#[test]
fn message_with_image_block_sets_has_images() {
let req = br#"{"messages":[{"role":"user","content":[{"type":"image"}]}]}"#;
let (_, _, has_images) = request_features(req);
assert!(has_images);
}
#[test]
fn prompt_hash_never_contains_raw_prompt_text() {
let hash = prompt_hash("salt", b"super secret prompt");
assert!(!hash.contains("secret"));
assert_eq!(hash.len(), 64);
}
#[test]
fn parse_model_request_preserves_content_verbatim_and_projects_text() {
let body = br#"{"model":"m","system":"sys","max_tokens":50,
"messages":[{"role":"user","content":[{"type":"text","text":"a"},{"type":"text","text":"b"}]},
{"role":"assistant","content":"c"}]}"#;
let req = parse_model_request(body).unwrap();
assert_eq!(req.system.as_deref(), Some("sys"));
assert_eq!(req.max_tokens, 50);
assert_eq!(req.messages.len(), 2);
assert_eq!(
req.messages[0].content,
serde_json::json!([{"type":"text","text":"a"},{"type":"text","text":"b"}])
);
assert_eq!(req.messages[1].content, Value::String("c".to_owned()));
assert_eq!(req.messages[0].text_view(), "a\nb");
assert_eq!(req.messages[1].text_view(), "c");
}
#[test]
fn tool_and_image_blocks_survive_the_request_round_trip() {
let body = br#"{"model":"m","max_tokens":50,"messages":[
{"role":"assistant","content":[{"type":"tool_use","id":"t1","name":"calc","input":{"x":1}}]},
{"role":"user","content":[
{"type":"tool_result","tool_use_id":"t1","content":"2"},
{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}
]}]}"#;
let req = parse_model_request(body).unwrap();
let round_tripped = serde_json::to_value(&req.messages).unwrap();
assert_eq!(
round_tripped,
serde_json::json!([
{"role":"assistant","content":[{"type":"tool_use","id":"t1","name":"calc","input":{"x":1}}]},
{"role":"user","content":[
{"type":"tool_result","tool_use_id":"t1","content":"2"},
{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}
]}
])
);
}
#[test]
fn text_message_serializes_byte_identical_to_a_plain_string() {
let m = ChatMessage::text("user", "hello");
assert_eq!(
serde_json::to_string(&m).unwrap(),
r#"{"role":"user","content":"hello"}"#
);
}
#[test]
fn parse_model_request_rejects_non_message_bodies() {
assert!(parse_model_request(b"not json").is_none());
assert!(parse_model_request(br#"{"no":"messages"}"#).is_none());
}
use crate::provider::{MockProvider, ModelResponse, Provider, ProviderError, ProviderRegistry};
use axum::extract::State;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::mpsc;
fn model_resp(model: &str, text: &str) -> ModelResponse {
ModelResponse {
model: model.to_owned(),
text: text.to_owned(),
in_tokens: 1000,
out_tokens: 400,
raw: serde_json::Value::Null,
}
}
fn enforce_state(
ladder: &[&str],
gates: &[&str],
outcomes: Vec<(&str, Result<ModelResponse, ProviderError>)>,
) -> (AppState, mpsc::Receiver<Trace>) {
let toml = format!(
"[[route]]\nmatch = {{}}\nmode = \"enforce\"\nladder = [{}]\ngates = [{}]\n",
ladder
.iter()
.map(|m| format!("\"{m}\""))
.collect::<Vec<_>>()
.join(", "),
gates
.iter()
.map(|g| format!("\"{g}\""))
.collect::<Vec<_>>()
.join(", "),
);
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_CONFIG_TOML" => Some(toml.clone()),
"FIRSTPASS_MODE" => Some("enforce".to_owned()),
_ => None,
})
.unwrap();
let mut outs = HashMap::new();
for (model, out) in outcomes {
outs.insert(model.to_owned(), out);
}
let mut map: HashMap<String, Arc<dyn Provider>> = HashMap::new();
map.insert(
"anthropic".to_owned(),
Arc::new(MockProvider::new("anthropic", outs)),
);
let providers = ProviderRegistry::from_map(map);
let (traces, rx) = mpsc::channel(64);
let state = AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers,
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter: None,
spill: None,
};
(state, rx)
}
fn user_body() -> Bytes {
Bytes::from_static(
br#"{"model":"ignored","max_tokens":64,"messages":[{"role":"user","content":"hi"}]}"#,
)
}
async fn body_json(resp: Response) -> Value {
let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
.await
.unwrap();
serde_json::from_slice(&bytes).unwrap()
}
#[tokio::test]
async fn enforce_serves_first_pass_and_returns_anthropic_shape() {
let (state, mut rx) = enforce_state(
&["anthropic/claude-haiku-4-5", "anthropic/claude-sonnet-5"],
&["non-empty"],
vec![(
"anthropic/claude-haiku-4-5",
Ok(model_resp("anthropic/claude-haiku-4-5", "hello")),
)],
);
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
user_body(),
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::OK);
let json = body_json(resp).await;
assert_eq!(json["type"], "message");
assert_eq!(json["content"][0]["text"], "hello");
assert_eq!(json["model"], "anthropic/claude-haiku-4-5");
let trace = rx.try_recv().expect("a trace was enqueued");
assert_eq!(trace.mode, Mode::Enforce);
assert_eq!(trace.final_.served_rung, Some(0));
assert_eq!(trace.attempts.len(), 1);
}
#[tokio::test]
async fn enforce_escalates_then_serves_and_traces_two_attempts() {
let (state, mut rx) = enforce_state(
&["anthropic/claude-haiku-4-5", "anthropic/claude-sonnet-5"],
&["non-empty"],
vec![
(
"anthropic/claude-haiku-4-5",
Ok(model_resp("anthropic/claude-haiku-4-5", " ")),
), (
"anthropic/claude-sonnet-5",
Ok(model_resp("anthropic/claude-sonnet-5", "answer")),
),
],
);
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
user_body(),
)
.await;
let json = body_json(resp).await;
assert_eq!(json["content"][0]["text"], "answer");
let trace = rx.try_recv().expect("trace enqueued");
assert_eq!(trace.attempts.len(), 2);
assert_eq!(trace.final_.escalations, 1);
assert_eq!(trace.final_.served_rung, Some(1));
}
#[tokio::test]
async fn enforce_all_rungs_error_returns_502() {
let (state, mut rx) = enforce_state(
&["anthropic/claude-haiku-4-5"],
&["non-empty"],
vec![(
"anthropic/claude-haiku-4-5",
Err(ProviderError::Transport("down".into())),
)],
);
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
user_body(),
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::BAD_GATEWAY);
assert!(rx.try_recv().is_ok());
}
#[tokio::test]
async fn no_routing_config_falls_through_to_observe_not_enforce() {
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_UPSTREAM_ANTHROPIC" => Some("http://127.0.0.1:1".to_owned()),
_ => None,
})
.unwrap();
let (traces, _rx) = mpsc::channel(64);
let state = AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::new("http://127.0.0.1:1", "http://127.0.0.1:1"),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter: None,
spill: None,
};
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
user_body(),
)
.await;
assert_ne!(resp.status(), axum::http::StatusCode::OK);
}
#[test]
fn detects_stream_requests() {
assert!(is_stream_request(br#"{"stream": true}"#));
assert!(!is_stream_request(br#"{"stream": false}"#));
assert!(!is_stream_request(br#"{"model":"m"}"#));
assert!(!is_stream_request(b"not json"));
}
#[test]
fn detects_tool_blocks_in_messages() {
let with =
br#"{"messages":[{"role":"user","content":[{"type":"tool_result","content":"42"}]}]}"#;
let without = br#"{"messages":[{"role":"user","content":"hi"}]}"#;
assert!(messages_have_tool_blocks(with));
assert!(!messages_have_tool_blocks(without));
}
#[test]
fn enforce_only_handles_plain_text() {
let plain =
Bytes::from_static(br#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#);
let tools = Bytes::from_static(
br#"{"model":"m","tools":[{"name":"t"}],"messages":[{"role":"user","content":"hi"}]}"#,
);
let f_plain = extract_features(&HeaderMap::new(), &plain);
let f_tools = extract_features(&HeaderMap::new(), &tools);
let anthropic_ladder = vec!["anthropic/claude-haiku-4-5".to_owned()];
let providers = test_registry();
assert!(enforce_can_handle(
&f_plain,
&plain,
false,
&anthropic_ladder,
&providers,
Dialect::Anthropic,
));
assert!(!enforce_can_handle(
&f_tools,
&tools,
false,
&anthropic_ladder,
&providers,
Dialect::Anthropic,
));
}
#[test]
fn structured_enforce_routes_tools_and_streaming() {
let tools = Bytes::from_static(
br#"{"model":"m","tools":[{"name":"t"}],"messages":[{"role":"user","content":"hi"}]}"#,
);
let streaming_tools = Bytes::from_static(
br#"{"model":"m","stream":true,"tools":[{"name":"t"}],"messages":[{"role":"user","content":"hi"}]}"#,
);
let f = extract_features(&HeaderMap::new(), &tools);
let anthropic_ladder = vec![
"anthropic/claude-haiku-4-5".to_owned(),
"anthropic/claude-sonnet-5".to_owned(),
];
let providers = test_registry();
assert!(enforce_can_handle(
&f,
&tools,
true,
&anthropic_ladder,
&providers,
Dialect::Anthropic,
));
assert!(enforce_can_handle(
&f,
&streaming_tools,
true,
&anthropic_ladder,
&providers,
Dialect::Anthropic,
));
}
fn test_registry() -> crate::provider::ProviderRegistry {
crate::provider::ProviderRegistry::new("http://localhost", "http://localhost")
}
#[test]
fn fidelity_guard_blocks_structured_on_non_verbatim_ladder() {
let tools = Bytes::from_static(
br#"{"model":"m","tools":[{"name":"t"}],"messages":[{"role":"user","content":"hi"}]}"#,
);
let plain =
Bytes::from_static(br#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#);
let f_tools = extract_features(&HeaderMap::new(), &tools);
let f_plain = extract_features(&HeaderMap::new(), &plain);
let providers = test_registry();
let mixed_ladder = vec![
"openai/gpt-4.1-mini".to_owned(),
"anthropic/claude-sonnet-5".to_owned(),
];
assert!(!enforce_can_handle(
&f_tools,
&tools,
true,
&mixed_ladder,
&providers,
Dialect::Anthropic,
));
assert!(enforce_can_handle(
&f_plain,
&plain,
true,
&mixed_ladder,
&providers,
Dialect::Anthropic,
));
}
#[test]
fn enforce_sse_reemission_preserves_text_and_tool_use() {
let resp = ModelResponse {
model: "anthropic/claude-haiku-4-5".to_owned(),
text: "let me check".to_owned(),
in_tokens: 5,
out_tokens: 7,
raw: serde_json::json!({
"content": [
{ "type": "text", "text": "let me check" },
{ "type": "tool_use", "id": "tu_1", "name": "get_weather", "input": { "city": "Paris" } }
]
}),
};
let sse = anthropic_sse_from_message(&anthropic_response_json(&resp));
let frames: Vec<Value> = sse
.lines()
.filter_map(|l| l.strip_prefix("data: "))
.map(|d| serde_json::from_str::<Value>(d).expect("each SSE data frame is valid JSON"))
.collect();
assert_eq!(frames.first().unwrap()["type"], "message_start");
assert_eq!(frames.last().unwrap()["type"], "message_stop");
assert!(frames.iter().any(|f| f["delta"]["type"] == "text_delta"
&& f["delta"]["text"] == "let me check"));
assert!(
frames
.iter()
.any(|f| f["content_block"]["type"] == "tool_use"
&& f["content_block"]["name"] == "get_weather"
&& f["content_block"]["id"] == "tu_1")
);
assert!(
frames
.iter()
.any(|f| f["delta"]["type"] == "input_json_delta"
&& f["delta"]["partial_json"] == r#"{"city":"Paris"}"#)
);
}
#[tokio::test]
async fn enforce_falls_back_to_observe_for_tool_requests() {
let toml = "[[route]]\nmatch = {}\nmode = \"enforce\"\nladder = [\"anthropic/m\"]\ngates = [\"non-empty\"]\n";
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_CONFIG_TOML" => Some(toml.to_owned()),
"FIRSTPASS_MODE" => Some("enforce".to_owned()),
"FIRSTPASS_UPSTREAM_ANTHROPIC" => Some("http://127.0.0.1:1".to_owned()),
_ => None,
})
.unwrap();
let mut outs = HashMap::new();
outs.insert(
"anthropic/m".to_owned(),
Ok(model_resp("anthropic/m", "hello")),
);
let mut map: HashMap<String, Arc<dyn Provider>> = HashMap::new();
map.insert(
"anthropic".to_owned(),
Arc::new(MockProvider::new("anthropic", outs)),
);
let (traces, _rx) = mpsc::channel(64);
let state = AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::from_map(map),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter: None,
spill: None,
};
let plain =
Bytes::from_static(br#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#);
let resp = messages(
State(state.clone()),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
plain,
)
.await;
assert_eq!(
resp.status(),
axum::http::StatusCode::OK,
"plain text should enforce"
);
let tools = Bytes::from_static(
br#"{"model":"m","tools":[{"name":"get_weather"}],"messages":[{"role":"user","content":"hi"}]}"#,
);
let resp = messages(
State(state.clone()),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
tools.clone(),
)
.await;
assert_eq!(
resp.status(),
axum::http::StatusCode::OK,
"tool request must route through enforce by default (ADR 0005 default-on)"
);
let toml_off = "[[route]]\nmatch = {}\nmode = \"enforce\"\nladder = [\"anthropic/m\"]\ngates = [\"non-empty\"]\n[escalation]\nenforce_structured = false\n";
let config_off = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_CONFIG_TOML" => Some(toml_off.to_owned()),
"FIRSTPASS_MODE" => Some("enforce".to_owned()),
"FIRSTPASS_UPSTREAM_ANTHROPIC" => Some("http://127.0.0.1:1".to_owned()),
_ => None,
})
.unwrap();
let state_off = AppState {
config: Arc::new(config_off),
..state.clone()
};
let resp = messages(
State(state_off),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
tools,
)
.await;
assert_ne!(
resp.status(),
axum::http::StatusCode::OK,
"with enforce_structured = false a tool request must fall back to observe"
);
let toolres = Bytes::from_static(
br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"x","content":"42"}]}]}"#,
);
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
toolres,
)
.await;
assert_eq!(
resp.status(),
axum::http::StatusCode::OK,
"tool_result blocks route through enforce by default too (verbatim carry)"
);
}
async fn feedback_state() -> (AppState, std::path::PathBuf, String) {
let db = std::env::temp_dir().join(format!("firstpass-feedback-{}.db", Uuid::now_v7()));
let (tx, handle) = crate::store::open(&db).unwrap();
let mut trace = build_error_trace(
&ProxyConfig::from_lookup(|_| None).unwrap(),
&Bytes::from_static(b"{}"),
5,
Some("sess-fb"),
);
trace.attempts.push(Attempt {
rung: 0,
model: "anthropic/claude-haiku-4-5".into(),
provider: "anthropic".into(),
in_tokens: 10,
out_tokens: 5,
cost_usd: 0.001,
latency_ms: 5,
gates: vec![],
verdict: Verdict::Pass,
});
let trace_id = trace.trace_id.to_string();
tx.try_send(trace).unwrap();
drop(tx);
handle.await.unwrap();
let db_str = db.to_string_lossy().into_owned();
let config = ProxyConfig::from_lookup(move |k| match k {
"FIRSTPASS_DB" => Some(db_str.clone()),
_ => None,
})
.unwrap();
let (traces, _rx) = mpsc::channel(64);
let state = AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::new("http://127.0.0.1:1", "http://127.0.0.1:1"),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter: None,
spill: None,
};
(state, db, trace_id)
}
#[tokio::test]
async fn feedback_nudges_the_adaptive_threshold() {
use firstpass_core::conformal::AdaptiveConformal;
let (mut state, _db, trace_id) = feedback_state().await;
let aci = Arc::new(std::sync::Mutex::new(AdaptiveConformal::new(0.1, 0.2, 0.5)));
state.adaptive = Some(aci.clone());
let before = aci.lock().unwrap().threshold();
let fail = Bytes::from(
serde_json::json!({ "trace_id": trace_id, "gate_id": "tests", "verdict": "fail", "reporter": "ci" })
.to_string(),
);
assert_eq!(
feedback(
State(state.clone()),
Extension(TenantId("default".to_owned())),
fail
)
.await
.status(),
axum::http::StatusCode::ACCEPTED
);
let after_fail = aci.lock().unwrap().threshold();
assert!(
after_fail > before,
"served fail should raise the live threshold: {before} -> {after_fail}"
);
let pass = Bytes::from(
serde_json::json!({ "trace_id": trace_id, "gate_id": "tests", "verdict": "pass", "reporter": "ci" })
.to_string(),
);
let _ = feedback(
State(state),
Extension(TenantId("default".to_owned())),
pass,
)
.await;
assert!(aci.lock().unwrap().threshold() < after_fail);
}
#[tokio::test]
async fn feedback_records_a_deferred_verdict_without_breaking_the_chain() {
let (state, db, trace_id) = feedback_state().await;
let body = Bytes::from(
serde_json::json!({
"trace_id": trace_id,
"gate_id": "tests",
"verdict": "pass",
"score": 1.0,
"reporter": "ci",
})
.to_string(),
);
let resp = feedback(
State(state),
Extension(TenantId("default".to_owned())),
body,
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::ACCEPTED);
let view = crate::store::load_trace_view(&db, "default", &trace_id)
.unwrap()
.unwrap();
assert_eq!(view.deferred.len(), 1);
assert_eq!(view.deferred[0].gate_id, "tests");
let traces = crate::store::load_all_traces(&db).unwrap();
firstpass_core::verify_chain(&traces, GENESIS_HASH).unwrap();
let _ = std::fs::remove_file(&db);
}
#[tokio::test]
async fn feedback_for_unknown_trace_is_404() {
let (state, db, _trace_id) = feedback_state().await;
let body = Bytes::from(
serde_json::json!({
"trace_id": "does-not-exist",
"gate_id": "tests",
"verdict": "pass",
"reporter": "ci",
})
.to_string(),
);
let resp = feedback(
State(state),
Extension(TenantId("default".to_owned())),
body,
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::NOT_FOUND);
let _ = std::fs::remove_file(&db);
}
#[tokio::test]
async fn feedback_across_tenants_is_404_not_403() {
let (state, db, trace_id) = feedback_state().await;
let body = Bytes::from(
serde_json::json!({
"trace_id": trace_id,
"gate_id": "tests",
"verdict": "pass",
"score": 1.0,
"reporter": "attacker",
})
.to_string(),
);
let resp = feedback(
State(state),
Extension(TenantId("tenant-b".to_owned())),
body,
)
.await;
assert_eq!(
resp.status(),
axum::http::StatusCode::NOT_FOUND,
"cross-tenant feedback must look exactly like a missing trace"
);
let _ = std::fs::remove_file(&db);
}
#[tokio::test]
async fn feedback_rejects_bad_verdict_and_score() {
let (state, db, trace_id) = feedback_state().await;
let bad_verdict = Bytes::from(
serde_json::json!({ "trace_id": trace_id, "gate_id": "g", "verdict": "maybe", "reporter": "x" })
.to_string(),
);
assert_eq!(
feedback(
State(state.clone()),
Extension(TenantId("default".to_owned())),
bad_verdict
)
.await
.status(),
axum::http::StatusCode::BAD_REQUEST
);
let bad_score = Bytes::from(
serde_json::json!({ "trace_id": trace_id, "gate_id": "g", "verdict": "pass", "score": 9.0, "reporter": "x" })
.to_string(),
);
assert_eq!(
feedback(
State(state),
Extension(TenantId("default".to_owned())),
bad_score
)
.await
.status(),
axum::http::StatusCode::BAD_REQUEST
);
let _ = std::fs::remove_file(&db);
}
#[tokio::test]
async fn metrics_endpoint_renders_after_a_real_request() {
use tower::ServiceExt;
let (state, mut rx) = enforce_state(
&["anthropic/claude-haiku-4-5"],
&["non-empty"],
vec![(
"anthropic/claude-haiku-4-5",
Ok(model_resp("anthropic/claude-haiku-4-5", "hello")),
)],
);
let router = app(state).expect("prometheus recorder installs");
let req = axum::http::Request::builder()
.method("POST")
.uri("/v1/messages")
.header("content-type", "application/json")
.body(Body::from(user_body()))
.unwrap();
let resp = router.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
rx.try_recv().expect("a trace was enqueued");
let metrics_req = axum::http::Request::builder()
.method("GET")
.uri("/metrics")
.body(Body::empty())
.unwrap();
let metrics_resp = router.oneshot(metrics_req).await.unwrap();
assert_eq!(metrics_resp.status(), axum::http::StatusCode::OK);
let bytes = axum::body::to_bytes(metrics_resp.into_body(), 1 << 20)
.await
.unwrap();
let body = String::from_utf8(bytes.to_vec()).unwrap();
assert!(
body.contains("firstpass_enforce_latency_ms"),
"metrics body missing enforce latency histogram: {body}"
);
assert!(
body.contains("firstpass_served_total"),
"metrics body missing served counter: {body}"
);
}
fn auth_state(require_auth: bool, keys_json: Option<String>) -> AppState {
auth_state_rated(require_auth, keys_json, None)
}
fn auth_state_rated(
require_auth: bool,
keys_json: Option<String>,
rate_per_sec: Option<u32>,
) -> AppState {
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_REQUIRE_AUTH" => require_auth.then(|| "true".to_owned()),
"FIRSTPASS_TENANT_KEYS_JSON" => keys_json.clone(),
"FIRSTPASS_TENANT_RATE_PER_SEC" => rate_per_sec.map(|n| n.to_string()),
_ => None,
})
.unwrap();
let (traces, _rx) = mpsc::channel(64);
std::mem::forget(_rx);
let providers: HashMap<String, Arc<dyn Provider>> = HashMap::new();
let tenant_rate_limiter = build_tenant_rate_limiter(&config);
AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::from_map(providers),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter,
spill: None,
}
}
fn cap_request(auth_header: Option<&str>) -> axum::http::Request<Body> {
let mut b = axum::http::Request::builder()
.method("GET")
.uri("/v1/capabilities");
if let Some(h) = auth_header {
b = b.header("authorization", h);
}
b.body(Body::empty()).unwrap()
}
#[tokio::test]
async fn auth_off_allows_unauthenticated_request() {
use tower::ServiceExt;
let router = app(auth_state(false, None)).expect("router");
let resp = router.oneshot(cap_request(None)).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
}
#[tokio::test]
async fn auth_on_missing_key_is_401_opaque() {
use tower::ServiceExt;
let hash = crate::tenant_auth::TenantKeys::hash_key("key-a").unwrap();
let keys = format!("{{\"tenant-a\": {hash:?}}}");
let router = app(auth_state(true, Some(keys))).expect("router");
let resp = router.oneshot(cap_request(None)).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::UNAUTHORIZED);
let json = body_json(resp).await;
assert_eq!(json["error"]["type"], "unauthorized");
let msg = json["error"]["message"].as_str().unwrap();
assert!(!msg.contains("tenant"), "no tenant oracle in body: {msg}");
}
#[tokio::test]
async fn auth_on_invalid_key_is_401() {
use tower::ServiceExt;
let hash = crate::tenant_auth::TenantKeys::hash_key("key-a").unwrap();
let keys = format!("{{\"tenant-a\": {hash:?}}}");
let router = app(auth_state(true, Some(keys))).expect("router");
let resp = router
.oneshot(cap_request(Some("Bearer wrong-key")))
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn auth_on_valid_key_proceeds() {
use tower::ServiceExt;
let hash = crate::tenant_auth::TenantKeys::hash_key("key-a").unwrap();
let keys = format!("{{\"tenant-a\": {hash:?}}}");
let router = app(auth_state(true, Some(keys))).expect("router");
let resp = router
.oneshot(cap_request(Some("Bearer tenant-a.key-a")))
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
}
fn two_tenant_state(rate_per_sec: Option<u32>) -> AppState {
let hash_a = crate::tenant_auth::TenantKeys::hash_key("key-a").unwrap();
let hash_b = crate::tenant_auth::TenantKeys::hash_key("key-b").unwrap();
let keys = format!("{{\"tenant-a\": {hash_a:?}, \"tenant-b\": {hash_b:?}}}");
auth_state_rated(true, Some(keys), rate_per_sec)
}
#[tokio::test]
async fn tenant_exceeding_rate_limit_gets_429_opaque() {
use tower::ServiceExt;
let router = app(two_tenant_state(Some(1))).expect("router");
let (r1, r2, r3, r4) = tokio::join!(
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
);
let responses = [r1.unwrap(), r2.unwrap(), r3.unwrap(), r4.unwrap()];
let ok = responses
.iter()
.filter(|r| r.status() == axum::http::StatusCode::OK)
.count();
assert!(ok >= 1, "the burst's first request must pass");
let limited: Vec<_> = responses
.into_iter()
.filter(|r| r.status() == axum::http::StatusCode::TOO_MANY_REQUESTS)
.collect();
assert!(
!limited.is_empty(),
"a 4-request burst against 1 req/sec must trip the limiter"
);
let json = body_json(limited.into_iter().next().unwrap()).await;
assert_eq!(json["error"]["type"], "rate_limited");
let msg = json["error"]["message"].as_str().unwrap();
assert!(!msg.contains('1'), "no limit value in body: {msg}");
}
#[tokio::test]
async fn rate_limit_buckets_are_independent_per_tenant() {
use tower::ServiceExt;
let router = app(two_tenant_state(Some(1))).expect("router");
let (a1, a2, a3) = tokio::join!(
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a"))),
);
let a_limited = [a1.unwrap(), a2.unwrap(), a3.unwrap()]
.iter()
.filter(|r| r.status() == axum::http::StatusCode::TOO_MANY_REQUESTS)
.count();
assert!(a_limited >= 1, "tenant A's burst must trip its limiter");
let b1 = router
.clone()
.oneshot(cap_request(Some("Bearer tenant-b.key-b")))
.await
.unwrap();
assert_eq!(b1.status(), axum::http::StatusCode::OK);
}
#[tokio::test]
async fn rate_limit_unset_never_429s() {
use tower::ServiceExt;
let router = app(two_tenant_state(None)).expect("router");
for _ in 0..20 {
let resp = router
.clone()
.oneshot(cap_request(Some("Bearer tenant-a.key-a")))
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
}
}
#[test]
fn u01_is_deterministic_and_in_range() {
let s1 = u01(0xDEAD_BEEF_CAFE_1234_u128);
let s2 = u01(0xDEAD_BEEF_CAFE_1234_u128);
assert_eq!(s1, s2, "u01 must be deterministic for the same seed");
assert!((0.0..1.0).contains(&s1), "u01 must return [0, 1), got {s1}");
let s3 = u01(0x1234_5678_9ABC_DEF0_u128);
assert_ne!(s1, s3, "different seeds should give different values");
for i in 0u64..256 {
let v = u01(i as u128);
assert!((0.0..1.0).contains(&v), "seed {i}: u01={v} out of [0,1)");
}
}
#[test]
fn epsilon_propensity_formula() {
let epsilon = 0.2_f64;
let k = 3_usize;
let greedy = 1_u32;
let p_greedy = epsilon_propensity(greedy, greedy, epsilon, k);
let expected_greedy = (1.0 - epsilon) + epsilon / k as f64;
assert!(
(p_greedy - expected_greedy).abs() < 1e-12,
"{p_greedy} != {expected_greedy}"
);
let p_other = epsilon_propensity(0, greedy, epsilon, k);
let expected_other = epsilon / k as f64;
assert!(
(p_other - expected_other).abs() < 1e-12,
"{p_other} != {expected_other}"
);
for chosen in 0..k as u32 {
let p = epsilon_propensity(chosen, greedy, epsilon, k);
assert!(
p > 0.0 && p <= 1.0,
"propensity {p} out of (0,1] for chosen={chosen}"
);
}
}
#[test]
fn epsilon_branch_and_greedy_branch_both_occur_over_many_seeds() {
let epsilon = 0.3_f64;
let mut saw_explore = false;
let mut saw_greedy = false;
for i in 0u64..200 {
let u = u01(i as u128);
if u < epsilon {
saw_explore = true;
} else {
saw_greedy = true;
}
if saw_explore && saw_greedy {
break;
}
}
assert!(
saw_explore,
"epsilon branch must fire with epsilon=0.3 over 200 seeds"
);
assert!(
saw_greedy,
"greedy branch must occur with epsilon=0.3 over 200 seeds"
);
}
#[test]
fn epsilon_propensity_sums_to_one_over_all_rungs() {
let epsilon = 0.15_f64;
let k = 4_usize;
let greedy = 2_u32;
let total: f64 = (0..k as u32)
.map(|r| epsilon_propensity(r, greedy, epsilon, k))
.sum();
assert!(
(total - 1.0).abs() < 1e-12,
"propensities must sum to 1, got {total}"
);
}
#[tokio::test(start_paused = true)]
async fn keepalive_stream_ticks_then_emits_final_frame() {
let (tx, rx) = tokio::sync::oneshot::channel::<Result<Value, ProxyError>>();
let mut ticks = tokio::time::interval(SSE_KEEPALIVE_EVERY);
ticks.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ticks.reset();
let mut stream = KeepaliveStream {
rx: Some(rx),
ticks,
format_message: anthropic_sse_from_message,
};
async fn next(
stream: &mut KeepaliveStream,
) -> Option<Option<Result<Bytes, std::convert::Infallible>>> {
std::future::poll_fn(|cx| {
std::task::Poll::Ready(
match futures_core::Stream::poll_next(std::pin::Pin::new(&mut *stream), cx) {
std::task::Poll::Ready(item) => Some(item),
std::task::Poll::Pending => None,
},
)
})
.await
}
assert!(
next(&mut stream).await.is_none(),
"no frame before an interval"
);
tokio::time::advance(SSE_KEEPALIVE_EVERY + Duration::from_millis(1)).await;
let frame = next(&mut stream)
.await
.expect("keepalive due")
.unwrap()
.unwrap();
assert!(
frame.starts_with(b": "),
"keepalive must be an SSE comment (ignored by every conforming parser)"
);
let message = serde_json::json!({
"id": "msg_1", "type": "message", "role": "assistant", "model": "m",
"content": [{ "type": "text", "text": "done" }],
"usage": { "input_tokens": 1, "output_tokens": 1 }
});
tx.send(Ok(message)).unwrap();
let frame = next(&mut stream)
.await
.expect("final frame")
.unwrap()
.unwrap();
let text = String::from_utf8(frame.to_vec()).unwrap();
assert!(text.contains("event: message_start"));
assert!(text.contains("event: message_stop"));
let eos = next(&mut stream)
.await
.expect("stream must end after the final frame");
assert!(eos.is_none(), "end-of-stream after the final frame");
}
#[tokio::test]
async fn streaming_enforce_serves_full_sse_sequence() {
let toml = "[[route]]\nmatch = {}\nmode = \"enforce\"\nladder = [\"anthropic/m\"]\ngates = [\"non-empty\"]\n";
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_CONFIG_TOML" => Some(toml.to_owned()),
"FIRSTPASS_MODE" => Some("enforce".to_owned()),
"FIRSTPASS_UPSTREAM_ANTHROPIC" => Some("http://127.0.0.1:1".to_owned()),
_ => None,
})
.unwrap();
let mut outs = HashMap::new();
outs.insert(
"anthropic/m".to_owned(),
Ok(model_resp("anthropic/m", "gated answer")),
);
let mut map: HashMap<String, Arc<dyn Provider>> = HashMap::new();
map.insert(
"anthropic".to_owned(),
Arc::new(MockProvider::new("anthropic", outs)),
);
let (traces, _rx) = mpsc::channel(64);
let state = AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::from_map(map),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter: None,
spill: None,
};
let body = Bytes::from_static(
br#"{"model":"m","stream":true,"messages":[{"role":"user","content":"hi"}]}"#,
);
let resp = messages(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
body,
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(
resp.headers()
.get(axum::http::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.is_some_and(|ct| ct.starts_with("text/event-stream")),
"streaming client must get SSE"
);
let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
.await
.unwrap();
let text = String::from_utf8(bytes.to_vec()).unwrap();
assert!(text.contains("event: message_start"));
assert!(text.contains("gated answer"));
assert!(text.contains("event: message_stop"));
}
#[test]
fn parse_openai_request_plain_text() {
let body = br#"{"model":"gpt-4o","max_tokens":256,"messages":[{"role":"user","content":"hello"}]}"#;
let req = parse_openai_request(body, false).expect("must parse");
assert_eq!(req.model, "gpt-4o");
assert_eq!(req.max_tokens, 256);
assert_eq!(req.messages.len(), 1);
assert_eq!(req.messages[0].role, "user");
assert_eq!(req.messages[0].content, Value::String("hello".to_owned()));
assert!(req.system.is_none());
assert_eq!(req.raw, Value::Null);
}
#[test]
fn parse_openai_request_system_message() {
let body = br#"{"model":"gpt-4o","messages":[{"role":"system","content":"be concise"},{"role":"user","content":"hi"}]}"#;
let req = parse_openai_request(body, false).expect("must parse");
assert_eq!(req.system.as_deref(), Some("be concise"));
assert_eq!(req.messages.len(), 1);
assert_eq!(req.messages[0].role, "user");
}
#[test]
fn parse_openai_request_tool_calls_translate_to_tool_use() {
let body = br#"{
"model":"gpt-4o",
"messages":[
{"role":"user","content":"what's the weather?"},
{"role":"assistant","content":null,"tool_calls":[
{"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"city\":\"Paris\"}"}}
]},
{"role":"tool","tool_call_id":"call_1","content":"15C, cloudy"}
]
}"#;
let req = parse_openai_request(body, false).expect("must parse");
assert_eq!(req.messages.len(), 3);
let asst = &req.messages[1];
assert_eq!(asst.role, "assistant");
let blocks = asst.content.as_array().expect("content array");
assert_eq!(blocks[0]["type"], "tool_use");
assert_eq!(blocks[0]["name"], "get_weather");
assert_eq!(blocks[0]["id"], "call_1");
assert_eq!(blocks[0]["input"]["city"], "Paris");
let tool_msg = &req.messages[2];
assert_eq!(tool_msg.role, "user");
let result_blocks = tool_msg.content.as_array().expect("result blocks");
assert_eq!(result_blocks[0]["type"], "tool_result");
assert_eq!(result_blocks[0]["tool_use_id"], "call_1");
}
#[test]
fn parse_openai_request_tools_translate_to_anthropic_format() {
let body = br#"{
"model":"gpt-4o",
"messages":[{"role":"user","content":"use a tool"}],
"tools":[{"type":"function","function":{"name":"search","description":"web search","parameters":{"type":"object","properties":{"q":{"type":"string"}}}}}]
}"#;
let req = parse_openai_request(body, false).expect("must parse");
let tools = req.tools.as_array().expect("tools array");
assert_eq!(tools.len(), 1);
assert_eq!(tools[0]["name"], "search");
assert_eq!(tools[0]["description"], "web search");
assert_eq!(tools[0]["input_schema"]["type"], "object");
}
#[test]
fn parse_openai_request_raw_carry_preserves_full_body() {
let body =
br#"{"model":"gpt-4o","max_tokens":100,"messages":[{"role":"user","content":"hi"}]}"#;
let req = parse_openai_request(body, true).expect("must parse");
assert!(req.raw.is_object(), "raw must be the full JSON object");
assert_eq!(req.raw["model"], "gpt-4o");
assert_eq!(req.raw["max_tokens"], 100);
assert!(req.tools.is_null(), "no tools in this request");
}
#[test]
fn parse_openai_request_http_image_returns_none() {
let body = br#"{"model":"gpt-4o","messages":[{"role":"user","content":[
{"type":"text","text":"describe this"},
{"type":"image_url","image_url":{"url":"https://example.com/cat.png"}}
]}]}"#;
let result = parse_openai_request(body, false);
assert!(result.is_none(), "http image URL must fail translation");
}
#[test]
fn parse_openai_request_data_url_image_translates_to_anthropic_base64() {
let body = br#"{"model":"gpt-4o","messages":[{"role":"user","content":[
{"type":"text","text":"describe"},
{"type":"image_url","image_url":{"url":"data:image/png;base64,iVBORw0KGgo="}}
]}]}"#;
let req = parse_openai_request(body, false).expect("data URL must parse");
let blocks = req.messages[0].content.as_array().expect("blocks");
let img = blocks
.iter()
.find(|b| b["type"] == "image")
.expect("image block");
assert_eq!(img["source"]["type"], "base64");
assert_eq!(img["source"]["media_type"], "image/png");
assert_eq!(img["source"]["data"], "iVBORw0KGgo=");
}
#[test]
fn openai_response_json_renders_text_response() {
let resp = ModelResponse {
model: "gpt-4o".to_owned(),
text: "Hello!".to_owned(),
in_tokens: 10,
out_tokens: 5,
raw: serde_json::json!({
"content": [{ "type": "text", "text": "Hello!" }]
}),
};
let json = openai_response_json(&resp);
assert_eq!(json["object"], "chat.completion");
assert_eq!(json["model"], "gpt-4o");
assert_eq!(json["choices"][0]["message"]["role"], "assistant");
assert_eq!(json["choices"][0]["message"]["content"], "Hello!");
assert_eq!(json["choices"][0]["finish_reason"], "stop");
assert_eq!(json["usage"]["prompt_tokens"], 10);
assert_eq!(json["usage"]["completion_tokens"], 5);
}
#[test]
fn openai_response_json_renders_tool_call() {
let resp = ModelResponse {
model: "gpt-4o".to_owned(),
text: String::new(),
in_tokens: 20,
out_tokens: 15,
raw: serde_json::json!({
"content": [{
"type": "tool_use",
"id": "call_abc",
"name": "search",
"input": {"q": "Rust async"}
}]
}),
};
let json = openai_response_json(&resp);
assert_eq!(json["choices"][0]["finish_reason"], "tool_calls");
let tc = &json["choices"][0]["message"]["tool_calls"][0];
assert_eq!(tc["id"], "call_abc");
assert_eq!(tc["type"], "function");
assert_eq!(tc["function"]["name"], "search");
assert_eq!(json["choices"][0]["message"]["content"], Value::Null);
}
#[test]
fn openai_sse_from_message_plain_text_has_role_then_content_then_stop() {
let resp = ModelResponse {
model: "gpt-4o".to_owned(),
text: "Hi there!".to_owned(),
in_tokens: 5,
out_tokens: 3,
raw: serde_json::json!({
"content": [{ "type": "text", "text": "Hi there!" }]
}),
};
let sse = openai_sse_from_message(&openai_response_json(&resp));
for line in sse.lines() {
assert!(
line.is_empty() || line.starts_with("data: "),
"bad SSE line: {line:?}"
);
}
let frames: Vec<&str> = sse
.lines()
.filter_map(|l| l.strip_prefix("data: "))
.collect();
assert_eq!(*frames.last().unwrap(), "[DONE]");
let role_frame: Value = serde_json::from_str(frames[0]).unwrap();
assert_eq!(role_frame["choices"][0]["delta"]["role"], "assistant");
assert!(frames.iter().any(|f| {
if *f == "[DONE]" {
return false;
}
serde_json::from_str::<Value>(f)
.ok()
.is_some_and(|v| v["choices"][0]["delta"]["content"] == "Hi there!")
}));
assert!(frames.iter().any(|f| {
if *f == "[DONE]" {
return false;
}
serde_json::from_str::<Value>(f)
.ok()
.is_some_and(|v| v["choices"][0]["finish_reason"] == "stop")
}));
}
#[test]
fn detects_openai_tool_calls() {
let with_tool_calls = Bytes::from_static(br#"{"messages":[
{"role":"user","content":"hi"},
{"role":"assistant","content":null,"tool_calls":[{"id":"c1","type":"function","function":{"name":"f","arguments":"{}"}}]}
]}"#);
let without = Bytes::from_static(br#"{"messages":[{"role":"user","content":"hi"}]}"#);
let with_tool_msg = Bytes::from_static(
br#"{"messages":[
{"role":"tool","tool_call_id":"c1","content":"result"}
]}"#,
);
assert!(openai_messages_have_tool_calls(&with_tool_calls));
assert!(!openai_messages_have_tool_calls(&without));
assert!(openai_messages_have_tool_calls(&with_tool_msg));
}
#[test]
fn detects_openai_http_images() {
let http_img = Bytes::from_static(
br#"{"messages":[{"role":"user","content":[
{"type":"image_url","image_url":{"url":"https://example.com/img.png"}}
]}]}"#,
);
let data_img = Bytes::from_static(
br#"{"messages":[{"role":"user","content":[
{"type":"image_url","image_url":{"url":"data:image/png;base64,abc"}}
]}]}"#,
);
let no_img = Bytes::from_static(br#"{"messages":[{"role":"user","content":"hi"}]}"#);
assert!(openai_has_http_images(&http_img));
assert!(!openai_has_http_images(&data_img));
assert!(!openai_has_http_images(&no_img));
}
#[test]
fn enforce_can_handle_openai_inbound_all_openai_ladder() {
let tools_body = Bytes::from_static(br#"{"model":"gpt-4o","messages":[{"role":"assistant","content":null,"tool_calls":[{"id":"c","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"#);
let f = extract_openai_features(&HeaderMap::new(), &tools_body);
let ladder = vec!["openai/gpt-4o-mini".to_owned(), "openai/gpt-4o".to_owned()];
let providers = crate::provider::ProviderRegistry::new("http://x", "http://x");
assert!(enforce_can_handle(
&f,
&tools_body,
true,
&ladder,
&providers,
Dialect::Openai
));
}
#[test]
fn enforce_can_handle_openai_inbound_all_anthropic_ladder_no_http_image() {
let tools_body = Bytes::from_static(br#"{"model":"gpt-4o","messages":[{"role":"assistant","content":null,"tool_calls":[{"id":"c","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"#);
let f = extract_openai_features(&HeaderMap::new(), &tools_body);
let ladder = vec!["anthropic/claude-haiku-4-5".to_owned()];
let providers = crate::provider::ProviderRegistry::new("http://x", "http://x");
assert!(enforce_can_handle(
&f,
&tools_body,
true,
&ladder,
&providers,
Dialect::Openai,
));
}
#[test]
fn enforce_can_handle_openai_inbound_http_image_falls_back() {
let img_body = Bytes::from_static(
br#"{"model":"gpt-4o","messages":[{"role":"user","content":[
{"type":"image_url","image_url":{"url":"https://example.com/img.png"}}
]}]}"#,
);
let f = extract_openai_features(&HeaderMap::new(), &img_body);
let ladder = vec!["anthropic/claude-haiku-4-5".to_owned()];
let providers = crate::provider::ProviderRegistry::new("http://x", "http://x");
assert!(!enforce_can_handle(
&f,
&img_body,
true,
&ladder,
&providers,
Dialect::Openai,
));
}
fn openai_enforce_state(mock_resp: ModelResponse) -> AppState {
let toml = "[[route]]\nmatch = {}\nmode = \"enforce\"\nladder = [\"mock/m\"]\ngates = [\"non-empty\"]\n";
let config = ProxyConfig::from_lookup(|k| match k {
"FIRSTPASS_CONFIG_TOML" => Some(toml.to_owned()),
"FIRSTPASS_MODE" => Some("enforce".to_owned()),
_ => None,
})
.unwrap();
let mut outs = HashMap::new();
outs.insert("mock/m".to_owned(), Ok(mock_resp));
let mut map: HashMap<String, Arc<dyn Provider>> = HashMap::new();
map.insert("mock".to_owned(), Arc::new(MockProvider::new("mock", outs)));
let (traces, _rx) = mpsc::channel(64);
std::mem::forget(_rx);
let tenant_rate_limiter = build_tenant_rate_limiter(&config);
AppState {
config: Arc::new(config),
http: reqwest::Client::new(),
providers: ProviderRegistry::from_map(map),
gate_health: Arc::new(GateHealthRegistry::new()),
traces,
adaptive: None,
bandit: None,
tenant_rate_limiter,
spill: None,
}
}
#[tokio::test]
async fn chat_completions_plain_text_enforce_returns_openai_shape() {
let mock = model_resp("mock/m", "gated answer");
let state = openai_enforce_state(mock);
let body = Bytes::from_static(
br#"{"model":"gpt-4o","messages":[{"role":"user","content":"hello"}]}"#,
);
let resp = chat_completions(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
body,
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::OK);
let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
.await
.unwrap();
let json: Value = serde_json::from_slice(&bytes).expect("must be JSON");
assert_eq!(json["object"], "chat.completion");
assert_eq!(json["choices"][0]["message"]["role"], "assistant");
assert_eq!(json["choices"][0]["message"]["content"], "gated answer");
assert_eq!(json["choices"][0]["finish_reason"], "stop");
assert!(
json["id"]
.as_str()
.is_some_and(|id| id.starts_with("chatcmpl-"))
);
}
#[tokio::test]
async fn chat_completions_stream_true_returns_sse_with_openai_chunks() {
let mock = model_resp("mock/m", "gated answer");
let state = openai_enforce_state(mock);
let body = Bytes::from_static(
br#"{"model":"gpt-4o","stream":true,"messages":[{"role":"user","content":"hello"}]}"#,
);
let resp = chat_completions(
State(state),
Extension(TenantId("default".to_owned())),
HeaderMap::new(),
body,
)
.await;
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(
resp.headers()
.get(axum::http::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.is_some_and(|ct| ct.starts_with("text/event-stream")),
"stream:true must return SSE"
);
let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
.await
.unwrap();
let text = String::from_utf8(bytes.to_vec()).unwrap();
assert!(
text.contains("chat.completion.chunk"),
"must have OpenAI chunk frames"
);
assert!(text.contains("[DONE]"), "must end with [DONE]");
assert!(
text.contains("gated answer"),
"content must be in the stream"
);
assert!(
!text.contains("message_start"),
"must not have Anthropic event types"
);
}
}