1pub mod auth;
57pub(crate) mod constrained;
58pub(crate) mod darklane;
62pub(crate) mod health;
66pub(crate) mod lanes {
67 pub use memra_lanes::*;
68}
69mod admit_predict;
74mod affinity;
81mod anthropic;
87#[allow(dead_code)] mod build_id;
92mod dsv4_serve;
93mod embed_api;
94pub mod metering;
101mod responses_api;
102mod surfaces;
103mod toolcall;
104mod ttft;
105mod worker;
106
107use std::collections::HashMap;
108use std::net::{SocketAddr, ToSocketAddrs};
109use std::sync::mpsc::Sender;
110use std::sync::{Arc, Mutex};
111
112use axum::{
113 Extension, Json, Router,
114 body::Body,
115 extract::{DefaultBodyLimit, FromRequest, Query, Request as AxumRequest, State},
116 http::{
117 HeaderMap, StatusCode,
118 header::{CONTENT_LENGTH, CONTENT_TYPE, TRANSFER_ENCODING},
119 },
120 middleware::{self, Next},
121 response::{
122 IntoResponse, Response,
123 sse::{Event as SseEvent, Sse},
124 },
125 routing::{get, post},
126};
127use futures_core::Stream as _;
128use serde::de::DeserializeOwned;
129use serde::{Deserialize, Serialize};
130use serde_json::json;
131use tower::ServiceExt as _;
132
133use memra_engine::decode::GenParams;
134use memra_engine::sampler::SamplerConfig;
135use memra_tokenizer::{
136 Tokenizer,
137 chat::{self, ThinkMode, ToolCall as TmplToolCall, Turn as TmplTurn},
138};
139use toolcall::{ParsedToolCall, Piece, ToolStreamParser};
140use worker::{Cmd, Event, ModelCaps, Request, SharedMetrics};
141
142const MAX_BODY_BYTES: usize = 192 * 1024 * 1024;
166const MAX_BODY_ADMISSIONS: usize = 4;
167const MAX_SMALL_BODY_ADMISSIONS: usize = 32;
168#[allow(clippy::identity_op)] const BODY_ADMISSION_BYPASS_BYTES: usize = 1 * 1024 * 1024;
173const BODY_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(90);
174const BODY_READ_RATE_BYTES_PER_SEC: u64 = 2 * 1024 * 1024;
175const BODY_READ_TIMEOUT_MAX: std::time::Duration = std::time::Duration::from_secs(180);
176const BODY_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
177const BODY_ADMISSION_RETRY_AFTER_S: u64 = 1;
178const MAX_STOP_SEQUENCES: usize = 16;
179const MAX_STOP_SEQUENCE_BYTES: usize = 1_024;
180const MAX_STOP_SEQUENCES_BYTES: usize = 4 * 1_024;
181const MAX_CLIENT_IDENTIFIER_BYTES: usize = 256;
182const MAX_HTTP_CONNECTIONS: usize = 1_024;
183const MAX_HTTP2_STREAMS_PER_CONNECTION: u32 = 128;
184const HTTP1_HEADER_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
185const HTTP_CONNECTION_MAX_LIFETIME: std::time::Duration = std::time::Duration::from_secs(300);
186
187fn body_admission_semaphore() -> Arc<tokio::sync::Semaphore> {
188 static SEMAPHORE: std::sync::OnceLock<Arc<tokio::sync::Semaphore>> = std::sync::OnceLock::new();
189 SEMAPHORE
190 .get_or_init(|| Arc::new(tokio::sync::Semaphore::new(MAX_BODY_ADMISSIONS)))
191 .clone()
192}
193
194fn small_body_admission_semaphore() -> Arc<tokio::sync::Semaphore> {
195 static SEMAPHORE: std::sync::OnceLock<Arc<tokio::sync::Semaphore>> = std::sync::OnceLock::new();
196 SEMAPHORE
197 .get_or_init(|| Arc::new(tokio::sync::Semaphore::new(MAX_SMALL_BODY_ADMISSIONS)))
198 .clone()
199}
200
201#[derive(Clone)]
202pub(crate) struct BodyAdmissionGuard {
203 permit: Arc<Mutex<Option<tokio::sync::OwnedSemaphorePermit>>>,
204}
205
206impl BodyAdmissionGuard {
207 fn new(permit: tokio::sync::OwnedSemaphorePermit) -> Self {
208 Self {
209 permit: Arc::new(Mutex::new(Some(permit))),
210 }
211 }
212
213 pub(crate) fn release(&self) {
214 if let Ok(mut permit) = self.permit.lock() {
215 permit.take();
216 }
217 }
218}
219
220pub(crate) struct BodyAdmissionLease(Option<BodyAdmissionGuard>);
221
222impl BodyAdmissionLease {
223 fn release(&mut self) {
224 if let Some(admission) = self.0.take() {
225 admission.release();
226 }
227 }
228
229 pub(crate) fn guard(&self) -> Option<&BodyAdmissionGuard> {
230 self.0.as_ref()
231 }
232}
233
234impl Drop for BodyAdmissionLease {
235 fn drop(&mut self) {
236 self.release();
237 }
238}
239
240pub(crate) struct AdmittedJson<T>(pub(crate) T, pub(crate) BodyAdmissionLease);
241
242#[axum::async_trait]
243impl<S, T> FromRequest<S> for AdmittedJson<T>
244where
245 S: Send + Sync,
246 T: DeserializeOwned,
247{
248 type Rejection = axum::extract::rejection::JsonRejection;
249
250 async fn from_request(req: AxumRequest, state: &S) -> Result<Self, Self::Rejection> {
251 let admission = req.extensions().get::<BodyAdmissionGuard>().cloned();
252 let parsed = Json::<T>::from_request(req, state).await;
253 parsed.map(|Json(value)| Self(value, BodyAdmissionLease(admission)))
254 }
255}
256
257fn declared_body_length(req: &AxumRequest) -> Option<usize> {
258 req.headers()
259 .get(CONTENT_LENGTH)
260 .and_then(|value| value.to_str().ok())
261 .and_then(|value| value.parse().ok())
262}
263
264fn body_requires_admission(req: &AxumRequest) -> bool {
265 if req.headers().contains_key(TRANSFER_ENCODING) {
269 return true;
270 }
271 declared_body_length(req).is_none_or(|length| length > BODY_ADMISSION_BYPASS_BYTES)
272}
273
274fn body_read_timeout(req: &AxumRequest) -> std::time::Duration {
278 let Some(length) = declared_body_length(req) else {
279 return BODY_READ_TIMEOUT;
280 };
281 let bytes = length as u64;
282 let extra_seconds =
283 bytes.saturating_add(BODY_READ_RATE_BYTES_PER_SEC - 1) / BODY_READ_RATE_BYTES_PER_SEC;
284 let seconds = BODY_READ_TIMEOUT
285 .as_secs()
286 .saturating_add(extra_seconds)
287 .min(BODY_READ_TIMEOUT_MAX.as_secs());
288 std::time::Duration::from_secs(seconds)
289}
290
291async fn shape_payload_too_large(req: AxumRequest, next: Next) -> Response {
296 let resp = next.run(req).await;
297 if resp.status() != StatusCode::PAYLOAD_TOO_LARGE {
298 return resp;
299 }
300 error_response_coded(
301 StatusCode::PAYLOAD_TOO_LARGE,
302 &format!(
303 "request body exceeds the {} MiB limit",
304 MAX_BODY_BYTES / (1024 * 1024)
305 ),
306 "invalid_request_error",
307 None,
308 Some("request_too_large"),
309 )
310}
311
312fn apply_body_limit(app: Router) -> Router {
315 app.layer(DefaultBodyLimit::max(MAX_BODY_BYTES))
316 .layer(middleware::from_fn(shape_payload_too_large))
317}
318
319fn protected_inference_path(path: &str) -> bool {
320 matches!(
321 path,
322 "/v1/auth/check"
323 | "/v1/completions"
324 | "/v1/chat/completions"
325 | "/v1/messages"
326 | "/v1/responses"
327 | "/v1/embeddings"
328 | "/v1/rerank"
329 )
330}
331
332async fn shape_inference_early_response(path: &str, response: Response) -> Response {
336 let request_id = Envelope::new(path != "/v1/completions");
337 if path == "/v1/messages" {
338 anthropic::with_anthropic_request_id(
339 &request_id.id,
340 anthropic::reshape_error(response, &request_id.id).await,
341 )
342 } else {
343 with_request_id(&request_id.id, response)
344 }
345}
346
347async fn authenticate_inference_before_body(
352 State(st): State<AppState>,
353 mut req: AxumRequest,
354 next: Next,
355) -> Response {
356 if !protected_inference_path(req.uri().path()) {
357 return next.run(req).await;
358 }
359 let path = req.uri().path().to_string();
360 if declared_body_length(&req).is_some_and(|length| length > MAX_BODY_BYTES) {
364 return shape_inference_early_response(
365 &path,
366 error_response_coded(
367 StatusCode::PAYLOAD_TOO_LARGE,
368 &format!(
369 "request body exceeds the {} MiB limit",
370 MAX_BODY_BYTES / (1024 * 1024)
371 ),
372 "invalid_request_error",
373 None,
374 Some("request_too_large"),
375 ),
376 )
377 .await;
378 }
379 let headers = req.headers();
380 let bearer = bearer_token(headers);
381 let auth = if matches!(path.as_str(), "/v1/messages" | "/v1/auth/check") {
382 let api_key = headers
383 .get("x-api-key")
384 .and_then(|value| value.to_str().ok());
385 surfaces::authenticate_candidates(&st.api_auth, &[bearer, api_key])
386 } else {
387 surfaces::authenticate_candidates(&st.api_auth, &[bearer])
388 };
389 if let Err(why) = auth {
390 return shape_inference_early_response(&path, authentication_error(why)).await;
391 }
392 let body_deadline = tokio::time::Instant::now() + body_read_timeout(&req);
400 let body_admission = if body_requires_admission(&req) {
401 body_admission_semaphore()
402 } else {
403 small_body_admission_semaphore()
404 };
405 let body_permit = match body_admission.try_acquire_owned() {
406 Ok(permit) => permit,
407 Err(tokio::sync::TryAcquireError::Closed) => {
408 let response = retry_contract_response(
409 error_response_coded(
410 StatusCode::SERVICE_UNAVAILABLE,
411 "request body admission is unavailable",
412 "server_error",
413 None,
414 Some("body_admission_unavailable"),
415 ),
416 Some(BODY_ADMISSION_RETRY_AFTER_S),
417 );
418 return shape_inference_early_response(&path, response).await;
419 }
420 Err(tokio::sync::TryAcquireError::NoPermits) => {
421 let response = retry_contract_response(
422 error_response_coded(
423 StatusCode::TOO_MANY_REQUESTS,
424 "request body admission is busy",
425 "rate_limit_error",
426 None,
427 Some("body_admission_busy"),
428 ),
429 Some(BODY_ADMISSION_RETRY_AFTER_S),
430 );
431 return shape_inference_early_response(&path, response).await;
432 }
433 };
434 let body_admission_guard = BodyAdmissionGuard::new(body_permit);
439 req.extensions_mut().insert(body_admission_guard.clone());
440 let body = std::mem::replace(req.body_mut(), Body::empty());
441 let mut body = Box::pin(body.into_data_stream());
442 let body_timed_out = Arc::new(std::sync::atomic::AtomicBool::new(false));
443 let body_timed_out_flag = body_timed_out.clone();
444 let guarded_body = async_stream::stream! {
445 loop {
446 let remaining = body_deadline.saturating_duration_since(tokio::time::Instant::now());
447 if remaining.is_zero() {
448 body_timed_out_flag.store(true, std::sync::atomic::Ordering::Release);
449 yield Err(std::io::Error::new(
450 std::io::ErrorKind::TimedOut,
451 "request body read deadline exceeded",
452 ));
453 break;
454 }
455 let poll = std::future::poll_fn(|cx| body.as_mut().poll_next(cx));
456 let frame = match tokio::time::timeout(BODY_IDLE_TIMEOUT.min(remaining), poll).await {
457 Ok(frame) => frame,
458 Err(_) => {
459 body_timed_out_flag.store(true, std::sync::atomic::Ordering::Release);
460 yield Err(std::io::Error::new(
461 std::io::ErrorKind::TimedOut,
462 "request body idle timeout exceeded",
463 ));
464 break;
465 }
466 };
467 match frame {
468 Some(Ok(bytes)) => yield Ok(bytes),
469 Some(Err(error)) => {
470 yield Err(std::io::Error::other(error.to_string()));
471 break;
472 }
473 None => break,
474 }
475 }
476 };
477 *req.body_mut() = Body::from_stream(guarded_body);
478 let response = next.run(req).await;
479 body_admission_guard.release();
480 if body_timed_out.load(std::sync::atomic::Ordering::Acquire) {
481 let request_id = Envelope::new(path != "/v1/completions");
482 let timeout = error_response_coded(
483 StatusCode::REQUEST_TIMEOUT,
484 "request body read timed out",
485 "invalid_request_error",
486 None,
487 Some("request_body_timeout"),
488 );
489 return if path == "/v1/messages" {
490 anthropic::with_anthropic_request_id(
491 &request_id.id,
492 anthropic::reshape_error(timeout, &request_id.id).await,
493 )
494 } else {
495 with_request_id(&request_id.id, timeout)
496 };
497 }
498 if path == "/v1/messages" && response.status() == StatusCode::PAYLOAD_TOO_LARGE {
499 let request_id = Envelope::new(true);
500 return anthropic::with_anthropic_request_id(
501 &request_id.id,
502 anthropic::reshape_error(response, &request_id.id).await,
503 );
504 }
505 response
506}
507
508#[cfg(test)]
509mod body_limit_tests {
510 use super::*;
511
512 static BODY_ADMISSION_TEST_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
513
514 fn test_app() -> Router {
518 let app = Router::new()
519 .route(
520 "/bytes",
521 post(|b: axum::body::Bytes| async move { b.len().to_string() }),
522 )
523 .route(
524 "/json",
525 post(
526 |AdmittedJson(v, _admission): AdmittedJson<serde_json::Value>| async move {
527 v["pad"].as_str().unwrap_or("").len().to_string()
528 },
529 ),
530 );
531 apply_body_limit(app)
532 }
533
534 fn streamed_body(chunks: usize) -> Body {
535 let chunk = axum::body::Bytes::from(vec![b'x'; 1024 * 1024]);
538 Body::from_stream(async_stream::stream! {
539 for _ in 0..chunks {
540 yield Ok::<_, std::io::Error>(chunk.clone());
541 }
542 })
543 }
544
545 #[tokio::test]
546 async fn bodies_past_the_old_2mib_default_are_accepted() {
547 for (path, body) in [
550 ("/bytes", Body::from(vec![b'x'; 3 * 1024 * 1024])),
551 (
552 "/json",
553 Body::from(
554 serde_json::to_vec(&json!({ "pad": "x".repeat(3 * 1024 * 1024) })).unwrap(),
555 ),
556 ),
557 ] {
558 let resp = test_app()
559 .oneshot(
560 axum::http::Request::post(path)
561 .header(CONTENT_TYPE, "application/json")
562 .body(body)
563 .unwrap(),
564 )
565 .await
566 .unwrap();
567 assert_eq!(resp.status(), StatusCode::OK, "{path}");
568 }
569 }
570
571 #[tokio::test]
572 async fn body_at_exactly_the_limit_is_accepted() {
573 let resp = test_app()
574 .oneshot(
575 axum::http::Request::post("/bytes")
576 .body(streamed_body(MAX_BODY_BYTES / (1024 * 1024)))
577 .unwrap(),
578 )
579 .await
580 .unwrap();
581 assert_eq!(resp.status(), StatusCode::OK);
582 let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
583 .await
584 .unwrap();
585 assert_eq!(body.as_ref(), MAX_BODY_BYTES.to_string().as_bytes());
586 }
587
588 #[tokio::test]
589 async fn oversize_body_is_a_clean_413_in_our_error_shape() {
590 for path in ["/bytes", "/json"] {
594 let resp = test_app()
595 .oneshot(
596 axum::http::Request::post(path)
597 .header(CONTENT_TYPE, "application/json")
598 .body(streamed_body(MAX_BODY_BYTES / (1024 * 1024) + 1))
599 .unwrap(),
600 )
601 .await
602 .unwrap();
603 assert_eq!(resp.status(), StatusCode::PAYLOAD_TOO_LARGE, "{path}");
604 assert_eq!(
605 resp.headers().get("x-should-retry").map(|v| v.as_bytes()),
606 Some(b"false".as_ref()),
607 "{path}: retrying identical bytes cannot fix a 413"
608 );
609 let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
610 .await
611 .unwrap();
612 let v: serde_json::Value = serde_json::from_slice(&body).expect("JSON error shape");
613 assert_eq!(v["error"]["type"], "invalid_request_error", "{path}");
614 assert_eq!(v["error"]["code"], "request_too_large", "{path}");
615 assert!(
616 v["error"]["message"].as_str().unwrap().contains("192 MiB"),
617 "{path}: message names the limit"
618 );
619 }
620 }
621
622 #[tokio::test]
623 async fn authenticated_body_admission_is_finite() {
624 let _test_lock = BODY_ADMISSION_TEST_LOCK.lock().await;
625 let semaphore = body_admission_semaphore();
626 let mut permits = Vec::new();
627 for _ in 0..MAX_BODY_ADMISSIONS {
628 permits.push(semaphore.clone().acquire_owned().await.unwrap());
629 }
630 assert!(
631 tokio::time::timeout(std::time::Duration::from_millis(20), semaphore.acquire())
632 .await
633 .is_err(),
634 "body parser admission must not be unbounded"
635 );
636 drop(permits);
637 assert!(semaphore.acquire().await.is_ok());
638 }
639
640 #[tokio::test]
641 async fn small_body_admission_is_finite_and_separate() {
642 let _test_lock = BODY_ADMISSION_TEST_LOCK.lock().await;
643 let large = body_admission_semaphore();
644 let small = small_body_admission_semaphore();
645 let mut small_permits = Vec::new();
646 for _ in 0..MAX_SMALL_BODY_ADMISSIONS {
647 small_permits.push(small.clone().acquire_owned().await.unwrap());
648 }
649 assert!(
650 tokio::time::timeout(std::time::Duration::from_millis(20), small.acquire())
651 .await
652 .is_err(),
653 "small body parser admission must be bounded"
654 );
655 assert!(
656 large.clone().try_acquire().is_ok(),
657 "small uploads must not consume large-upload permits"
658 );
659 drop(small_permits);
660 assert!(small.acquire().await.is_ok());
661 }
662
663 #[test]
664 fn small_declared_bodies_bypass_large_upload_admission() {
665 let request = axum::http::Request::post("/v1/chat/completions")
666 .header(CONTENT_LENGTH, "2048")
667 .body(Body::empty())
668 .unwrap();
669 assert!(!body_requires_admission(&request));
670
671 let request = axum::http::Request::post("/v1/chat/completions")
672 .header(
673 CONTENT_LENGTH,
674 (BODY_ADMISSION_BYPASS_BYTES + 1).to_string(),
675 )
676 .body(Body::empty())
677 .unwrap();
678 assert!(body_requires_admission(&request));
679
680 let request = axum::http::Request::post("/v1/chat/completions")
681 .header(CONTENT_LENGTH, "2048")
682 .header(TRANSFER_ENCODING, "chunked")
683 .body(Body::empty())
684 .unwrap();
685 assert!(body_requires_admission(&request));
686 }
687
688 #[test]
689 fn declared_body_timeout_scales_with_upload_size_and_has_a_cap() {
690 let unknown = axum::http::Request::post("/v1/chat/completions")
691 .body(Body::empty())
692 .unwrap();
693 assert_eq!(body_read_timeout(&unknown), BODY_READ_TIMEOUT);
694
695 let large = axum::http::Request::post("/v1/chat/completions")
696 .header(CONTENT_LENGTH, MAX_BODY_BYTES.to_string())
697 .body(Body::empty())
698 .unwrap();
699 assert!(body_read_timeout(&large) > BODY_READ_TIMEOUT);
700 assert_eq!(body_read_timeout(&large), BODY_READ_TIMEOUT_MAX);
701
702 let absurd = axum::http::Request::post("/v1/chat/completions")
703 .header(CONTENT_LENGTH, u64::MAX.to_string())
704 .body(Body::empty())
705 .unwrap();
706 assert_eq!(body_read_timeout(&absurd), BODY_READ_TIMEOUT_MAX);
707 }
708
709 #[tokio::test]
710 async fn early_body_refusals_keep_dialect_ids_and_retry_contracts() {
711 let too_large = shape_inference_early_response(
712 "/v1/messages",
713 error_response_coded(
714 StatusCode::PAYLOAD_TOO_LARGE,
715 "request body exceeds the 192 MiB limit",
716 "invalid_request_error",
717 None,
718 Some("request_too_large"),
719 ),
720 )
721 .await;
722 assert_eq!(too_large.status(), StatusCode::PAYLOAD_TOO_LARGE);
723 let house_id = too_large.headers()["x-request-id"].clone();
724 assert_eq!(too_large.headers()["request-id"], house_id);
725 assert_eq!(too_large.headers()["x-should-retry"], "false");
726 let body = axum::body::to_bytes(too_large.into_body(), usize::MAX)
727 .await
728 .unwrap();
729 let payload: serde_json::Value = serde_json::from_slice(&body).unwrap();
730 assert_eq!(payload["type"], "error");
731 assert_eq!(payload["request_id"], house_id.to_str().unwrap());
732
733 let busy = shape_inference_early_response(
734 "/v1/chat/completions",
735 retry_contract_response(
736 error_response_coded(
737 StatusCode::TOO_MANY_REQUESTS,
738 "request body admission is busy",
739 "rate_limit_error",
740 None,
741 Some("body_admission_busy"),
742 ),
743 Some(BODY_ADMISSION_RETRY_AFTER_S),
744 ),
745 )
746 .await;
747 assert_eq!(busy.status(), StatusCode::TOO_MANY_REQUESTS);
748 assert!(!busy.headers()["x-request-id"].is_empty());
749 assert_eq!(busy.headers()["retry-after"], "1");
750 assert_eq!(busy.headers()["retry-after-ms"], "1000");
751 assert!(busy.headers().get("x-should-retry").is_none());
752 let body = axum::body::to_bytes(busy.into_body(), usize::MAX)
753 .await
754 .unwrap();
755 let payload: serde_json::Value = serde_json::from_slice(&body).unwrap();
756 assert_eq!(payload["error"]["code"], "body_admission_busy");
757 }
758
759 #[tokio::test]
760 async fn vision_preprocess_admission_is_fail_fast_and_retryable() {
761 let semaphore = Box::leak(Box::new(tokio::sync::Semaphore::new(1)));
762 let held = semaphore.try_acquire().unwrap();
763 let busy = try_vision_preprocess_with(true, semaphore).unwrap_err();
764 assert_eq!(busy.status(), StatusCode::TOO_MANY_REQUESTS);
765 assert_eq!(busy.headers()["retry-after"], "1");
766 drop(held);
767 assert!(
768 try_vision_preprocess_with(true, semaphore)
769 .unwrap()
770 .is_some()
771 );
772 assert!(
773 try_vision_preprocess_with(false, semaphore)
774 .unwrap()
775 .is_none()
776 );
777 }
778
779 #[tokio::test]
780 async fn typed_json_retains_body_admission_until_handler_validation_releases_it() {
781 #[derive(Clone)]
782 struct Signals {
783 parsed: Arc<tokio::sync::Notify>,
784 finish: Arc<tokio::sync::Notify>,
785 semaphore: Arc<tokio::sync::Semaphore>,
786 }
787
788 let semaphore = Arc::new(tokio::sync::Semaphore::new(1));
789 let guard = BodyAdmissionGuard::new(semaphore.clone().try_acquire_owned().unwrap());
790 let signals = Signals {
791 parsed: Arc::new(tokio::sync::Notify::new()),
792 finish: Arc::new(tokio::sync::Notify::new()),
793 semaphore: semaphore.clone(),
794 };
795 let app = Router::new()
796 .route(
797 "/",
798 post(
799 |Extension(signals): Extension<Signals>,
800 AdmittedJson(_, mut admission): AdmittedJson<serde_json::Value>| async move {
801 assert_eq!(
802 signals.semaphore.available_permits(),
803 0,
804 "typed deserialization alone must not release post-parse admission"
805 );
806 admission.release();
807 assert_eq!(signals.semaphore.available_permits(), 1);
808 signals.parsed.notify_one();
809 signals.finish.notified().await;
810 "ok"
811 },
812 ),
813 )
814 .layer(Extension(signals.clone()))
815 .layer(Extension(guard));
816 let response = tokio::spawn(
817 app.oneshot(
818 axum::http::Request::post("/")
819 .header(CONTENT_TYPE, "application/json")
820 .body(Body::from(r#"{"value":1}"#))
821 .unwrap(),
822 ),
823 );
824 signals.parsed.notified().await;
825 assert_eq!(
826 semaphore.available_permits(),
827 1,
828 "validated work must release admission before generation waits"
829 );
830 signals.finish.notify_one();
831 assert_eq!(response.await.unwrap().unwrap().status(), StatusCode::OK);
832 }
833
834 #[tokio::test]
835 async fn transport_closes_stalled_headers_and_caps_connections() {
836 use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
837
838 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
839 let address = listener.local_addr().unwrap();
840 let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
841 let server = tokio::spawn(serve_bounded_http_with_limits(
842 listener,
843 Router::new().route("/", get(|| async { "ok" })).route(
844 "/slow",
845 get(|| async {
846 tokio::time::sleep(std::time::Duration::from_millis(140)).await;
847 "slow-ok"
848 }),
849 ),
850 async move {
851 let _ = shutdown_rx.await;
852 },
853 std::time::Duration::from_millis(30),
854 1,
855 std::time::Duration::from_millis(80),
856 ));
857
858 let mut stalled = tokio::net::TcpStream::connect(address).await.unwrap();
859 stalled.write_all(b"GET / HT").await.unwrap();
860 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
861 let mut excess = tokio::net::TcpStream::connect(address).await.unwrap();
862 let mut bytes = Vec::new();
863 tokio::time::timeout(
864 std::time::Duration::from_millis(250),
865 excess.read_to_end(&mut bytes),
866 )
867 .await
868 .expect("connection beyond the cap must be closed promptly")
869 .unwrap();
870
871 bytes.clear();
872 tokio::time::timeout(
873 std::time::Duration::from_millis(500),
874 stalled.read_to_end(&mut bytes),
875 )
876 .await
877 .expect("stalled request headers must hit the configured deadline")
878 .unwrap();
879
880 let mut idle = tokio::net::TcpStream::connect(address).await.unwrap();
881 idle.write_all(b"GET / HTTP/1.1\r\nHost: local\r\n\r\n")
882 .await
883 .unwrap();
884 bytes.clear();
885 tokio::time::timeout(
886 std::time::Duration::from_millis(500),
887 idle.read_to_end(&mut bytes),
888 )
889 .await
890 .expect("an idle keep-alive connection must hit the maximum lifetime")
891 .unwrap();
892 assert!(String::from_utf8_lossy(&bytes).contains("200 OK"));
893
894 let mut active = tokio::net::TcpStream::connect(address).await.unwrap();
895 active
896 .write_all(b"GET /slow HTTP/1.1\r\nHost: local\r\n\r\n")
897 .await
898 .unwrap();
899 bytes.clear();
900 tokio::time::timeout(
901 std::time::Duration::from_millis(500),
902 active.read_to_end(&mut bytes),
903 )
904 .await
905 .expect("an active response must finish across the connection age boundary")
906 .unwrap();
907 let active_response = String::from_utf8_lossy(&bytes);
908 assert!(active_response.contains("200 OK"), "{active_response}");
909 assert!(active_response.contains("slow-ok"), "{active_response}");
910
911 let h2_stream = tokio::net::TcpStream::connect(address).await.unwrap();
915 let (mut h2_client, h2_connection) = h2::client::handshake(h2_stream).await.unwrap();
916 let h2_driver = tokio::spawn(h2_connection);
917 let request = axum::http::Request::builder()
918 .uri(format!("http://{address}/"))
919 .body(())
920 .unwrap();
921 let (response, _) = h2_client.send_request(request, true).unwrap();
922 let response = tokio::time::timeout(std::time::Duration::from_millis(500), response)
923 .await
924 .expect("HTTP/2 handshake and response must complete")
925 .expect("HTTP/2 connection must stay alive through the response");
926 assert_eq!(response.status(), StatusCode::OK);
927 drop(h2_client);
928 h2_driver.abort();
929 let _ = h2_driver.await;
930
931 let _ = shutdown_tx.send(());
932 server.await.unwrap().unwrap();
933 }
934}
935
936#[derive(Clone, Default)]
937struct TtftRequestTrace(Option<Arc<ttft::Trace>>);
938
939fn is_sse_data_frame(bytes: &[u8]) -> bool {
940 bytes
941 .windows(b"data:".len())
942 .any(|window| window == b"data:")
943}
944
945async fn ttft_request_start(mut req: AxumRequest, next: Next) -> Response {
946 let trace = ttft::start(req.uri().path());
947 req.extensions_mut().insert(TtftRequestTrace(trace.clone()));
948 let response = next.run(req).await;
949 let Some(trace) = trace else {
950 return response;
951 };
952 let is_sse = response
953 .headers()
954 .get(CONTENT_TYPE)
955 .and_then(|value| value.to_str().ok())
956 .is_some_and(|value| value.starts_with("text/event-stream"));
957 if !is_sse {
958 return response;
959 }
960
961 let (parts, body) = response.into_parts();
964 let mut body = Box::pin(body.into_data_stream());
965 let stream = async_stream::stream! {
966 while let Some(frame) =
967 std::future::poll_fn(|cx| body.as_mut().poll_next(cx)).await
968 {
969 if frame
970 .as_ref()
971 .is_ok_and(|bytes| is_sse_data_frame(bytes))
972 {
973 trace.mark_first_sse_byte();
974 }
975 yield frame;
976 }
977 };
978 Response::from_parts(parts, Body::from_stream(stream))
979}
980
981const OPENROUTER_SCHEMA_VERSION: &str = "2.4";
982const JSON_SAFE_INTEGER_MAX: u64 = 9_007_199_254_740_991;
983
984#[derive(Debug, Clone, Default, Deserialize)]
985#[serde(deny_unknown_fields)]
986struct OpenRouterMetadataFile {
987 #[serde(default)]
988 models: HashMap<String, OpenRouterModelMetadata>,
989 #[serde(default)]
992 planned_models: HashMap<String, OpenRouterModelMetadata>,
993 #[serde(default)]
996 provider: Option<ProviderMetadata>,
997}
998
999#[derive(Debug, Clone, Deserialize)]
1003#[serde(deny_unknown_fields)]
1004struct ProviderMetadata {
1005 id: String,
1006 #[serde(default)]
1007 status_url: Option<String>,
1008 #[serde(default)]
1009 support_contact: Option<String>,
1010 #[serde(default)]
1011 incident_contact: Option<String>,
1012 #[serde(default)]
1013 regions: Vec<String>,
1014}
1015
1016#[derive(Debug, Clone, Default, Deserialize)]
1018#[serde(deny_unknown_fields)]
1019struct LifecycleMetadata {
1020 #[serde(default)]
1021 status: Option<String>,
1022 #[serde(default)]
1023 deprecation_at: Option<String>,
1024 #[serde(default)]
1025 retirement_at: Option<String>,
1026 #[serde(default)]
1027 replacement_model_id: Option<String>,
1028}
1029
1030#[derive(Debug, Clone, Default, Deserialize)]
1032#[serde(deny_unknown_fields)]
1033struct ReliabilityMetadata {
1034 #[serde(default)]
1035 first_token_timeout_seconds: Option<u64>,
1036 #[serde(default)]
1037 completion_timeout_seconds: Option<u64>,
1038 #[serde(default)]
1039 stream_idle_timeout_seconds: Option<u64>,
1040 #[serde(default)]
1041 capacity_scope: Option<String>,
1042}
1043
1044#[derive(Debug, Clone, Default, Deserialize)]
1045#[serde(deny_unknown_fields)]
1046struct OpenRouterModelMetadata {
1047 #[serde(default)]
1049 owned_by: Option<String>,
1050 #[serde(default)]
1051 lifecycle: Option<LifecycleMetadata>,
1052 #[serde(default)]
1053 reliability: Option<ReliabilityMetadata>,
1054 #[serde(default)]
1055 hugging_face_id: Option<String>,
1056 #[serde(default)]
1057 created: Option<u64>,
1058 #[serde(default)]
1059 quantization: Option<String>,
1060 #[serde(default)]
1061 description: Option<String>,
1062 #[serde(default)]
1063 max_prompt_length: Option<u64>,
1064 #[serde(default)]
1065 max_output_length: Option<u64>,
1066 #[serde(default)]
1069 default_output_length: Option<u64>,
1070 #[serde(default)]
1071 pricing: OpenRouterPricing,
1072 #[serde(default)]
1073 capacity: OpenRouterCapacity,
1074 #[serde(default)]
1075 is_ready: Option<bool>,
1076 #[serde(default)]
1077 is_free: Option<bool>,
1078 #[serde(default)]
1079 discount_to_user: Option<f64>,
1080 #[serde(default)]
1081 openrouter_slug: Option<String>,
1082 #[serde(default)]
1083 datacenters: Vec<OpenRouterDatacenter>,
1084 #[serde(default)]
1088 input_modalities: Vec<String>,
1089 #[serde(default)]
1104 surface: Option<String>,
1105 #[serde(default)]
1106 zdr: Option<bool>,
1107 #[serde(default)]
1108 hipaa: Option<bool>,
1109 #[serde(default)]
1119 default_reasoning_effort: Option<String>,
1120 #[serde(default)]
1138 default_temperature: Option<f32>,
1139 #[serde(default)]
1140 default_top_p: Option<f32>,
1141 #[serde(default)]
1143 default_top_k: Option<usize>,
1144 #[serde(default)]
1145 default_min_p: Option<f32>,
1146 #[serde(default)]
1147 default_presence_penalty: Option<f32>,
1148 #[serde(default)]
1149 default_frequency_penalty: Option<f32>,
1150 #[serde(default)]
1152 default_repetition_penalty: Option<f32>,
1153 #[serde(default)]
1174 non_thinking_sampling: Option<SamplingArmMetadata>,
1175}
1176
1177#[derive(Debug, Clone, Default, Deserialize)]
1183#[serde(deny_unknown_fields)]
1184struct SamplingArmMetadata {
1185 #[serde(default)]
1186 temperature: Option<f32>,
1187 #[serde(default)]
1188 top_p: Option<f32>,
1189 #[serde(default)]
1190 top_k: Option<usize>,
1191 #[serde(default)]
1192 min_p: Option<f32>,
1193 #[serde(default)]
1194 presence_penalty: Option<f32>,
1195 #[serde(default)]
1196 frequency_penalty: Option<f32>,
1197 #[serde(default)]
1198 repetition_penalty: Option<f32>,
1199}
1200
1201impl SamplingArmMetadata {
1202 fn is_empty(&self) -> bool {
1203 self.temperature.is_none()
1204 && self.top_p.is_none()
1205 && self.top_k.is_none()
1206 && self.min_p.is_none()
1207 && self.presence_penalty.is_none()
1208 && self.frequency_penalty.is_none()
1209 && self.repetition_penalty.is_none()
1210 }
1211}
1212
1213#[derive(Debug, Clone, Default, Deserialize)]
1214#[serde(deny_unknown_fields)]
1215struct OpenRouterPricing {
1216 #[serde(default)]
1217 prompt: Option<String>,
1218 #[serde(default)]
1219 cached_prompt: Option<String>,
1220 #[serde(default)]
1221 cache_write: Option<String>,
1222 #[serde(default)]
1223 completion: Option<String>,
1224 #[serde(default)]
1225 internal_reasoning: Option<String>,
1226 #[serde(default)]
1227 request: Option<String>,
1228}
1229
1230#[derive(Debug, Clone, Default, Deserialize)]
1231#[serde(deny_unknown_fields)]
1232struct OpenRouterCapacity {
1233 #[serde(default)]
1234 prompt_tpm: Option<u64>,
1235 #[serde(default)]
1236 cached_prompt_tpm: Option<u64>,
1237 #[serde(default)]
1238 completion_tpm: Option<u64>,
1239 #[serde(default)]
1240 request_rpm: Option<u64>,
1241 #[serde(default)]
1242 concurrency: Option<u64>,
1243}
1244
1245#[derive(Debug, Clone, Deserialize, Serialize)]
1246#[serde(deny_unknown_fields)]
1247struct OpenRouterDatacenter {
1248 country_code: String,
1249 #[serde(default, skip_serializing_if = "Option::is_none")]
1250 region: Option<String>,
1251}
1252
1253impl OpenRouterMetadataFile {
1254 fn parse(
1255 text: &str,
1256 ) -> Result<
1257 (
1258 HashMap<String, OpenRouterModelMetadata>,
1259 Option<ProviderMetadata>,
1260 ),
1261 String,
1262 > {
1263 let file: Self =
1264 toml::from_str(text).map_err(|e| format!("models metadata TOML parse: {e}"))?;
1265 for (alias, metadata) in &file.models {
1266 validate_openrouter_metadata(alias, metadata)?;
1267 }
1268 for (alias, metadata) in &file.planned_models {
1269 validate_openrouter_metadata(alias, metadata)?;
1270 if file.models.contains_key(alias) {
1271 return Err(format!(
1272 "model alias {alias:?} appears in both models and planned_models"
1273 ));
1274 }
1275 }
1276 if let Some(provider) = &file.provider {
1277 if provider.id.is_empty() {
1278 return Err("provider.id must be a non-empty slug".into());
1279 }
1280 for (field, value) in [
1282 ("provider.support_contact", &provider.support_contact),
1283 ("provider.incident_contact", &provider.incident_contact),
1284 ] {
1285 if let Some(value) = value
1286 && !value.contains(':')
1287 {
1288 return Err(format!(
1289 "{field} must be a URI (mailto:… or https://…), got {value:?}"
1290 ));
1291 }
1292 }
1293 }
1294 Ok((file.models, file.provider))
1295 }
1296
1297 #[cfg(test)]
1298 fn from_toml(text: &str) -> Result<HashMap<String, OpenRouterModelMetadata>, String> {
1299 Self::parse(text).map(|(models, _)| models)
1300 }
1301}
1302
1303fn per_million_price(per_token: &str) -> Option<String> {
1307 if !valid_price_string(per_token) {
1308 return None;
1309 }
1310 let (whole, frac) = match per_token.split_once('.') {
1311 Some((whole, frac)) => (whole, frac),
1312 None => (per_token, ""),
1313 };
1314 let mut digits = format!("{whole}{frac}");
1315 let point = whole.len() + 6;
1316 while digits.len() < point {
1317 digits.push('0');
1318 }
1319 let (int_part, frac_part) = digits.split_at(point);
1320 let int_part = int_part.trim_start_matches('0');
1321 let int_part = if int_part.is_empty() { "0" } else { int_part };
1322 let mut frac_out = frac_part.trim_end_matches('0').to_string();
1323 while frac_out.len() < 2 {
1324 frac_out.push('0');
1325 }
1326 Some(format!("{int_part}.{frac_out}"))
1327}
1328
1329fn valid_price_string(value: &str) -> bool {
1330 let mut parts = value.split('.');
1331 let whole = parts.next().unwrap_or_default();
1332 let fraction = parts.next();
1333 !whole.is_empty()
1334 && whole.bytes().all(|b| b.is_ascii_digit())
1335 && fraction.is_none_or(|v| !v.is_empty() && v.bytes().all(|b| b.is_ascii_digit()))
1336 && parts.next().is_none()
1337}
1338
1339fn validate_openrouter_metadata(
1340 alias: &str,
1341 metadata: &OpenRouterModelMetadata,
1342) -> Result<(), String> {
1343 if alias.is_empty() {
1344 return Err("models metadata contains an empty model alias".into());
1345 }
1346 if let Some(effort) = metadata.default_reasoning_effort.as_deref()
1349 && !matches!(effort, "none" | "minimal" | "low" | "medium" | "high")
1350 {
1351 return Err(format!(
1352 "model {alias:?}: default_reasoning_effort {effort:?} is not a \
1353 reasoning_effort level (none|minimal|low|medium|high)"
1354 ));
1355 }
1356 validate_sampling_defaults(alias, metadata)?;
1357 for m in &metadata.input_modalities {
1358 if m != "image" && m != "video" {
1359 return Err(format!(
1360 "model {alias:?}: input_modalities entry {m:?} not served (image/video)"
1361 ));
1362 }
1363 }
1364 if let Some(sfc) = metadata.surface.as_deref()
1365 && !matches!(sfc, "chat" | "embedding" | "rerank")
1366 {
1367 return Err(format!(
1368 "model {alias:?}: surface {sfc:?} is not a served surface (chat|embedding|rerank)"
1369 ));
1370 }
1371 if let Some(q) = metadata.quantization.as_deref()
1372 && !matches!(
1373 q,
1374 "int4"
1375 | "int8"
1376 | "fp4"
1377 | "mxfp4"
1378 | "nvfp4"
1379 | "fp6"
1380 | "fp8"
1381 | "mxfp8"
1382 | "fp16"
1383 | "bf16"
1384 | "fp32"
1385 )
1386 {
1387 return Err(format!(
1388 "model {alias:?}: quantization {q:?} is not in the OpenRouter schema 2.4 enum"
1389 ));
1390 }
1391 for (field, value) in [
1392 ("pricing.prompt", metadata.pricing.prompt.as_deref()),
1393 (
1394 "pricing.cached_prompt",
1395 metadata.pricing.cached_prompt.as_deref(),
1396 ),
1397 (
1398 "pricing.cache_write",
1399 metadata.pricing.cache_write.as_deref(),
1400 ),
1401 ("pricing.completion", metadata.pricing.completion.as_deref()),
1402 (
1403 "pricing.internal_reasoning",
1404 metadata.pricing.internal_reasoning.as_deref(),
1405 ),
1406 ("pricing.request", metadata.pricing.request.as_deref()),
1407 ] {
1408 if let Some(value) = value
1409 && !valid_price_string(value)
1410 {
1411 return Err(format!(
1412 "model {alias:?}: {field} must be a non-negative per-unit USD decimal string"
1413 ));
1414 }
1415 }
1416 for (field, value) in [
1417 ("created", metadata.created),
1418 ("max_prompt_length", metadata.max_prompt_length),
1419 ("max_output_length", metadata.max_output_length),
1420 ("default_output_length", metadata.default_output_length),
1421 ("capacity.prompt_tpm", metadata.capacity.prompt_tpm),
1422 (
1423 "capacity.cached_prompt_tpm",
1424 metadata.capacity.cached_prompt_tpm,
1425 ),
1426 ("capacity.completion_tpm", metadata.capacity.completion_tpm),
1427 ("capacity.request_rpm", metadata.capacity.request_rpm),
1428 ("capacity.concurrency", metadata.capacity.concurrency),
1429 ] {
1430 if let Some(value) = value
1431 && value > JSON_SAFE_INTEGER_MAX
1432 {
1433 return Err(format!(
1434 "model {alias:?}: {field} exceeds OpenRouter's JSON safe-integer maximum"
1435 ));
1436 }
1437 }
1438 for (field, value) in [
1439 ("max_prompt_length", metadata.max_prompt_length),
1440 ("max_output_length", metadata.max_output_length),
1441 ("default_output_length", metadata.default_output_length),
1442 ("capacity.prompt_tpm", metadata.capacity.prompt_tpm),
1443 (
1444 "capacity.cached_prompt_tpm",
1445 metadata.capacity.cached_prompt_tpm,
1446 ),
1447 ("capacity.completion_tpm", metadata.capacity.completion_tpm),
1448 ("capacity.request_rpm", metadata.capacity.request_rpm),
1449 ("capacity.concurrency", metadata.capacity.concurrency),
1450 ] {
1451 if value == Some(0) {
1452 return Err(format!(
1453 "model {alias:?}: {field} must be greater than zero when declared"
1454 ));
1455 }
1456 }
1457 if let (Some(default), Some(maximum)) =
1458 (metadata.default_output_length, metadata.max_output_length)
1459 && default > maximum
1460 {
1461 return Err(format!(
1462 "model {alias:?}: default_output_length {default} exceeds max_output_length {maximum}"
1463 ));
1464 }
1465 if metadata.default_output_length.is_some() && metadata.max_output_length.is_none() {
1466 return Err(format!(
1467 "model {alias:?}: default_output_length requires max_output_length"
1468 ));
1469 }
1470 if let Some(discount) = metadata.discount_to_user
1471 && (!discount.is_finite() || discount >= 1.0)
1472 {
1473 return Err(format!(
1474 "model {alias:?}: discount_to_user must be finite and less than 1"
1475 ));
1476 }
1477 if metadata
1478 .openrouter_slug
1479 .as_deref()
1480 .is_some_and(str::is_empty)
1481 {
1482 return Err(format!(
1483 "model {alias:?}: openrouter_slug must not be empty when declared"
1484 ));
1485 }
1486 for dc in &metadata.datacenters {
1487 if dc.country_code.len() != 2 || !dc.country_code.bytes().all(|b| b.is_ascii_uppercase()) {
1488 return Err(format!(
1489 "model {alias:?}: datacenter country_code {:?} must be two uppercase ASCII letters",
1490 dc.country_code
1491 ));
1492 }
1493 }
1494 Ok(())
1495}
1496
1497fn validate_sampling_defaults(
1512 alias: &str,
1513 metadata: &OpenRouterModelMetadata,
1514) -> Result<(), String> {
1515 validate_sampling_arm(
1516 alias,
1517 &[
1518 "default_temperature",
1519 "default_top_p",
1520 "default_min_p",
1521 "default_presence_penalty",
1522 "default_frequency_penalty",
1523 "default_repetition_penalty",
1524 ],
1525 metadata.default_temperature,
1526 metadata.default_top_p,
1527 metadata.default_min_p,
1528 metadata.default_presence_penalty,
1529 metadata.default_frequency_penalty,
1530 metadata.default_repetition_penalty,
1531 )?;
1532 if let Some(arm) = &metadata.non_thinking_sampling {
1533 if arm.is_empty() {
1537 return Err(format!(
1538 "model {alias:?}: non_thinking_sampling declares no fields — declare at \
1539 least one vendor recommendation or delete the table"
1540 ));
1541 }
1542 validate_sampling_arm(
1543 alias,
1544 &[
1545 "non_thinking_sampling.temperature",
1546 "non_thinking_sampling.top_p",
1547 "non_thinking_sampling.min_p",
1548 "non_thinking_sampling.presence_penalty",
1549 "non_thinking_sampling.frequency_penalty",
1550 "non_thinking_sampling.repetition_penalty",
1551 ],
1552 arm.temperature,
1553 arm.top_p,
1554 arm.min_p,
1555 arm.presence_penalty,
1556 arm.frequency_penalty,
1557 arm.repetition_penalty,
1558 )?;
1559 }
1560 Ok(())
1561}
1562
1563#[allow(clippy::too_many_arguments)]
1569fn validate_sampling_arm(
1570 alias: &str,
1571 keys: &[&str; 6],
1572 temperature: Option<f32>,
1573 top_p: Option<f32>,
1574 min_p: Option<f32>,
1575 presence_penalty: Option<f32>,
1576 frequency_penalty: Option<f32>,
1577 repetition_penalty: Option<f32>,
1578) -> Result<(), String> {
1579 if let Some(t) = temperature
1580 && (!t.is_finite() || t <= 0.0 || t > 2.0)
1581 {
1582 return Err(format!(
1583 "model {alias:?}: {} {t} must be finite and in (0, 2]. \
1584 A zero DEFAULT would make greedy decoding the deployment-wide behavior for \
1585 every request that omits temperature (owner ruling 2026-08-19: we serve the \
1586 vendor recommendation, not greedy); clients reach greedy by sending an \
1587 explicit temperature 0.",
1588 keys[0]
1589 ));
1590 }
1591 if let Some(p) = top_p
1592 && (!p.is_finite() || p <= 0.0 || p > 1.0)
1593 {
1594 return Err(format!(
1595 "model {alias:?}: {} {p} must be finite and in (0, 1] (1.0 = disabled)",
1596 keys[1]
1597 ));
1598 }
1599 if let Some(m) = min_p
1600 && (!m.is_finite() || !(0.0..1.0).contains(&m))
1601 {
1602 return Err(format!(
1603 "model {alias:?}: {} {m} must be finite and in [0, 1) (0.0 = disabled)",
1604 keys[2]
1605 ));
1606 }
1607 for (field, value) in [(keys[3], presence_penalty), (keys[4], frequency_penalty)] {
1608 if let Some(v) = value
1609 && (!v.is_finite() || !(-2.0..=2.0).contains(&v))
1610 {
1611 return Err(format!(
1612 "model {alias:?}: {field} {v} must be finite and in [-2, 2]"
1613 ));
1614 }
1615 }
1616 if let Some(r) = repetition_penalty
1617 && (!r.is_finite() || r <= 0.0)
1618 {
1619 return Err(format!(
1620 "model {alias:?}: {} {r} must be finite and \
1621 greater than zero (1.0 = off)",
1622 keys[5]
1623 ));
1624 }
1625 Ok(())
1626}
1627
1628fn load_openrouter_metadata(
1629 models: &[(String, String, Option<String>)],
1630) -> Result<
1631 (
1632 HashMap<String, OpenRouterModelMetadata>,
1633 Option<ProviderMetadata>,
1634 ),
1635 String,
1636> {
1637 let path = match std::env::var("MEMRA_MODEL_METADATA") {
1638 Ok(path) => path,
1639 Err(_) => return Ok((HashMap::new(), None)),
1640 };
1641 let p = std::path::Path::new(&path);
1642 if !p.is_file() {
1643 return Err(format!(
1644 "MEMRA_MODEL_METADATA={path:?} is not an existing TOML file"
1645 ));
1646 }
1647 let text =
1648 std::fs::read_to_string(p).map_err(|e| format!("MEMRA_MODEL_METADATA {path:?}: {e}"))?;
1649 let (metadata, provider) = OpenRouterMetadataFile::parse(&text)
1650 .map_err(|e| format!("MEMRA_MODEL_METADATA {path:?}: {e}"))?;
1651 for alias in metadata.keys() {
1652 if !models.iter().any(|(name, _, _)| name == alias) {
1653 return Err(format!(
1654 "MEMRA_MODEL_METADATA {path:?}: model alias {alias:?} is not present in MEMRA_MODELS"
1655 ));
1656 }
1657 }
1658 eprintln!(
1659 "[server] OpenRouter metadata loaded: {} model(s) from {path}",
1660 metadata.len()
1661 );
1662 Ok((metadata, provider))
1663}
1664
1665#[derive(Clone)]
1666struct AppState {
1667 cmd_tx: Sender<Cmd>,
1668 models: Arc<Vec<String>>,
1669 caps: Arc<HashMap<String, ModelCaps>>,
1670 openrouter_metadata: Arc<HashMap<String, OpenRouterModelMetadata>>,
1671 provider_metadata: Arc<Option<ProviderMetadata>>,
1673 metering: Option<Arc<dyn metering::Metering>>,
1679 budget_tokenizers: Option<Arc<HashMap<String, Arc<Tokenizer>>>>,
1682 api_auth: ApiAuth,
1685 metrics_auth: MetricsAuth,
1687 metrics: SharedMetrics,
1688 inflight: InflightCounts,
1692 tenant_inflight: TenantGauge,
1695 health: health::SharedHealth,
1699 bg: Option<(Arc<darklane::BgJobState>, &'static str)>,
1703}
1704
1705impl AppState {
1706 fn sampling_defaults(&self, model: &str) -> ModelSamplingDefaults {
1720 ModelSamplingDefaults::resolve(self.openrouter_metadata.get(model), self.caps.get(model))
1721 }
1722}
1723
1724#[derive(Clone, Default)]
1725struct ApiAuth {
1726 keyring: Option<&'static auth::KeyStore>,
1727 single_key: Option<Arc<str>>,
1728}
1729
1730impl ApiAuth {
1731 fn from_env() -> Result<ApiAuth, String> {
1732 let single_key = match std::env::var("MEMRA_API_KEY") {
1733 Ok(key) if key.is_empty() => return Err("MEMRA_API_KEY must not be empty".into()),
1734 Ok(key) => Some(Arc::from(key)),
1735 Err(std::env::VarError::NotPresent) => None,
1736 Err(std::env::VarError::NotUnicode(_)) => {
1737 return Err("MEMRA_API_KEY must be valid UTF-8".into());
1738 }
1739 };
1740 Ok(ApiAuth {
1741 keyring: auth::global(),
1742 single_key,
1743 })
1744 }
1745
1746 fn configured(&self) -> bool {
1747 self.keyring.is_some() || self.single_key.is_some()
1748 }
1749}
1750
1751#[derive(Clone, Default)]
1752struct MetricsAuth {
1753 required: bool,
1754 token: Option<Arc<str>>,
1755}
1756
1757impl MetricsAuth {
1758 fn new(bind_loopback: bool, api_auth_configured: bool, token: Option<String>) -> MetricsAuth {
1759 let token = token.map(Arc::from);
1760 MetricsAuth {
1761 required: !bind_loopback || api_auth_configured || token.is_some(),
1762 token,
1763 }
1764 }
1765}
1766
1767fn resolve_bind_addr(addr: &str) -> Result<(SocketAddr, bool), String> {
1768 let mut resolved = addr
1769 .to_socket_addrs()
1770 .map_err(|e| format!("MEMRA_ADDR={addr:?} cannot be resolved: {e}"))?;
1771 let first = resolved
1772 .next()
1773 .ok_or_else(|| format!("MEMRA_ADDR={addr:?} resolved to no socket addresses"))?;
1774 let mut loopback = first.ip().to_canonical().is_loopback();
1775 for socket in resolved {
1776 loopback &= socket.ip().to_canonical().is_loopback();
1777 }
1778 Ok((first, loopback))
1779}
1780
1781fn bind_is_loopback(addr: &str) -> Result<bool, String> {
1782 resolve_bind_addr(addr).map(|(_, loopback)| loopback)
1783}
1784
1785fn validate_bind_security(
1786 addr: &str,
1787 api_auth_configured: bool,
1788 allow_open_bind: bool,
1789) -> Result<bool, String> {
1790 let loopback = bind_is_loopback(addr)?;
1791 if !loopback && !api_auth_configured && !allow_open_bind {
1792 return Err(format!(
1793 "refusing unauthenticated non-loopback bind {addr:?}; configure MEMRA_API_KEY or \
1794 MEMRA_API_KEYS, or set MEMRA_ALLOW_OPEN_BIND=1 for an explicit development override"
1795 ));
1796 }
1797 Ok(loopback)
1798}
1799
1800type InflightCounts = Arc<[std::sync::atomic::AtomicUsize; 3]>;
1818
1819type TenantGauge = Arc<std::sync::Mutex<HashMap<String, usize>>>;
1822
1823struct InflightGuard {
1827 counts: InflightCounts,
1828 idx: usize,
1829 tenants: TenantGauge,
1830 tenant: String,
1831}
1832
1833impl InflightGuard {
1834 fn try_acquire(
1838 counts: InflightCounts,
1839 lane: lanes::Lane,
1840 tenants: TenantGauge,
1841 tenant: &str,
1842 tenant_cap: Option<usize>,
1843 ) -> Result<(Self, usize, usize), usize> {
1844 let idx = lane.idx();
1845 let nt = {
1846 let mut m = tenants.lock().unwrap();
1847 let e = m.entry(tenant.to_string()).or_insert(0);
1848 if tenant_cap.is_some_and(|cap| *e >= cap) {
1849 return Err(*e);
1850 }
1851 *e += 1;
1852 *e
1853 };
1854 let n = counts[idx].fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
1855 Ok((
1856 InflightGuard {
1857 counts,
1858 idx,
1859 tenants,
1860 tenant: tenant.to_string(),
1861 },
1862 n,
1863 nt,
1864 ))
1865 }
1866}
1867
1868impl Drop for InflightGuard {
1869 fn drop(&mut self) {
1870 self.counts[self.idx].fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
1871 let mut m = self.tenants.lock().unwrap();
1872 if let Some(e) = m.get_mut(&self.tenant) {
1873 *e -= 1;
1874 if *e == 0 {
1875 m.remove(&self.tenant);
1876 }
1877 }
1878 }
1879}
1880
1881fn lane_cap(lane: lanes::Lane) -> usize {
1885 static CAPS: std::sync::OnceLock<[usize; 3]> = std::sync::OnceLock::new();
1886 CAPS.get_or_init(|| {
1887 let batching = std::env::var("MEMRA_SERVE_BATCH")
1888 .map(|v| v != "0")
1889 .unwrap_or(true);
1890 let interactive = if batching {
1891 std::env::var("MEMRA_MAX_SESSIONS")
1892 .ok()
1893 .and_then(|v| v.parse().ok())
1894 .unwrap_or(64)
1895 } else {
1896 worker::MAX_ACTIVE
1897 };
1898 let p = lanes::LanePolicy::from_env();
1899 [interactive, p.max_sessions[1], p.max_sessions[2]]
1900 })[lane.idx()]
1901}
1902
1903fn reset_estimate_s(m: &worker::Metrics) -> u64 {
1906 if m.completed > 0 && m.step_p50_ms > 0.0 {
1907 let mean_toks = m.tokens_out as f64 / m.completed as f64;
1908 return ((mean_toks * m.step_p50_ms as f64 / 1000.0).ceil() as u64).clamp(1, 600);
1909 }
1910 static D: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
1911 *D.get_or_init(|| {
1912 std::env::var("MEMRA_RL_RESET_S")
1913 .ok()
1914 .and_then(|v| v.parse().ok())
1915 .unwrap_or(2)
1916 })
1917}
1918
1919pub(crate) const TIMEOUT_MS_MIN: u64 = 1_000;
1936pub(crate) const TIMEOUT_MS_MAX: u64 = 90_000;
1937pub(crate) const TIMEOUT_MS_DEFAULT: u64 = 90_000;
1938
1939pub(crate) fn timeout_ms_max() -> u64 {
1949 static V: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
1950 *V.get_or_init(|| {
1951 std::env::var("MEMRA_TIMEOUT_MS_MAX")
1952 .ok()
1953 .and_then(|s| s.parse::<u64>().ok())
1954 .filter(|&ms| ms >= TIMEOUT_MS_MIN)
1955 .unwrap_or(TIMEOUT_MS_MAX)
1956 })
1957}
1958
1959pub(crate) fn parse_timeout_ms(v: Option<&serde_json::Value>) -> Result<u64, String> {
1963 let max = timeout_ms_max();
1964 let Some(v) = v.filter(|v| !v.is_null()) else {
1965 return Ok(max);
1967 };
1968 let Some(ms) = v.as_u64() else {
1969 return Err(format!(
1970 "timeout_ms must be an integer number of milliseconds in \
1971 {TIMEOUT_MS_MIN}..={max}, got {v}; for work longer than \
1972 {max} ms use \"stream\": true — the deadline then bounds only the \
1973 time to first token and the stream may run as long as it needs"
1974 ));
1975 };
1976 if !(TIMEOUT_MS_MIN..=max).contains(&ms) {
1977 return Err(format!(
1978 "timeout_ms {ms} is outside the accepted range \
1979 {TIMEOUT_MS_MIN}..={max} (milliseconds). {max} is a \
1980 platform ceiling, not a preference: the fronting proxy fails a non-streaming \
1981 response whose headers take ~100 s (HTTP 524), so promising more would be a \
1982 lie. For work longer than {max} ms use \"stream\": true — the \
1983 deadline then bounds only the time to first token and the stream may run as \
1984 long as it needs"
1985 ));
1986 }
1987 Ok(ms)
1988}
1989
1990#[derive(Clone, Copy)]
1993pub(crate) struct RequestDeadline {
1994 pub(crate) at: tokio::time::Instant,
1995 pub(crate) ms: u64,
1996}
1997
1998impl RequestDeadline {
1999 pub(crate) fn starting_now(ms: u64) -> Self {
2000 Self {
2001 at: tokio::time::Instant::now() + std::time::Duration::from_millis(ms),
2002 ms,
2003 }
2004 }
2005
2006 pub(crate) fn remaining(&self) -> std::time::Duration {
2007 self.at
2008 .saturating_duration_since(tokio::time::Instant::now())
2009 }
2010}
2011
2012pub(crate) fn deadline_exceeded_response(ms: u64, stream: bool) -> Response {
2018 let what = if stream {
2019 "the first token was produced"
2020 } else {
2021 "the response completed"
2022 };
2023 let msg = format!(
2024 "deadline of {ms} ms (timeout_ms; default {TIMEOUT_MS_DEFAULT}) elapsed before \
2025 {what}; generation was cancelled and this request is not billed"
2026 );
2027 error_response_coded(
2028 StatusCode::REQUEST_TIMEOUT,
2029 &msg,
2030 "timeout",
2031 Some("timeout_ms"),
2032 Some("deadline_exceeded"),
2033 )
2034}
2035
2036pub(crate) const PREFILL_FLOOR_TOK_S: u64 = 2_000;
2076
2077pub(crate) const DECODE_FLOOR_TOK_S: u64 = 60;
2082
2083pub(crate) const DEADLINE_INFEASIBLE_MARGIN_PCT: u64 = 150;
2087
2088fn env_flag_on(name: &'static str, default_on: bool) -> bool {
2094 match std::env::var(name) {
2095 Ok(v) => !matches!(
2096 v.trim().to_ascii_lowercase().as_str(),
2097 "0" | "off" | "false"
2098 ),
2099 Err(_) => default_on,
2100 }
2101}
2102
2103pub(crate) fn env_u64(name: &'static str, default: u64) -> u64 {
2106 std::env::var(name)
2107 .ok()
2108 .and_then(|v| v.parse::<u64>().ok())
2109 .filter(|v| *v > 0)
2110 .unwrap_or(default)
2111}
2112
2113const CHARS_PER_TOKEN_FLOOR: usize = 6;
2124
2125pub(crate) fn prompt_tokens_estimate(
2126 request: &worker::Request,
2127 tokenizer: Option<&Tokenizer>,
2128) -> u64 {
2129 if !request.prompt_ids.is_empty() {
2130 return request.prompt_ids.len() as u64;
2131 }
2132 let mut text = String::new();
2133 text.push_str(&request.prompt_text);
2134 for turn in &request.chat_turns {
2135 text.push_str(&turn.content);
2136 }
2137 for tool in &request.tools_json {
2138 text.push_str(tool);
2139 }
2140 if let Some(tokenizer) = tokenizer {
2141 return tokenizer.encode(text.as_str(), false).len() as u64;
2142 }
2143 (text.len() / CHARS_PER_TOKEN_FLOOR) as u64
2144}
2145
2146pub(crate) fn deadline_fitting_max_tokens(prompt_tokens: u64, remaining_ms: u64) -> Option<u64> {
2150 let prefill_ms = prompt_tokens
2151 .saturating_mul(1_000)
2152 .checked_div(env_u64("MEMRA_PREFILL_FLOOR_TOK_S", PREFILL_FLOOR_TOK_S))
2153 .unwrap_or(u64::MAX);
2154 let decode_ms = remaining_ms.checked_sub(prefill_ms)?;
2155 if decode_ms == 0 {
2156 return None;
2157 }
2158 Some(decode_ms.saturating_mul(env_u64("MEMRA_DECODE_FLOOR_TOK_S", DECODE_FLOOR_TOK_S)) / 1_000)
2159}
2160
2161pub(crate) fn nonstream_deadline_gate(
2171 request: &worker::Request,
2172 stream: bool,
2173 deadline: RequestDeadline,
2174 caller_declared_max_tokens: bool,
2175 tokenizer: Option<&Tokenizer>,
2176) -> Result<(), String> {
2177 if stream || !env_flag_on("MEMRA_NONSTREAM_DEADLINE_GATE", true) {
2178 return Ok(());
2179 }
2180 let max_new = request.params.max_new as u64;
2181 if !caller_declared_max_tokens || max_new == worker::MAX_NEW_CTX_BOUNDED as u64 || max_new == 0
2188 {
2189 return Ok(());
2190 }
2191 let prompt_tokens = prompt_tokens_estimate(request, tokenizer);
2192 let remaining_ms = deadline.remaining().as_millis() as u64;
2193 let prefill_ms = prompt_tokens.saturating_mul(1_000)
2194 / env_u64("MEMRA_PREFILL_FLOOR_TOK_S", PREFILL_FLOOR_TOK_S).max(1);
2195 let decode_ms = max_new.saturating_mul(1_000)
2196 / env_u64("MEMRA_DECODE_FLOOR_TOK_S", DECODE_FLOOR_TOK_S).max(1);
2197 let est_ms = prefill_ms.saturating_add(decode_ms);
2198 let bound_ms = remaining_ms.saturating_mul(DEADLINE_INFEASIBLE_MARGIN_PCT) / 100;
2199 if est_ms <= bound_ms {
2200 return Ok(());
2201 }
2202 let fits = deadline_fitting_max_tokens(prompt_tokens, remaining_ms);
2203 let advice = match fits {
2204 Some(fits) if fits > 0 => format!(
2205 "lower max_tokens to about {fits} for this prompt, or set \"stream\": true — a \
2206 stream's deadline bounds only the time to first token, so it may run as long \
2207 as it needs"
2208 ),
2209 _ => format!(
2210 "this prompt ({prompt_tokens} tok) needs most of the deadline before the first \
2211 token, so no max_tokens fits: set \"stream\": true"
2212 ),
2213 };
2214 Err(format!(
2215 "a non-streaming request for {max_new} tokens on a ~{prompt_tokens}-token prompt \
2216 needs an estimated ~{}s, which does not fit the {remaining_ms} ms timeout_ms \
2217 deadline (max {TIMEOUT_MS_MAX} ms — a platform ceiling: the fronting proxy fails \
2218 a non-streaming response whose headers take ~100 s). Refused before any GPU work \
2219 rather than after the deadline: {advice}",
2220 est_ms / 1_000,
2221 ))
2222}
2223
2224fn max_queue_depth(cap: usize) -> usize {
2228 static D: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
2229 D.get_or_init(|| {
2230 std::env::var("MEMRA_MAX_QUEUE_DEPTH")
2231 .ok()
2232 .and_then(|v| v.parse().ok())
2233 })
2234 .unwrap_or(cap.saturating_mul(4))
2235}
2236
2237fn queue_wait_ceiling_s() -> u64 {
2244 static S: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
2245 *S.get_or_init(|| {
2246 std::env::var("MEMRA_QUEUE_WAIT_CEILING_S")
2247 .ok()
2248 .and_then(|v| v.parse().ok())
2249 .unwrap_or(0)
2250 })
2251}
2252
2253pub(crate) struct PendingAdmissionGuard {
2274 reserved: bool,
2275 lane: lanes::Lane,
2276}
2277
2278impl PendingAdmissionGuard {
2279 pub(crate) fn commit(mut self) {
2283 self.reserved = false;
2284 std::mem::forget(self);
2285 }
2286}
2287
2288impl Drop for PendingAdmissionGuard {
2289 fn drop(&mut self) {
2290 if self.reserved {
2291 worker::release_pending_admit();
2292 worker::release_admission_reservation(self.lane);
2293 }
2294 }
2295}
2296
2297#[allow(clippy::result_large_err)] pub(crate) fn reserve_pending_admit(
2299 st: &AppState,
2300 lane: lanes::Lane,
2301 rl: &RateLimit,
2302 deadline: RequestDeadline,
2303) -> Result<PendingAdmissionGuard, (Response, &'static str)> {
2304 reserve_pending_admit_with_ceiling(st, lane, rl, deadline, queue_wait_ceiling_s())
2305}
2306
2307#[allow(clippy::result_large_err)] fn reserve_pending_admit_with_ceiling(
2312 st: &AppState,
2313 lane: lanes::Lane,
2314 rl: &RateLimit,
2315 deadline: RequestDeadline,
2316 ceiling_s: u64,
2317) -> Result<PendingAdmissionGuard, (Response, &'static str)> {
2318 let cap = lane_cap(lane).max(1);
2323 let bound = max_queue_depth(cap);
2324 let reservations_for_lane = &worker::ADMISSION_RESERVATIONS[lane.idx()];
2325 loop {
2326 let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
2327 let reservations = reservations_for_lane.load(std::sync::atomic::Ordering::Acquire);
2328 let backlog = reservations;
2332 let est_wait_s = reset_estimate_s(&m).saturating_mul((backlog / cap + 1) as u64);
2333 if backlog >= bound {
2334 let msg = format!(
2335 "{} queue is at its bound ({backlog} queued, bound {bound}); this \
2336 request was not admitted and is not billed; retry after ~{est_wait_s}s (a \
2337 coarse estimate, not a promise)",
2338 lane.as_str()
2339 );
2340 let resp = retry_contract_response(
2341 (
2342 StatusCode::TOO_MANY_REQUESTS,
2343 Json(error_body(
2344 &msg,
2345 "rate_limit_error",
2346 None,
2347 Some("shed_queue"),
2348 )),
2349 )
2350 .into_response(),
2351 Some(est_wait_s),
2352 );
2353 return Err((resp, "shed_queue"));
2354 }
2355 let remaining_ms = deadline.remaining().as_millis() as u64;
2356 let waits_for_capacity = rl.remaining == 0 || backlog > 0;
2360 if lane == lanes::Lane::Interactive
2361 && waits_for_capacity
2362 && est_wait_s.saturating_mul(1_000) > remaining_ms
2363 {
2364 let msg = format!(
2365 "estimated queue wait ~{est_wait_s}s exceeds this request's remaining \
2366 timeout_ms deadline ({remaining_ms} ms); this request was not admitted and \
2367 is not billed; retry after ~{est_wait_s}s or raise timeout_ms (a coarse \
2368 estimate, not a promise)"
2369 );
2370 let resp = retry_contract_response(
2371 (
2372 StatusCode::TOO_MANY_REQUESTS,
2373 Json(error_body(
2374 &msg,
2375 "rate_limit_error",
2376 None,
2377 Some("shed_deadline"),
2378 )),
2379 )
2380 .into_response(),
2381 Some(est_wait_s),
2382 );
2383 return Err((resp, "shed_deadline"));
2384 }
2385 if lane == lanes::Lane::Interactive
2393 && waits_for_capacity
2394 && ceiling_s > 0
2395 && est_wait_s > ceiling_s
2396 {
2397 let msg = format!(
2398 "estimated queue wait ~{est_wait_s}s exceeds this deployment's queue-wait \
2399 ceiling ({ceiling_s}s); this request was not admitted and is not billed; \
2400 retry after ~{est_wait_s}s (a coarse estimate, not a promise)"
2401 );
2402 let resp = retry_contract_response(
2403 (
2404 StatusCode::TOO_MANY_REQUESTS,
2405 Json(error_body(
2406 &msg,
2407 "rate_limit_error",
2408 None,
2409 Some("shed_queue_wait"),
2410 )),
2411 )
2412 .into_response(),
2413 Some(est_wait_s),
2414 );
2415 return Err((resp, "shed_queue_wait"));
2416 }
2417 if reservations_for_lane
2418 .compare_exchange(
2419 reservations,
2420 reservations.saturating_add(1),
2421 std::sync::atomic::Ordering::AcqRel,
2422 std::sync::atomic::Ordering::Acquire,
2423 )
2424 .is_ok()
2425 {
2426 worker::PENDING_ADMITS.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
2430 return Ok(PendingAdmissionGuard {
2431 reserved: true,
2432 lane,
2433 });
2434 }
2435 }
2436}
2437
2438static DRAINING: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
2449
2450fn draining() -> bool {
2451 DRAINING.load(std::sync::atomic::Ordering::SeqCst)
2452}
2453
2454fn drain_deadline_s() -> u64 {
2456 static D: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
2457 *D.get_or_init(|| {
2458 std::env::var("MEMRA_DRAIN_S")
2459 .ok()
2460 .and_then(|v| v.parse().ok())
2461 .unwrap_or(30)
2462 })
2463}
2464
2465fn drain_response() -> Response {
2475 let resp = (
2476 StatusCode::SERVICE_UNAVAILABLE,
2477 Json(error_body(
2478 "server is draining (shutdown in progress); retry",
2479 "server_error",
2480 None,
2481 Some("draining"),
2482 )),
2483 )
2484 .into_response();
2485 retry_contract_response(resp, Some(drain_deadline_s()))
2486}
2487
2488struct RateLimit {
2490 limit: usize,
2491 remaining: usize,
2492 reset_s: u64,
2493}
2494
2495impl RateLimit {
2496 fn at_admit(
2501 lane: lanes::Lane,
2502 n_inflight: usize,
2503 metrics: &SharedMetrics,
2504 tenant: &auth::TenantCtx,
2505 n_tenant: usize,
2506 ) -> Self {
2507 let global = lane_cap(lane);
2508 let Some(t) = tenant.rate_limit.filter(|&t| t < global) else {
2509 return Self::compute(global, n_inflight, metrics);
2510 };
2511 let headroom = t
2512 .saturating_sub(n_tenant)
2513 .min(global.saturating_sub(n_inflight));
2514 Self::compute(t, t - headroom, metrics)
2516 }
2517
2518 fn compute(limit: usize, n_inflight: usize, metrics: &SharedMetrics) -> Self {
2519 let remaining = limit.saturating_sub(n_inflight);
2520 let reset_s = if remaining > 0 {
2521 0
2522 } else {
2523 let m = metrics.lock().map(|m| m.clone()).unwrap_or_default();
2524 reset_estimate_s(&m)
2525 };
2526 RateLimit {
2527 limit,
2528 remaining,
2529 reset_s,
2530 }
2531 }
2532
2533 fn attach(&self, mut resp: Response) -> Response {
2535 let h = resp.headers_mut();
2536 for (k, v) in [
2537 ("x-ratelimit-limit", self.limit as u64),
2538 ("x-ratelimit-remaining", self.remaining as u64),
2539 ("x-ratelimit-reset", self.reset_s),
2540 ] {
2541 if let Ok(v) = axum::http::HeaderValue::from_str(&v.to_string()) {
2542 h.insert(axum::http::HeaderName::from_static(k), v);
2543 }
2544 }
2545 resp
2546 }
2547}
2548
2549#[allow(clippy::result_large_err)] fn acquire_request_slot(
2554 st: &AppState,
2555 lane: lanes::Lane,
2556 tenant: &auth::TenantCtx,
2557 env: &Envelope,
2558) -> Result<(InflightGuard, RateLimit), Response> {
2559 let global = lane_cap(lane);
2560 let tenant_cap = tenant.rate_limit.filter(|&cap| cap < global);
2561 match InflightGuard::try_acquire(
2562 st.inflight.clone(),
2563 lane,
2564 st.tenant_inflight.clone(),
2565 &tenant.tenant,
2566 tenant_cap,
2567 ) {
2568 Ok((guard, n_inflight, n_tenant)) => {
2569 let rl = RateLimit::at_admit(lane, n_inflight, &st.metrics, tenant, n_tenant);
2570 Ok((guard, rl))
2571 }
2572 Err(n_tenant) => {
2573 let n_inflight = st.inflight[lane.idx()].load(std::sync::atomic::Ordering::SeqCst);
2574 let rl = RateLimit::at_admit(lane, n_inflight, &st.metrics, tenant, n_tenant);
2575 let error =
2576 worker::EngineError::rate_limit("api key concurrent request limit reached; retry");
2577 Err(rl.attach(with_request_id(&env.id, engine_error_response(&error))))
2578 }
2579 }
2580}
2581
2582#[derive(Deserialize)]
2584struct CompletionReq {
2585 model: String,
2586 #[serde(default)]
2587 prompt: String,
2588 #[serde(default)]
2590 prompt_ids: Vec<u32>,
2591 #[serde(default)]
2594 max_tokens: Option<usize>,
2595 #[serde(default)]
2608 temperature: Option<f32>,
2609 #[serde(default)]
2610 top_p: Option<f32>,
2611 #[serde(default)]
2613 top_k: Option<usize>,
2614 #[serde(default)]
2616 min_p: Option<f32>,
2617 #[serde(default)]
2619 frequency_penalty: Option<f32>,
2620 #[serde(default)]
2621 presence_penalty: Option<f32>,
2622 #[serde(default)]
2624 repetition_penalty: Option<f32>,
2625 #[serde(default)]
2631 seed: Option<u64>,
2632 #[serde(default)]
2633 stop: StopSequences,
2634 #[serde(default)]
2637 logit_bias: Option<serde_json::Value>,
2638 #[serde(default)]
2639 logprobs: Option<serde_json::Value>,
2640 #[serde(default)]
2641 n: Option<usize>,
2642 #[serde(default)]
2643 best_of: Option<usize>,
2644 #[serde(default)]
2646 chat: bool,
2647 #[serde(default)]
2649 stream: bool,
2650 #[serde(default)]
2652 max_ctx: Option<usize>,
2653 #[serde(default)]
2655 trace_id: Option<String>,
2656 #[serde(default)]
2660 cache_salt: Option<String>,
2661 #[serde(default)]
2665 session_id: Option<String>,
2666 #[serde(default)]
2667 user: Option<String>,
2668 #[serde(default)]
2672 timeout_ms: Option<serde_json::Value>,
2673}
2674
2675#[derive(Deserialize)]
2676struct ChatMessage {
2677 role: String,
2678 #[serde(default)]
2680 content: serde_json::Value,
2681 #[serde(default)]
2683 tool_calls: Vec<ReqToolCall>,
2684 #[serde(default)]
2687 tool_call_id: Option<String>,
2688 #[serde(default)]
2691 name: Option<String>,
2692 #[serde(default, alias = "reasoning_content")]
2702 reasoning: Option<String>,
2703}
2704
2705#[derive(Deserialize)]
2706struct ReqToolCall {
2707 #[serde(default)]
2708 #[allow(dead_code)]
2709 id: Option<String>,
2710 function: ReqToolFunction,
2711}
2712
2713#[derive(Deserialize)]
2714struct ReqToolFunction {
2715 name: String,
2716 #[serde(default)]
2718 arguments: serde_json::Value,
2719}
2720
2721#[derive(Clone, Default, Deserialize)]
2722#[serde(untagged)]
2723enum StopSequences {
2724 One(String),
2725 Many(Vec<String>),
2726 #[default]
2727 None,
2728}
2729
2730impl StopSequences {
2731 fn into_vec(self) -> Vec<String> {
2737 let stops = match self {
2738 Self::One(stop) => vec![stop],
2739 Self::Many(stops) => stops,
2740 Self::None => Vec::new(),
2741 };
2742 stops.into_iter().filter(|s| !s.is_empty()).collect()
2743 }
2744
2745 fn validate(&self) -> Result<(), String> {
2746 let stops: &[String] = match self {
2747 Self::One(stop) => std::slice::from_ref(stop),
2748 Self::Many(stops) => stops,
2749 Self::None => &[],
2750 };
2751 if stops.len() > MAX_STOP_SEQUENCES {
2752 return Err(format!(
2753 "stop accepts at most {MAX_STOP_SEQUENCES} sequences"
2754 ));
2755 }
2756 let mut total = 0usize;
2757 for stop in stops {
2758 let bytes = stop.len();
2759 if bytes > MAX_STOP_SEQUENCE_BYTES {
2760 return Err(format!(
2761 "each stop sequence must be at most {MAX_STOP_SEQUENCE_BYTES} UTF-8 bytes"
2762 ));
2763 }
2764 total = total
2765 .checked_add(bytes)
2766 .ok_or_else(|| "stop sequence byte count overflowed".to_string())?;
2767 }
2768 if total > MAX_STOP_SEQUENCES_BYTES {
2769 return Err(format!(
2770 "stop sequences must total at most {MAX_STOP_SEQUENCES_BYTES} UTF-8 bytes"
2771 ));
2772 }
2773 Ok(())
2774 }
2775}
2776
2777#[derive(Deserialize)]
2783struct ChatCompletionReq {
2784 model: String,
2785 messages: Vec<ChatMessage>,
2786 #[serde(default, alias = "max_completion_tokens")]
2789 max_tokens: Option<usize>,
2790 #[serde(default)]
2793 temperature: Option<f32>,
2794 #[serde(default)]
2795 top_p: Option<f32>,
2796 #[serde(default)]
2800 top_k: Option<usize>,
2801 #[serde(default)]
2803 min_p: Option<f32>,
2804 #[serde(default)]
2806 frequency_penalty: Option<f32>,
2807 #[serde(default)]
2808 presence_penalty: Option<f32>,
2809 #[serde(default)]
2811 repetition_penalty: Option<f32>,
2812 #[serde(default)]
2814 seed: Option<u64>,
2815 #[serde(default)]
2816 stop: StopSequences,
2817 #[serde(default)]
2818 stream: bool,
2819 #[serde(default)]
2820 max_ctx: Option<usize>,
2821 #[serde(default)]
2826 response_format: Option<serde_json::Value>,
2827 #[serde(default)]
2828 logit_bias: Option<serde_json::Value>,
2829 #[serde(default)]
2830 logprobs: Option<serde_json::Value>,
2831 #[serde(default)]
2832 top_logprobs: Option<usize>,
2833 #[serde(default)]
2834 n: Option<usize>,
2835 #[serde(default)]
2837 tools: Vec<serde_json::Value>,
2838 #[serde(default)]
2840 tool_choice: Option<serde_json::Value>,
2841 #[serde(default)]
2847 reasoning_effort: Option<String>,
2848 #[serde(default)]
2857 reasoning: Option<serde_json::Value>,
2858 #[serde(default)]
2873 include_reasoning: Option<bool>,
2874 #[serde(default)]
2883 enable_thinking: Option<bool>,
2884 #[serde(default)]
2890 chat_template_kwargs: Option<serde_json::Value>,
2891 #[serde(default)]
2895 cache_salt: Option<String>,
2896 #[serde(default)]
2898 session_id: Option<String>,
2899 #[serde(default)]
2900 user: Option<String>,
2901 #[serde(default)]
2906 timeout_ms: Option<serde_json::Value>,
2907}
2908fn one() -> f32 {
2909 1.0
2910}
2911fn default_temperature() -> f32 {
2917 1.0
2918}
2919
2920#[derive(Debug, Clone, Copy, Default, PartialEq)]
2940struct SamplingDefaults {
2941 temperature: Option<f32>,
2942 top_p: Option<f32>,
2943 top_k: Option<usize>,
2944 min_p: Option<f32>,
2945 frequency_penalty: Option<f32>,
2946 presence_penalty: Option<f32>,
2947 repetition_penalty: Option<f32>,
2948}
2949
2950impl SamplingDefaults {
2951 fn resolve(metadata: Option<&OpenRouterModelMetadata>, caps: Option<&ModelCaps>) -> Self {
2954 SamplingDefaults {
2955 temperature: metadata
2956 .and_then(|m| m.default_temperature)
2957 .or_else(|| caps.and_then(|c| c.chat_temperature_default)),
2958 top_p: metadata
2959 .and_then(|m| m.default_top_p)
2960 .or_else(|| caps.and_then(|c| c.chat_top_p_default)),
2961 top_k: metadata.and_then(|m| m.default_top_k),
2962 min_p: metadata.and_then(|m| m.default_min_p),
2963 frequency_penalty: metadata.and_then(|m| m.default_frequency_penalty),
2964 presence_penalty: metadata.and_then(|m| m.default_presence_penalty),
2965 repetition_penalty: metadata.and_then(|m| m.default_repetition_penalty),
2966 }
2967 }
2968}
2969
2970#[derive(Debug, Clone, Copy, Default, PartialEq)]
2983struct ModelSamplingDefaults {
2984 thinking: SamplingDefaults,
2985 non_thinking: Option<SamplingDefaults>,
2986}
2987
2988impl ModelSamplingDefaults {
2989 fn resolve(metadata: Option<&OpenRouterModelMetadata>, caps: Option<&ModelCaps>) -> Self {
2990 ModelSamplingDefaults {
2991 thinking: SamplingDefaults::resolve(metadata, caps),
2992 non_thinking: metadata
2997 .and_then(|m| m.non_thinking_sampling.as_ref())
2998 .map(|arm| SamplingDefaults {
2999 temperature: arm.temperature,
3000 top_p: arm.top_p,
3001 top_k: arm.top_k,
3002 min_p: arm.min_p,
3003 frequency_penalty: arm.frequency_penalty,
3004 presence_penalty: arm.presence_penalty,
3005 repetition_penalty: arm.repetition_penalty,
3006 }),
3007 }
3008 }
3009
3010 fn for_mode(&self, think: ThinkMode) -> &SamplingDefaults {
3023 match (think, &self.non_thinking) {
3024 (ThinkMode::NoThink, Some(non_thinking)) => non_thinking,
3025 _ => &self.thinking,
3026 }
3027 }
3028
3029 #[cfg(test)] fn single(thinking: SamplingDefaults) -> Self {
3033 ModelSamplingDefaults {
3034 thinking,
3035 non_thinking: None,
3036 }
3037 }
3038}
3039
3040#[derive(Debug, Clone, Copy, Default)]
3046struct ClientSampling {
3047 temperature: Option<f32>,
3048 top_p: Option<f32>,
3049 top_k: Option<usize>,
3050 min_p: Option<f32>,
3051 frequency_penalty: Option<f32>,
3052 presence_penalty: Option<f32>,
3053 repetition_penalty: Option<f32>,
3054 seed: Option<u64>,
3055}
3056
3057impl From<&CompletionReq> for ClientSampling {
3058 fn from(r: &CompletionReq) -> Self {
3059 ClientSampling {
3060 temperature: r.temperature,
3061 top_p: r.top_p,
3062 top_k: r.top_k,
3063 min_p: r.min_p,
3064 frequency_penalty: r.frequency_penalty,
3065 presence_penalty: r.presence_penalty,
3066 repetition_penalty: r.repetition_penalty,
3067 seed: r.seed,
3068 }
3069 }
3070}
3071
3072impl From<&ChatCompletionReq> for ClientSampling {
3073 fn from(r: &ChatCompletionReq) -> Self {
3074 ClientSampling {
3075 temperature: r.temperature,
3076 top_p: r.top_p,
3077 top_k: r.top_k,
3078 min_p: r.min_p,
3079 frequency_penalty: r.frequency_penalty,
3080 presence_penalty: r.presence_penalty,
3081 repetition_penalty: r.repetition_penalty,
3082 seed: r.seed,
3083 }
3084 }
3085}
3086
3087fn resolve_sampler_config(client: ClientSampling, defaults: &SamplingDefaults) -> SamplerConfig {
3093 sampler_config(
3094 client
3095 .temperature
3096 .or(defaults.temperature)
3097 .unwrap_or_else(default_temperature),
3098 client.top_k.or(defaults.top_k).unwrap_or(0),
3099 client.top_p.or(defaults.top_p).unwrap_or_else(one),
3100 client.min_p.or(defaults.min_p).unwrap_or(0.0),
3101 client
3102 .frequency_penalty
3103 .or(defaults.frequency_penalty)
3104 .unwrap_or(0.0),
3105 client
3106 .presence_penalty
3107 .or(defaults.presence_penalty)
3108 .unwrap_or(0.0),
3109 client
3110 .repetition_penalty
3111 .or(defaults.repetition_penalty)
3112 .unwrap_or_else(one),
3113 client.seed,
3114 )
3115}
3116
3117#[derive(Serialize)]
3118struct CompletionResp {
3119 model: String,
3120 text: String,
3121 tokens: Vec<u32>,
3122 stop_reason: String,
3126 #[serde(default, skip_serializing_if = "Option::is_none")]
3131 error: Option<serde_json::Value>,
3132 n_tokens: usize,
3133 prompt_tokens: usize,
3136 cached_tokens: usize,
3137 elapsed_s: f64,
3138}
3139
3140fn usage_json(
3148 n_prompt: usize,
3149 n_tokens: usize,
3150 n_cached: usize,
3151 elapsed_s: f64,
3152 spec: Option<worker::SpecUsage>,
3153) -> serde_json::Value {
3154 let mut u = json!({
3155 "prompt_tokens": n_prompt,
3156 "completion_tokens": n_tokens,
3157 "total_tokens": n_prompt + n_tokens,
3158 "prompt_tokens_details": { "cached_tokens": n_cached },
3159 "elapsed_s": elapsed_s,
3160 });
3161 if let Some(sp) = spec {
3162 u["spec"] = json!({
3163 "rounds": sp.rounds,
3164 "drafted": sp.drafted,
3165 "accepted": sp.accepted,
3166 "acceptance_rate": if sp.drafted > 0 {
3167 sp.accepted as f64 / sp.drafted as f64 } else { 0.0 },
3168 });
3169 }
3170 u
3171}
3172
3173pub const SYSTEM_FINGERPRINT: &str = concat!(
3205 "memra-",
3206 env!("CARGO_PKG_VERSION"),
3207 "-",
3208 env!("MEMRA_BUILD_ID")
3209);
3210
3211pub const BUILD_ID_SRC: &str = env!("MEMRA_BUILD_ID_SRC");
3213
3214pub const BUILD_ID_NOTE: &str = env!("MEMRA_BUILD_ID_NOTE");
3216
3217pub const BUILD_GIT_SHA: &str = env!("MEMRA_BUILD_SHA");
3221
3222pub fn build_identity_line() -> String {
3225 format!("[server] build: {SYSTEM_FINGERPRINT} (id: {BUILD_ID_SRC}, git: {BUILD_GIT_SHA})")
3226}
3227
3228fn gen_hex128() -> String {
3231 use std::hash::{BuildHasher, Hasher};
3232 static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
3233 let n = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3234 let t = std::time::SystemTime::now()
3235 .duration_since(std::time::UNIX_EPOCH)
3236 .map(|d| d.as_nanos() as u64)
3237 .unwrap_or(0);
3238 let mut h1 = std::collections::hash_map::RandomState::new().build_hasher();
3239 h1.write_u64(n);
3240 h1.write_u64(t);
3241 let mut h2 = std::collections::hash_map::RandomState::new().build_hasher();
3242 h2.write_u64(t.rotate_left(17));
3243 h2.write_u64(n);
3244 format!("{:016x}{:016x}", h1.finish(), h2.finish())
3245}
3246
3247#[derive(Clone)]
3250struct Envelope {
3251 id: String,
3252 created: u64,
3253}
3254
3255impl Envelope {
3256 fn new(chat: bool) -> Self {
3257 Envelope {
3258 id: format!(
3259 "{}-{}",
3260 if chat { "chatcmpl" } else { "cmpl" },
3261 gen_hex128()
3262 ),
3263 created: std::time::SystemTime::now()
3264 .duration_since(std::time::UNIX_EPOCH)
3265 .map(|d| d.as_secs())
3266 .unwrap_or(0),
3267 }
3268 }
3269
3270 fn capture_child(&self, index: usize) -> Self {
3287 Envelope {
3288 id: format!("{}.{index}", self.id),
3289 created: self.created,
3290 }
3291 }
3292
3293 fn stamp(&self, mut v: serde_json::Value) -> serde_json::Value {
3295 v["id"] = json!(self.id);
3296 v["created"] = json!(self.created);
3297 v["system_fingerprint"] = json!(SYSTEM_FINGERPRINT);
3298 v
3299 }
3300}
3301
3302fn with_request_id(id: &str, mut resp: Response) -> Response {
3304 if let Ok(v) = axum::http::HeaderValue::from_str(id) {
3305 resp.headers_mut()
3306 .insert(axum::http::HeaderName::from_static("x-request-id"), v);
3307 }
3308 resp
3309}
3310
3311fn openai_compat() -> bool {
3319 static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3320 *C.get_or_init(|| match std::env::var("MEMRA_COMPAT").as_deref() {
3321 Ok("openai") => true,
3322 Ok(_) => false,
3323 Err(_) => std::env::var("MEMRA_API_KEY").is_ok(),
3324 })
3325}
3326
3327fn cache_namespace(cache_salt: &Option<String>) -> String {
3339 cache_salt.clone().unwrap_or_default()
3340}
3341
3342const CACHE_SALT_MAX_BYTES: usize = 64;
3343
3344fn validate_cache_namespace(
3345 cache_salt: &Option<String>,
3346 keyring_configured: bool,
3347) -> Result<String, &'static str> {
3348 let raw = cache_namespace(cache_salt);
3349 if raw.len() > CACHE_SALT_MAX_BYTES {
3350 return Err("cache_salt must be at most 64 bytes");
3351 }
3352 if !keyring_configured && raw.starts_with("t:") {
3353 return Err("cache_salt must not use the reserved t: prefix without a keyring");
3354 }
3355 if !raw
3356 .bytes()
3357 .all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'_' | b'.' | b'+' | b'/' | b'='))
3358 {
3359 return Err("cache_salt contains unsupported characters");
3360 }
3361 Ok(raw)
3362}
3363
3364fn affinity_key(
3380 session_id: &Option<String>,
3381 user: &Option<String>,
3382 headers: &axum::http::HeaderMap,
3383) -> Result<Option<String>, String> {
3384 let clean = |s: &str| -> Result<Option<String>, String> {
3385 let t = s.trim();
3386 if t.is_empty() {
3387 Ok(None)
3388 } else if t.len() > MAX_CLIENT_IDENTIFIER_BYTES {
3389 Err(format!(
3390 "session identity must be at most {MAX_CLIENT_IDENTIFIER_BYTES} UTF-8 bytes"
3391 ))
3392 } else if t.chars().any(char::is_control) {
3393 Err("session identity must not contain control characters".into())
3394 } else {
3395 Ok(Some(t.to_string()))
3396 }
3397 };
3398 if let Some(value) = session_id.as_deref()
3399 && let Some(value) = clean(value)?
3400 {
3401 return Ok(Some(value));
3402 }
3403 if let Some(value) = user.as_deref()
3404 && let Some(value) = clean(value)?
3405 {
3406 return Ok(Some(value));
3407 }
3408 match headers.get("x-session-id") {
3409 Some(value) => clean(
3410 value
3411 .to_str()
3412 .map_err(|_| "x-session-id must contain visible ASCII or UTF-8 text")?,
3413 ),
3414 None => Ok(None),
3415 }
3416}
3417
3418fn validate_client_identifier(value: Option<&str>, name: &str) -> Result<(), String> {
3419 let Some(value) = value else {
3420 return Ok(());
3421 };
3422 if value.len() > MAX_CLIENT_IDENTIFIER_BYTES {
3423 return Err(format!(
3424 "{name} must be at most {MAX_CLIENT_IDENTIFIER_BYTES} UTF-8 bytes"
3425 ));
3426 }
3427 if value.chars().any(char::is_control) {
3428 return Err(format!("{name} must not contain control characters"));
3429 }
3430 Ok(())
3431}
3432
3433fn error_body(
3438 message: &str,
3439 etype: &str,
3440 param: Option<&str>,
3441 code: Option<&str>,
3442) -> serde_json::Value {
3443 json!({ "error": {
3444 "message": message,
3445 "type": etype,
3446 "param": param,
3447 "code": code,
3448 } })
3449}
3450
3451fn error_response(status: StatusCode, message: &str, etype: &str, param: Option<&str>) -> Response {
3452 error_response_coded(status, message, etype, param, None)
3453}
3454
3455fn error_response_coded(
3460 status: StatusCode,
3461 message: &str,
3462 etype: &str,
3463 param: Option<&str>,
3464 code: Option<&str>,
3465) -> Response {
3466 let mut resp = (status, Json(error_body(message, etype, param, code))).into_response();
3467 if status.is_client_error()
3468 && status != StatusCode::TOO_MANY_REQUESTS
3469 && status != StatusCode::REQUEST_TIMEOUT
3470 && status != StatusCode::CONFLICT
3471 {
3472 resp.headers_mut().insert(
3473 "x-should-retry",
3474 axum::http::HeaderValue::from_static("false"),
3475 );
3476 }
3477 resp
3478}
3479
3480fn bad_request(message: &str, param: Option<&str>) -> Response {
3481 error_response(
3482 StatusCode::BAD_REQUEST,
3483 message,
3484 "invalid_request_error",
3485 param,
3486 )
3487}
3488
3489const 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>) {
3518 use worker::ErrClass as C;
3519 match class {
3520 C::InvalidRequest => (StatusCode::BAD_REQUEST, "invalid_request_error", None),
3521 C::ContextLength => (
3522 StatusCode::BAD_REQUEST,
3523 "invalid_request_error",
3524 Some("context_length_exceeded"),
3525 ),
3526 C::ModelNotFound => (
3527 StatusCode::BAD_REQUEST,
3528 "invalid_request_error",
3529 Some("model_not_found"),
3530 ),
3531 C::RateLimit => (
3532 StatusCode::TOO_MANY_REQUESTS,
3533 "rate_limit_error",
3534 Some("rate_limit_exceeded"),
3535 ),
3536 C::Overloaded => (
3537 StatusCode::SERVICE_UNAVAILABLE,
3538 "server_error",
3539 Some("overloaded"),
3540 ),
3541 C::Engine => (
3542 StatusCode::INTERNAL_SERVER_ERROR,
3543 "server_error",
3544 Some("engine_error"),
3545 ),
3546 }
3547}
3548
3549fn class_retry_after_s(class: worker::ErrClass) -> Option<u64> {
3551 use worker::ErrClass as C;
3552 match class {
3553 C::RateLimit => Some(RETRY_AFTER_S_RATE_LIMIT),
3554 C::Overloaded => Some(RETRY_AFTER_S_OVERLOADED),
3555 C::Engine | C::InvalidRequest | C::ContextLength | C::ModelNotFound => None,
3559 }
3560}
3561
3562fn engine_error_body(e: &worker::EngineError) -> serde_json::Value {
3565 let (_, etype, code) = class_http(e.class);
3566 error_body(&e.message, etype, e.param, code)
3567}
3568
3569fn engine_error_response(e: &worker::EngineError) -> Response {
3575 engine_error_response_with_retry_after(
3576 e,
3577 e.retry_after_s.or_else(|| class_retry_after_s(e.class)),
3578 )
3579}
3580
3581fn engine_error_response_with_retry_after(
3582 e: &worker::EngineError,
3583 retry_after_s: Option<u64>,
3584) -> Response {
3585 let (status, _, _) = class_http(e.class);
3586 let resp = (status, Json(engine_error_body(e))).into_response();
3587 retry_contract_response(resp, retry_after_s)
3588}
3589
3590fn retry_contract_response(mut resp: Response, retry_after_s: Option<u64>) -> Response {
3592 let status = resp.status();
3593 let h = resp.headers_mut();
3594 match retry_after_s {
3595 Some(secs) => {
3596 let secs = secs.clamp(1, 60);
3598 if let Ok(v) = axum::http::HeaderValue::from_str(&secs.to_string()) {
3599 h.insert(axum::http::header::RETRY_AFTER, v);
3600 }
3601 if let Ok(v) = axum::http::HeaderValue::from_str(&(secs * 1000).to_string()) {
3602 h.insert("retry-after-ms", v);
3603 }
3604 }
3605 None if status.is_client_error() => {
3606 h.insert(
3609 "x-should-retry",
3610 axum::http::HeaderValue::from_static("false"),
3611 );
3612 }
3613 None => {}
3614 }
3615 resp
3616}
3617
3618fn worker_unavailable_response() -> Response {
3619 engine_error_response_with_retry_after(
3620 &worker::EngineError::overloaded("worker unavailable"),
3621 Some(worker::WORKER_RESPAWN_BACKOFF_BASE_S),
3622 )
3623}
3624
3625fn stop_reason_to_finish(r: &str) -> &'static str {
3626 match r {
3627 "Eos" | "Callback" => "stop",
3628 "MaxNew" | "ContextFull" => "length",
3629 _ => "stop",
3630 }
3631}
3632
3633fn content_to_text(v: &serde_json::Value) -> Result<String, String> {
3637 match v {
3638 serde_json::Value::Null => Ok(String::new()),
3639 serde_json::Value::String(s) => Ok(s.clone()),
3640 serde_json::Value::Array(parts) => {
3641 let mut out = String::new();
3642 for p in parts {
3643 match p.get("type").and_then(|t| t.as_str()) {
3644 Some("text") | None => match p.get("text").and_then(|t| t.as_str()) {
3645 Some(t) => out.push_str(t),
3646 None => return Err("content part has no text field".into()),
3647 },
3648 Some(other) => {
3649 return Err(format!(
3650 "unsupported content part type {other:?} (text only)"
3651 ));
3652 }
3653 }
3654 }
3655 Ok(out)
3656 }
3657 _ => Err("content must be a string, null, or an array of text parts".into()),
3658 }
3659}
3660
3661pub(crate) static VISION_PLACEMENT_SERVING: std::sync::atomic::AtomicBool =
3679 std::sync::atomic::AtomicBool::new(true);
3680
3681fn vision_placement_serving() -> bool {
3682 VISION_PLACEMENT_SERVING.load(std::sync::atomic::Ordering::Acquire)
3683}
3684
3685fn vision_media_admissible(placement: bool, kind: &str) -> Result<(), String> {
3691 if placement {
3692 Ok(())
3693 } else {
3694 Err(format!(
3695 "{kind} input is not enabled on this deployment (vision overlay placement \
3696 inadmissible at boot: see the worker's IMAGE INPUT DISABLED line)"
3697 ))
3698 }
3699}
3700
3701fn vision_placement_admits(kind: &str) -> Result<(), String> {
3702 vision_media_admissible(vision_placement_serving(), kind)
3703}
3704
3705fn vision_enabled() -> bool {
3709 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3710 *ON.get_or_init(|| {
3711 std::env::var("MEMRA_VISION_DIR").is_ok()
3712 && std::env::var("MEMRA_VISION").as_deref() != Ok("0")
3713 })
3714}
3715
3716fn gemma_vision_enabled() -> bool {
3721 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3722 *ON.get_or_init(|| {
3723 std::env::var("MEMRA_GEMMA_VISION").as_deref() == Ok("1")
3724 && std::env::var("MEMRA_GEMMA_MMPROJ").is_ok()
3725 })
3726}
3727
3728pub(crate) static GLM5_VISION_SERVING: std::sync::atomic::AtomicBool =
3738 std::sync::atomic::AtomicBool::new(false);
3739
3740fn glm5_vision_enabled() -> bool {
3743 GLM5_VISION_SERVING.load(std::sync::atomic::Ordering::Acquire)
3744}
3745
3746fn step_vision_enabled() -> bool {
3754 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3755 *ON.get_or_init(|| {
3756 std::env::var("MEMRA_STEP_VISION_DIR").is_ok()
3757 && std::env::var("MEMRA_STEP_VISION").as_deref() != Ok("0")
3758 })
3759}
3760
3761const VISION_MAX_IMAGES: usize = 8;
3763
3764pub(crate) const MAX_VISION_PATCH_BYTES: usize = 1 << 30; static VISION_PATCH_BYTES_IN_USE: std::sync::atomic::AtomicUsize =
3770 std::sync::atomic::AtomicUsize::new(0);
3771pub(crate) static VISION_PREPROCESS_SEMAPHORE: tokio::sync::Semaphore =
3775 tokio::sync::Semaphore::const_new(1);
3776
3777#[allow(clippy::result_large_err)]
3780pub(crate) fn try_vision_preprocess(
3781 required: bool,
3782) -> Result<Option<tokio::sync::SemaphorePermit<'static>>, Response> {
3783 try_vision_preprocess_with(required, &VISION_PREPROCESS_SEMAPHORE)
3784}
3785
3786#[allow(clippy::result_large_err)]
3787fn try_vision_preprocess_with(
3788 required: bool,
3789 semaphore: &'static tokio::sync::Semaphore,
3790) -> Result<Option<tokio::sync::SemaphorePermit<'static>>, Response> {
3791 if !required {
3792 return Ok(None);
3793 }
3794 match semaphore.try_acquire() {
3795 Ok(permit) => Ok(Some(permit)),
3796 Err(tokio::sync::TryAcquireError::NoPermits) => Err(retry_contract_response(
3797 error_response_coded(
3798 StatusCode::TOO_MANY_REQUESTS,
3799 "vision preprocessing is busy",
3800 "rate_limit_error",
3801 Some("messages"),
3802 Some("vision_preprocess_busy"),
3803 ),
3804 Some(BODY_ADMISSION_RETRY_AFTER_S),
3805 )),
3806 Err(tokio::sync::TryAcquireError::Closed) => Err(error_response_coded(
3807 StatusCode::SERVICE_UNAVAILABLE,
3808 "vision preprocessing is unavailable",
3809 "server_error",
3810 Some("messages"),
3811 Some("vision_preprocess_unavailable"),
3812 )),
3813 }
3814}
3815
3816pub(crate) struct VisionMemoryPermit {
3817 bytes: usize,
3818}
3819
3820#[derive(Debug)]
3821pub(crate) enum VisionMemoryError {
3822 Request(String),
3823 Capacity(String),
3824}
3825
3826impl Drop for VisionMemoryPermit {
3827 fn drop(&mut self) {
3828 if self.bytes != 0 {
3829 VISION_PATCH_BYTES_IN_USE.fetch_sub(self.bytes, std::sync::atomic::Ordering::AcqRel);
3830 }
3831 }
3832}
3833
3834fn try_reserve_vision_memory(
3835 bytes: usize,
3836) -> Result<Option<VisionMemoryPermit>, VisionMemoryError> {
3837 if bytes == 0 {
3838 return Ok(None);
3839 }
3840 if bytes > MAX_VISION_PATCH_BYTES {
3841 return Err(VisionMemoryError::Request(format!(
3842 "vision preprocessing requires {bytes} bytes of patch memory, exceeding the {} MiB request limit",
3843 MAX_VISION_PATCH_BYTES / (1024 * 1024)
3844 )));
3845 }
3846 let mut in_use = VISION_PATCH_BYTES_IN_USE.load(std::sync::atomic::Ordering::Acquire);
3847 loop {
3848 let Some(next) = in_use.checked_add(bytes) else {
3849 return Err(VisionMemoryError::Capacity(
3850 "vision patch memory reservation overflowed".into(),
3851 ));
3852 };
3853 if next > MAX_VISION_PATCH_BYTES {
3854 return Err(VisionMemoryError::Capacity(format!(
3855 "vision preprocessing is at capacity ({} MiB reserved; request needs {} MiB)",
3856 in_use / (1024 * 1024),
3857 bytes / (1024 * 1024)
3858 )));
3859 }
3860 match VISION_PATCH_BYTES_IN_USE.compare_exchange_weak(
3861 in_use,
3862 next,
3863 std::sync::atomic::Ordering::AcqRel,
3864 std::sync::atomic::Ordering::Acquire,
3865 ) {
3866 Ok(_) => return Ok(Some(VisionMemoryPermit { bytes })),
3867 Err(actual) => in_use = actual,
3868 }
3869 }
3870}
3871
3872pub(crate) fn vision_memory_error_response(
3873 error: VisionMemoryError,
3874 param: Option<&str>,
3875) -> Response {
3876 match error {
3877 VisionMemoryError::Request(message) => bad_request(&message, param),
3878 VisionMemoryError::Capacity(message) => retry_contract_response(
3879 error_response_coded(
3880 StatusCode::SERVICE_UNAVAILABLE,
3881 &message,
3882 "server_error",
3883 None,
3884 Some("vision_memory_busy"),
3885 ),
3886 Some(RETRY_AFTER_S_OVERLOADED),
3887 ),
3888 }
3889}
3890
3891enum PendingVisionUnit {
3897 Still {
3898 bytes: Vec<u8>,
3899 gh: usize,
3900 gw: usize,
3901 },
3902 Video {
3903 bytes: Vec<u8>,
3904 groups: Vec<memra_engine::vision_pre::PlannedVideoGroup>,
3905 video: usize,
3906 },
3907}
3908
3909struct PendingGemmaImage {
3911 bytes: Vec<u8>,
3912 gw: usize,
3913 gh: usize,
3914}
3915
3916struct PendingGlm5Image {
3919 bytes: Vec<u8>,
3920 gh: usize,
3921 gw: usize,
3922}
3923
3924struct PendingStepImage {
3927 bytes: Vec<u8>,
3928 plan: memra_engine::vision_step::StepImagePlan,
3929}
3930
3931fn content_to_text_vision_step(
3941 v: &serde_json::Value,
3942 step_images: &mut Vec<PendingStepImage>,
3943) -> Result<String, String> {
3944 use memra_engine::vision_step::{SV_MAIN_ROWS, SV_TILE_ROWS};
3945 let parts = match v {
3946 serde_json::Value::Array(parts) => parts,
3947 _ => return content_to_text(v),
3948 };
3949 let mut out = String::new();
3950 let mut needs_sep = false;
3951 for p in parts {
3952 match p.get("type").and_then(|t| t.as_str()) {
3953 Some("text") | None => match p.get("text").and_then(|t| t.as_str()) {
3954 Some(t) => {
3955 if needs_sep {
3956 out.push(' ');
3957 }
3958 out.push_str(t);
3959 needs_sep = true;
3960 }
3961 None => return Err("content part has no text field".into()),
3962 },
3963 Some("image_url") => {
3964 vision_placement_admits("image")?;
3965 let url = p
3966 .get("image_url")
3967 .and_then(|u| {
3968 if u.is_string() {
3969 u.as_str()
3970 } else {
3971 u.get("url").and_then(|x| x.as_str())
3972 }
3973 })
3974 .ok_or("image_url part has no url")?;
3975 if !url.starts_with("data:") {
3976 return Err(
3977 "image_url must be a base64 data URI (http(s) fetch is disabled)".into(),
3978 );
3979 }
3980 if step_images.len() >= VISION_MAX_IMAGES {
3981 return Err(format!("too many images (max {VISION_MAX_IMAGES})"));
3982 }
3983 let bytes = memra_engine::vision_pre::decode_data_uri(url)
3987 .map_err(|e| format!("image {}: {e}", step_images.len() + 1))?;
3988 let plan = memra_engine::vision_step::step_plan_image(&bytes)
3989 .map_err(|e| format!("image {}: {e}", step_images.len() + 1))?;
3990 for i in 0..plan.n_tiles {
3991 out.push_str("<patch_start>");
3992 for _ in 0..SV_TILE_ROWS {
3993 out.push_str("<im_patch>");
3994 }
3995 out.push_str("<patch_end>");
3996 if plan.newline_mask[i] {
3997 out.push_str("<patch_newline>");
3998 }
3999 }
4000 out.push_str("<im_start>");
4001 for _ in 0..SV_MAIN_ROWS {
4002 out.push_str("<im_patch>");
4003 }
4004 out.push_str("<im_end>");
4005 step_images.push(PendingStepImage { bytes, plan });
4006 needs_sep = false;
4007 }
4008 Some("video_url") => {
4009 return Err("step37 has no video input (image-only processor)".into());
4010 }
4011 Some(other) => {
4012 return Err(format!("unsupported content part type {other:?}"));
4013 }
4014 }
4015 }
4016 Ok(out)
4017}
4018
4019fn content_to_text_vision(
4028 v: &serde_json::Value,
4029 images: &mut Vec<PendingVisionUnit>,
4030 gemma_images: &mut Vec<PendingGemmaImage>,
4031 glm5_images: &mut Vec<PendingGlm5Image>,
4032 step_images: &mut Vec<PendingStepImage>,
4033 next_video: &mut usize,
4034) -> Result<String, String> {
4035 if step_vision_enabled() {
4039 return content_to_text_vision_step(v, step_images);
4040 }
4041 let parts = match v {
4042 serde_json::Value::Array(parts) => parts,
4043 _ => return content_to_text(v),
4044 };
4045 let mut out = String::new();
4046 for p in parts {
4047 match p.get("type").and_then(|t| t.as_str()) {
4048 Some("text") | None => match p.get("text").and_then(|t| t.as_str()) {
4049 Some(t) => out.push_str(t),
4050 None => return Err("content part has no text field".into()),
4051 },
4052 Some("image_url") if glm5_vision_enabled() => {
4053 let url = p
4054 .get("image_url")
4055 .and_then(|u| {
4056 if u.is_string() {
4057 u.as_str()
4058 } else {
4059 u.get("url").and_then(|x| x.as_str())
4060 }
4061 })
4062 .ok_or("image_url part has no url")?;
4063 if !url.starts_with("data:") {
4064 return Err(
4065 "image_url must be a base64 data URI (http(s) fetch is disabled)".into(),
4066 );
4067 }
4068 if glm5_images.len() >= VISION_MAX_IMAGES {
4069 return Err(format!("too many images (max {VISION_MAX_IMAGES})"));
4070 }
4071 let bytes = memra_engine::vision_pre::decode_data_uri(url)
4075 .map_err(|e| format!("image {}: {e}", glm5_images.len() + 1))?;
4076 let (gh, gw) = memra_engine::vision_glm5::glm5_plan_image(&bytes)
4077 .map_err(|e| format!("image {}: {e}", glm5_images.len() + 1))?;
4078 out.push_str("<|begin_of_image|>");
4082 for _ in 0..memra_engine::vision_glm5::n_merged_for_grid(gh, gw) {
4083 out.push_str("<|image|>");
4084 }
4085 out.push_str("<|end_of_image|>");
4086 glm5_images.push(PendingGlm5Image { bytes, gh, gw });
4087 }
4088 Some("video_url") if glm5_vision_enabled() => {
4089 return Err(
4090 "glm5 video input is not served (tensor census only; image input is the \
4091 supported surface)"
4092 .into(),
4093 );
4094 }
4095 Some("image_url") if gemma_vision_enabled() => {
4096 vision_placement_admits("image")?;
4097 let url = p
4098 .get("image_url")
4099 .and_then(|u| {
4100 if u.is_string() {
4101 u.as_str()
4102 } else {
4103 u.get("url").and_then(|x| x.as_str())
4104 }
4105 })
4106 .ok_or("image_url part has no url")?;
4107 if !url.starts_with("data:") {
4108 return Err(
4109 "image_url must be a base64 data URI (http(s) fetch is disabled)".into(),
4110 );
4111 }
4112 if gemma_images.len() >= VISION_MAX_IMAGES {
4113 return Err(format!("too many images (max {VISION_MAX_IMAGES})"));
4114 }
4115 let bytes = memra_engine::vision_gemma::gemma_decode_data_uri(url)
4119 .map_err(|e| format!("image {}: {e}", gemma_images.len() + 1))?;
4120 let (gw, gh) = memra_engine::vision_gemma::gemma_plan_image(&bytes)
4121 .map_err(|e| format!("image {}: {e}", gemma_images.len() + 1))?;
4122 out.push_str("<|image>");
4124 for _ in 0..memra_engine::vision_gemma::n_soft_for_grid(gw, gh) {
4125 out.push_str("<|image|>");
4126 }
4127 out.push_str("<image|>");
4128 gemma_images.push(PendingGemmaImage { bytes, gw, gh });
4129 }
4130 Some("image_url") => {
4131 if !vision_enabled() {
4132 return Err("image input is not enabled on this deployment".into());
4133 }
4134 vision_placement_admits("image")?;
4135 let url = p
4136 .get("image_url")
4137 .and_then(|u| {
4138 if u.is_string() {
4139 u.as_str()
4140 } else {
4141 u.get("url").and_then(|x| x.as_str())
4142 }
4143 })
4144 .ok_or("image_url part has no url")?;
4145 if !url.starts_with("data:") {
4146 return Err(
4147 "image_url must be a base64 data URI (http(s) fetch is disabled)".into(),
4148 );
4149 }
4150 if images
4151 .iter()
4152 .filter(|u| matches!(u, PendingVisionUnit::Still { .. }))
4153 .count()
4154 >= VISION_MAX_IMAGES
4155 {
4156 return Err(format!("too many images (max {VISION_MAX_IMAGES})"));
4157 }
4158 let bytes = memra_engine::vision_pre::decode_data_uri(url)
4163 .map_err(|e| format!("image {}: {e}", images.len() + 1))?;
4164 let (gh, gw) = memra_engine::vision_pre::plan_image_bytes(&bytes)
4165 .map_err(|e| format!("image {}: {e}", images.len() + 1))?;
4166 out.push_str("<|vision_start|>");
4167 for _ in 0..memra_engine::vision_pre::n_tokens_for_grid(gh, gw) {
4168 out.push_str("<|image_pad|>");
4169 }
4170 out.push_str("<|vision_end|>");
4171 images.push(PendingVisionUnit::Still { bytes, gh, gw });
4172 }
4173 Some("video_url") if gemma_vision_enabled() => {
4174 return Err("gemma-4 has no video input (image-only projector)".into());
4175 }
4176 Some("video_url") => {
4177 if !vision_enabled() {
4178 return Err("video input is not enabled on this deployment".into());
4179 }
4180 vision_placement_admits("video")?;
4181 let url = p
4182 .get("video_url")
4183 .and_then(|u| {
4184 if u.is_string() {
4185 u.as_str()
4186 } else {
4187 u.get("url").and_then(|x| x.as_str())
4188 }
4189 })
4190 .ok_or("video_url part has no url")?;
4191 if !url.starts_with("data:") {
4192 return Err(
4193 "video_url must be a base64 data URI (http(s) fetch is disabled)".into(),
4194 );
4195 }
4196 if *next_video >= 2 {
4197 return Err("too many videos (max 2)".into());
4198 }
4199 let bytes = memra_engine::vision_pre::decode_data_uri(url)?;
4202 let vid = memra_engine::vision_pre::plan_video_gif(&bytes)
4203 .map_err(|e| format!("video: {e}"))?;
4204 let vidx = *next_video;
4205 *next_video += 1;
4206 for group in &vid.groups {
4208 out.push_str(&format!("<{:.1} seconds>", group.timestamp));
4209 out.push_str("<|vision_start|>");
4210 for _ in 0..memra_engine::vision_pre::n_tokens_for_grid(group.gh, group.gw) {
4211 out.push_str("<|video_pad|>");
4212 }
4213 out.push_str("<|vision_end|>");
4214 }
4215 images.push(PendingVisionUnit::Video {
4218 bytes,
4219 groups: vid.groups,
4220 video: vidx,
4221 });
4222 }
4223 Some(other) => {
4224 return Err(format!("unsupported content part type {other:?}"));
4225 }
4226 }
4227 }
4228 Ok(out)
4229}
4230
4231fn pyjson(v: &serde_json::Value, out: &mut String) {
4235 match v {
4236 serde_json::Value::Object(m) => {
4237 out.push('{');
4238 for (i, (k, val)) in m.iter().enumerate() {
4239 if i > 0 {
4240 out.push_str(", ");
4241 }
4242 out.push_str(&serde_json::Value::String(k.clone()).to_string());
4243 out.push_str(": ");
4244 pyjson(val, out);
4245 }
4246 out.push('}');
4247 }
4248 serde_json::Value::Array(a) => {
4249 out.push('[');
4250 for (i, val) in a.iter().enumerate() {
4251 if i > 0 {
4252 out.push_str(", ");
4253 }
4254 pyjson(val, out);
4255 }
4256 out.push(']');
4257 }
4258 scalar => out.push_str(&scalar.to_string()),
4259 }
4260}
4261
4262fn pyjson_str(v: &serde_json::Value) -> String {
4263 let mut s = String::new();
4264 pyjson(v, &mut s);
4265 s
4266}
4267
4268#[allow(clippy::too_many_arguments)] fn sampler_config(
4276 temperature: f32,
4277 top_k: usize,
4278 top_p: f32,
4279 min_p: f32,
4280 frequency_penalty: f32,
4281 presence_penalty: f32,
4282 repetition_penalty: f32,
4283 seed: Option<u64>,
4284) -> SamplerConfig {
4285 let penalties_on =
4286 frequency_penalty != 0.0 || presence_penalty != 0.0 || repetition_penalty != 1.0;
4287 SamplerConfig {
4288 temperature,
4289 top_k,
4290 top_p,
4291 min_p,
4292 penalty_last_n: if penalties_on {
4293 memra_engine::spec::PEN_WINDOW_MAX
4294 } else {
4295 0
4296 },
4297 penalty_repeat: repetition_penalty,
4298 penalty_freq: frequency_penalty,
4299 penalty_present: presence_penalty,
4300 seed: seed.unwrap_or_else(fresh_seed),
4303 }
4304}
4305
4306fn fresh_seed() -> u64 {
4311 use std::sync::atomic::{AtomicU64, Ordering};
4312 static COUNTER: AtomicU64 = AtomicU64::new(0);
4313 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
4314 let nanos = std::time::SystemTime::now()
4315 .duration_since(std::time::UNIX_EPOCH)
4316 .map(|d| d.as_nanos() as u64)
4317 .unwrap_or(0);
4318 let mut z = nanos
4319 .wrapping_add(n.wrapping_mul(0x9E3779B97F4A7C15))
4320 .wrapping_add(0x9E3779B97F4A7C15);
4321 z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
4322 z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
4323 z ^= z >> 31;
4324 if z == 0 { 0x9E3779B97F4A7C15 } else { z }
4327}
4328
4329fn reject_unsupported(fields: &[(&str, bool, &str)]) -> Result<(), (String, String)> {
4334 for (param, present, why) in fields {
4335 if *present {
4336 return Err((format!("{param} is not supported{why}"), param.to_string()));
4337 }
4338 }
4339 Ok(())
4340}
4341
4342#[derive(PartialEq)]
4343enum ToolChoice {
4344 Auto,
4345 None,
4346}
4347
4348fn parse_tool_choice(v: &Option<serde_json::Value>) -> Result<ToolChoice, String> {
4349 match v {
4350 None | Some(serde_json::Value::Null) => Ok(ToolChoice::Auto),
4351 Some(serde_json::Value::String(s)) => match s.as_str() {
4352 "auto" => Ok(ToolChoice::Auto),
4353 "none" => Ok(ToolChoice::None),
4354 "required" => Err("tool_choice \"required\" is not supported (no constrained \
4355 decoding); use \"auto\""
4356 .into()),
4357 other => Err(format!("bad tool_choice {other:?} (auto|none)")),
4358 },
4359 Some(serde_json::Value::Object(_)) => {
4360 Err("named-function tool_choice is not supported; use \"auto\"".into())
4361 }
4362 Some(other) => Err(format!("bad tool_choice: {other}")),
4363 }
4364}
4365
4366fn parse_think(
4416 reasoning_effort: &Option<String>,
4417 reasoning: &Option<serde_json::Value>,
4418 vllm_switch: Option<bool>,
4419 suppress_switch: Option<bool>,
4420 default_effort: Option<&str>,
4421 max_tier: bool,
4422) -> Result<(ThinkMode, Option<String>, bool), String> {
4423 let mut effort = reasoning_effort.clone();
4424 let ReasoningObject {
4425 mut enabled,
4426 effort: object_effort,
4427 exclude,
4428 } = parse_reasoning_object(reasoning)?;
4429 if let Some(e) = object_effort {
4430 effort = Some(e);
4431 }
4432 match (enabled, vllm_switch) {
4437 (Some(a), Some(b)) if a != b => {
4438 return Err(format!(
4439 "contradictory reasoning switches: reasoning.enabled={a} and \
4440 enable_thinking={b} — send one"
4441 ));
4442 }
4443 (None, Some(b)) => enabled = Some(b),
4444 _ => {}
4445 }
4446 let suppress = match (exclude, suppress_switch) {
4458 (Some(true), _) | (_, Some(false)) => Some(false),
4459 _ => None,
4460 };
4461 match (enabled, suppress) {
4462 (Some(true), Some(false)) => {
4463 return Err(
4464 "contradictory reasoning switches: reasoning is enabled but \
4465 include_reasoning:false / reasoning.exclude:true asks for no reasoning — \
4466 on this server not delivering reasoning means not generating it, so send one"
4467 .into(),
4468 );
4469 }
4470 (None, Some(b)) => enabled = Some(b),
4471 _ => {}
4472 }
4473 let client_explicit = effort.is_some() || enabled.is_some();
4477 if effort.is_none() && enabled.is_none() {
4482 effort = default_effort.map(str::to_string);
4483 }
4484 let effort_arm = match effort.as_deref() {
4489 None => None,
4490 Some(raw) => {
4491 let level = canonical_effort_for(raw, max_tier).ok_or_else(|| {
4492 format!(
4493 "bad reasoning_effort {raw:?} \
4494 (none|minimal|low|medium|high; xhigh/max/ultra clamp to the \
4495 highest level this model's template distinguishes)"
4496 )
4497 })?;
4498 Some(match level {
4499 "none" | "minimal" => (ThinkMode::NoThink, "low"),
4500 "low" => (ThinkMode::Think, "low"),
4501 "medium" => (ThinkMode::Think, "medium"),
4502 "max" => (ThinkMode::Think, "max"),
4503 _ => (ThinkMode::Think, "high"),
4504 })
4505 }
4506 };
4507 let (think, level) = match (enabled, effort_arm) {
4508 (Some(false), _) => (ThinkMode::NoThink, Some("low".to_string())),
4511 (Some(true), arm) => (ThinkMode::Think, arm.map(|(_, level)| level.to_string())),
4512 (None, Some((think, level))) => (think, Some(level.to_string())),
4513 (None, None) => (ThinkMode::Default, None),
4514 };
4515 Ok((think, level, client_explicit))
4516}
4517
4518struct ReasoningObject {
4520 enabled: Option<bool>,
4521 effort: Option<String>,
4522 exclude: Option<bool>,
4523}
4524
4525fn parse_reasoning_object(
4544 reasoning: &Option<serde_json::Value>,
4545) -> Result<ReasoningObject, String> {
4546 let mut out = ReasoningObject {
4547 enabled: None,
4548 effort: None,
4549 exclude: None,
4550 };
4551 let Some(v) = reasoning else { return Ok(out) };
4552 let obj = match v {
4553 serde_json::Value::Null => return Ok(out),
4554 serde_json::Value::Object(obj) => obj,
4555 _ => return Err("reasoning must be an object".into()),
4556 };
4557 for (key, value) in obj {
4558 match key.as_str() {
4566 "enabled" => {
4567 if !value.is_null() {
4568 out.enabled = Some(
4569 value
4570 .as_bool()
4571 .ok_or("reasoning.enabled must be true or false")?,
4572 );
4573 }
4574 }
4575 "exclude" => {
4576 if !value.is_null() {
4577 out.exclude = Some(
4578 value
4579 .as_bool()
4580 .ok_or("reasoning.exclude must be true or false")?,
4581 );
4582 }
4583 }
4584 "effort" => {
4585 if !value.is_null() {
4586 out.effort = Some(
4587 value
4588 .as_str()
4589 .ok_or("reasoning.effort must be a string")?
4590 .to_string(),
4591 );
4592 }
4593 }
4594 "max_tokens" => {
4595 return Err(
4596 "reasoning.max_tokens is not supported by this server: reasoning tokens \
4597 are output tokens here, and max_tokens is the ONE output budget covering \
4598 reasoning and content together — there is no separate reasoning budget to \
4599 spend against, so honouring this field is impossible rather than merely \
4600 unimplemented. Use max_tokens for the budget, and reasoning.effort (or \
4601 reasoning.enabled:false) to spend less of it on reasoning"
4602 .into(),
4603 );
4604 }
4605 other => {
4606 return Err(format!(
4607 "reasoning.{other} is not a field this server implements (it would change \
4608 nothing about the request); the supported keys are enabled, effort and \
4609 exclude"
4610 ));
4611 }
4612 }
4613 }
4614 Ok(out)
4615}
4616
4617fn parse_template_kwargs(kwargs: &Option<serde_json::Value>) -> Result<Option<bool>, String> {
4638 let Some(v) = kwargs else { return Ok(None) };
4639 let obj = match v {
4640 serde_json::Value::Null => return Ok(None),
4641 serde_json::Value::Object(obj) => obj,
4642 _ => return Err("chat_template_kwargs must be an object".into()),
4643 };
4644 let mut switch = None;
4645 for (key, value) in obj {
4646 match key.as_str() {
4647 "enable_thinking" => {
4648 switch = Some(
4649 value
4650 .as_bool()
4651 .ok_or("chat_template_kwargs.enable_thinking must be true or false")?,
4652 );
4653 }
4654 "preserve_thinking" => {
4655 let preserve = value
4656 .as_bool()
4657 .ok_or("chat_template_kwargs.preserve_thinking must be true or false")?;
4658 if !preserve {
4659 return Err(
4660 "chat_template_kwargs.preserve_thinking:false is not supported by this \
4661 server: the renderer implements the vendor DEFAULT (replay every prior \
4662 assistant turn's <think> block, empty when no reasoning was sent) but \
4663 not the strip arm — serving replay bytes under a strip request would \
4664 misdescribe the prompt. Omit the flag or send true"
4665 .into(),
4666 );
4667 }
4668 }
4670 other => {
4671 return Err(format!(
4672 "chat_template_kwargs.{other} is not supported by this server's \
4673 template renderer (it would change nothing about the prompt); the only \
4674 supported key is enable_thinking (preserve_thinking is RECOGNISED but \
4675 refuses in both directions — see its own message)"
4676 ));
4677 }
4678 }
4679 }
4680 Ok(switch)
4681}
4682
4683fn resolve_vllm_think_switch(
4687 enable_thinking: Option<bool>,
4688 kwargs: &Option<serde_json::Value>,
4689) -> Result<Option<bool>, String> {
4690 let from_kwargs = parse_template_kwargs(kwargs)?;
4691 match (enable_thinking, from_kwargs) {
4692 (Some(a), Some(b)) if a != b => Err(format!(
4693 "contradictory reasoning switches: enable_thinking={a} and \
4694 chat_template_kwargs.enable_thinking={b} — send one"
4695 )),
4696 (Some(a), _) => Ok(Some(a)),
4697 (None, b) => Ok(b),
4698 }
4699}
4700
4701pub(crate) fn canonical_effort_for(value: &str, max_tier: bool) -> Option<&'static str> {
4721 match value {
4722 "none" => Some("none"),
4723 "minimal" => Some("minimal"),
4724 "low" => Some("low"),
4725 "medium" => Some("medium"),
4726 "high" => Some("high"),
4727 "xhigh" | "max" | "ultra" => Some(if max_tier { "max" } else { "high" }),
4733 _ => None,
4734 }
4735}
4736
4737pub(crate) fn canonical_effort(value: &str) -> Option<&'static str> {
4740 canonical_effort_for(value, false)
4741}
4742
4743fn json_to_val(v: &serde_json::Value) -> chat::Val {
4746 match v {
4747 serde_json::Value::Null => chat::Val::Null,
4748 serde_json::Value::Bool(b) => chat::Val::Bool(*b),
4749 serde_json::Value::Number(n) => chat::Val::Num(n.to_string()),
4750 serde_json::Value::String(s) => chat::Val::Str(s.clone()),
4751 serde_json::Value::Array(a) => chat::Val::Arr(a.iter().map(json_to_val).collect()),
4752 serde_json::Value::Object(o) => chat::Val::Obj(
4755 o.iter()
4756 .map(|(k, val)| (k.clone(), json_to_val(val)))
4757 .collect(),
4758 ),
4759 }
4760}
4761
4762#[allow(clippy::type_complexity)]
4766fn prepare_tools(
4767 tools: &[serde_json::Value],
4768) -> Result<
4769 (
4770 Vec<String>,
4771 Vec<chat::Val>,
4772 HashMap<String, HashMap<String, String>>,
4773 ),
4774 String,
4775> {
4776 let mut tools_json = Vec::with_capacity(tools.len());
4777 let mut tools_struct = Vec::with_capacity(tools.len());
4778 let mut schemas: HashMap<String, HashMap<String, String>> = HashMap::new();
4779 for t in tools {
4780 let f = t
4781 .get("function")
4782 .ok_or("each tool needs a function object")?;
4783 let name = f
4784 .get("name")
4785 .and_then(|n| n.as_str())
4786 .ok_or("each tool needs function.name")?;
4787 let mut params: HashMap<String, String> = HashMap::new();
4788 if let Some(props) = f
4789 .get("parameters")
4790 .and_then(|p| p.get("properties"))
4791 .and_then(|p| p.as_object())
4792 {
4793 for (p, def) in props {
4794 if let Some(ty) = def.get("type").and_then(|t| t.as_str()) {
4795 params.insert(p.clone(), ty.to_string());
4796 }
4797 }
4798 }
4799 schemas.insert(name.to_string(), params);
4800 tools_json.push(pyjson_str(t));
4801 tools_struct.push(json_to_val(f));
4803 }
4804 Ok((tools_json, tools_struct, schemas))
4805}
4806
4807fn render_req_tool_call(tc: &ReqToolCall) -> Result<TmplToolCall, String> {
4812 let parsed: serde_json::Value = match &tc.function.arguments {
4813 serde_json::Value::Null => json!({}),
4814 serde_json::Value::String(s) if s.trim().is_empty() => json!({}),
4815 serde_json::Value::String(s) => serde_json::from_str(s)
4816 .map_err(|e| format!("tool_calls arguments is not valid JSON: {e}"))?,
4817 v @ serde_json::Value::Object(_) => v.clone(),
4818 _ => return Err("tool_calls arguments must be a JSON object".into()),
4819 };
4820 let obj = parsed
4821 .as_object()
4822 .ok_or("tool_calls arguments must decode to a JSON object")?;
4823 let params = obj
4824 .iter()
4825 .map(|(k, v)| {
4826 let rendered = match v {
4827 serde_json::Value::String(s) => s.clone(),
4828 v @ (serde_json::Value::Object(_) | serde_json::Value::Array(_)) => pyjson_str(v),
4829 scalar => scalar.to_string(),
4830 };
4831 (k.clone(), rendered)
4832 })
4833 .collect();
4834 let args = obj
4837 .iter()
4838 .map(|(k, v)| (k.clone(), json_to_val(v)))
4839 .collect();
4840 Ok(TmplToolCall {
4841 name: tc.function.name.clone(),
4842 params,
4843 args,
4844 id: tc.id.clone(),
4845 })
4846}
4847
4848fn tool_call_json(c: &ParsedToolCall) -> serde_json::Value {
4850 json!({ "id": c.id, "type": "function",
4851 "function": { "name": c.name, "arguments": c.arguments } })
4852}
4853
4854async fn serve_bounded_http_with_limits<F>(
4858 listener: tokio::net::TcpListener,
4859 app: Router,
4860 shutdown: F,
4861 header_read_timeout: std::time::Duration,
4862 max_connections: usize,
4863 connection_max_lifetime: std::time::Duration,
4864) -> std::io::Result<()>
4865where
4866 F: std::future::Future<Output = ()> + Send,
4867{
4868 let connections = Arc::new(tokio::sync::Semaphore::new(max_connections));
4869 let (connection_shutdown, _) = tokio::sync::watch::channel(false);
4870 let mut connection_tasks = tokio::task::JoinSet::new();
4871 let mut shutdown = Box::pin(shutdown);
4872
4873 loop {
4874 tokio::select! {
4875 _ = &mut shutdown => break,
4876 joined = connection_tasks.join_next(), if !connection_tasks.is_empty() => {
4877 if let Some(Err(error)) = joined {
4878 eprintln!("[server] connection task failed: {error}");
4879 }
4880 }
4881 accepted = listener.accept() => {
4882 let (stream, _) = match accepted {
4883 Ok(connection) => connection,
4884 Err(error) => {
4885 eprintln!("[server] accept failed: {error}");
4886 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
4887 continue;
4888 }
4889 };
4890 let permit = match connections.clone().try_acquire_owned() {
4891 Ok(permit) => permit,
4892 Err(_) => {
4893 drop(stream);
4894 continue;
4895 }
4896 };
4897 let service = app.clone().map_request(
4898 |request: hyper::Request<hyper::body::Incoming>| request.map(Body::new),
4899 );
4900 let service = hyper_util::service::TowerToHyperService::new(service);
4901 let io = hyper_util::rt::TokioIo::new(stream);
4902 let mut builder = hyper_util::server::conn::auto::Builder::new(
4903 hyper_util::rt::TokioExecutor::new(),
4904 );
4905 builder
4906 .http1()
4907 .timer(hyper_util::rt::TokioTimer::new())
4908 .header_read_timeout(header_read_timeout)
4909 .max_headers(64);
4910 builder
4911 .http2()
4912 .timer(hyper_util::rt::TokioTimer::new())
4913 .max_concurrent_streams(MAX_HTTP2_STREAMS_PER_CONNECTION)
4914 .keep_alive_interval(Some(std::time::Duration::from_secs(30)))
4915 .keep_alive_timeout(std::time::Duration::from_secs(10));
4916 let mut connection = Box::pin(builder
4917 .serve_connection_with_upgrades(io, service)
4918 .into_owned());
4919 let mut shutdown_rx = connection_shutdown.subscribe();
4920 connection_tasks.spawn(async move {
4921 let _permit = permit;
4922 tokio::select! {
4923 result = connection.as_mut() => {
4924 let _ = result;
4925 }
4926 _ = tokio::time::sleep(connection_max_lifetime) => {
4927 connection.as_mut().graceful_shutdown();
4932 let _ = connection.await;
4933 }
4934 _ = shutdown_rx.changed() => {
4935 connection.as_mut().graceful_shutdown();
4936 let _ = connection.await;
4937 }
4938 }
4939 });
4940 }
4941 }
4942 }
4943 drop(listener);
4944 let _ = connection_shutdown.send(true);
4945 let drained = tokio::time::timeout(std::time::Duration::from_secs(5), async {
4946 while connection_tasks.join_next().await.is_some() {}
4947 })
4948 .await;
4949 if drained.is_err() {
4950 connection_tasks.abort_all();
4951 eprintln!("[server] WARN: HTTP connections exceeded the 5s graceful close deadline");
4952 }
4953 Ok(())
4954}
4955
4956async fn serve_bounded_http<F>(
4957 listener: tokio::net::TcpListener,
4958 app: Router,
4959 shutdown: F,
4960) -> std::io::Result<()>
4961where
4962 F: std::future::Future<Output = ()> + Send,
4963{
4964 serve_bounded_http_with_limits(
4965 listener,
4966 app,
4967 shutdown,
4968 HTTP1_HEADER_READ_TIMEOUT,
4969 MAX_HTTP_CONNECTIONS,
4970 HTTP_CONNECTION_MAX_LIFETIME,
4971 )
4972 .await
4973}
4974
4975#[tokio::main]
4976pub async fn serve_main() -> Result<(), Box<dyn std::error::Error>> {
4977 serve_with(ServerWiring::stock()).await
4978}
4979
4980enum MeteringWiring {
4982 Stock,
4986 Custom(metering::MeteringFactory),
4991}
4992
4993pub struct ServerWiring {
4997 metering: MeteringWiring,
4998 on_ready: Option<Box<dyn FnOnce(RuntimeHandles) + Send>>,
5001 claimed_env: Vec<&'static str>,
5005}
5006
5007impl ServerWiring {
5008 pub fn stock() -> Self {
5010 ServerWiring {
5011 metering: MeteringWiring::Stock,
5012 on_ready: None,
5013 claimed_env: Vec::new(),
5014 }
5015 }
5016
5017 pub fn with_metering(factory: metering::MeteringFactory) -> Self {
5020 ServerWiring {
5021 metering: MeteringWiring::Custom(factory),
5022 on_ready: None,
5023 claimed_env: Vec::new(),
5024 }
5025 }
5026
5027 pub fn claiming(mut self, var: &'static str) -> Self {
5031 self.claimed_env.push(var);
5032 self
5033 }
5034
5035 pub fn on_ready(mut self, hook: impl FnOnce(RuntimeHandles) + Send + 'static) -> Self {
5036 self.on_ready = Some(Box::new(hook));
5037 self
5038 }
5039}
5040
5041pub struct RuntimeHandles {
5044 pub trim: TrimHandle,
5045 pub purge: PurgeHandle,
5048 pub kv_handoff: HostHandoffHandle,
5054 pub shutdown: tokio::sync::watch::Receiver<bool>,
5059}
5060
5061#[derive(Clone)]
5064pub struct TrimHandle {
5065 cmd_tx: Sender<Cmd>,
5066}
5067
5068impl TrimHandle {
5069 pub async fn trim(&self) -> Result<serde_json::Value, String> {
5071 let (tx, rx) = tokio::sync::oneshot::channel();
5072 if self.cmd_tx.send(Cmd::TrimPools(tx)).is_err() {
5073 return Err("worker is down".into());
5074 }
5075 match tokio::time::timeout(std::time::Duration::from_secs(30), rx).await {
5076 Ok(Ok(report)) => Ok(json!(report)),
5077 _ => Err("worker did not answer the trim within 30s".into()),
5078 }
5079 }
5080}
5081
5082#[derive(Clone)]
5091pub struct PurgeHandle {
5092 cmd_tx: Sender<Cmd>,
5093}
5094
5095impl PurgeHandle {
5096 pub async fn purge_tenant(&self, tenant: &str) -> Result<serde_json::Value, String> {
5098 let (tx, rx) = tokio::sync::oneshot::channel();
5099 let cmd = Cmd::PurgeTenantHost {
5100 tenant: tenant.to_string(),
5101 tx,
5102 };
5103 if self.cmd_tx.send(cmd).is_err() {
5104 return Err("worker is down".into());
5105 }
5106 match tokio::time::timeout(std::time::Duration::from_secs(30), rx).await {
5107 Ok(Ok(report)) => Ok(json!(report)),
5108 _ => Err("worker did not answer the purge within 30s".into()),
5109 }
5110 }
5111}
5112
5113#[derive(Clone)]
5122pub struct HostHandoffHandle {
5123 cmd_tx: Sender<Cmd>,
5124}
5125
5126impl HostHandoffHandle {
5127 pub async fn export(&self, force: bool) -> Result<serde_json::Value, String> {
5130 let (tx, rx) = tokio::sync::oneshot::channel();
5131 if self
5132 .cmd_tx
5133 .send(Cmd::ExportHostHandoff { force, tx })
5134 .is_err()
5135 {
5136 return Err("worker is down".into());
5137 }
5138 match tokio::time::timeout(std::time::Duration::from_secs(900), rx).await {
5139 Ok(Ok(Ok(report))) => Ok(json!(report)),
5140 Ok(Ok(Err(refused))) => Err(refused),
5141 _ => Err("worker did not answer the export within 900s".into()),
5142 }
5143 }
5144
5145 pub async fn import(&self) -> Result<serde_json::Value, String> {
5148 let (tx, rx) = tokio::sync::oneshot::channel();
5149 if self.cmd_tx.send(Cmd::ImportHostHandoff { tx }).is_err() {
5150 return Err("worker is down".into());
5151 }
5152 match tokio::time::timeout(std::time::Duration::from_secs(30), rx).await {
5153 Ok(Ok(Ok(start))) => Ok(json!(start)),
5154 Ok(Ok(Err(refused))) => Err(refused),
5155 _ => Err("worker did not answer the import within 30s".into()),
5156 }
5157 }
5158}
5159
5160pub async fn serve_with(wiring: ServerWiring) -> Result<(), Box<dyn std::error::Error>> {
5161 let args: Vec<String> = std::env::args().skip(1).collect();
5164 if args.iter().any(|a| a == "--version" || a == "-V") {
5169 println!("memra-server {}", env!("CARGO_PKG_VERSION"));
5170 println!("system_fingerprint {SYSTEM_FINGERPRINT}");
5171 println!("build_id_src {BUILD_ID_SRC}");
5172 println!("git_sha {BUILD_GIT_SHA}");
5173 if !BUILD_ID_NOTE.is_empty() {
5174 println!("degraded {BUILD_ID_NOTE}");
5175 }
5176 return Ok(());
5177 }
5178 if let Some(code) = auth::run_cli(&args) {
5179 std::process::exit(code);
5180 }
5181 eprintln!("{}", build_identity_line());
5185 if BUILD_ID_SRC != build_id::BUILD_ID_SRC_TREE {
5186 eprintln!(
5187 "[server] WARNING: build identity is DEGRADED: {BUILD_ID_NOTE}. \
5188 system_fingerprint {SYSTEM_FINGERPRINT} carries a version-only id, so it does \
5189 NOT identify the source this binary was compiled from and published \
5190 performance pins cannot be verified against it (darklanes \
5191 tools/check-claim-builds.mjs --live). Rebuild where the workspace source tree \
5192 is readable."
5193 );
5194 }
5195 auth::init_from_env();
5198 let api_auth = match ApiAuth::from_env() {
5199 Ok(auth) => auth,
5200 Err(err) => {
5201 eprintln!("[server] FATAL: {err}");
5202 std::process::exit(1);
5203 }
5204 };
5205 let addr = std::env::var("MEMRA_ADDR").unwrap_or_else(|_| "127.0.0.1:8080".into());
5206 let allow_open_bind = std::env::var("MEMRA_ALLOW_OPEN_BIND").as_deref() == Ok("1");
5207 let (bind_addr, bind_loopback) = match resolve_bind_addr(&addr) {
5208 Ok(resolved) => resolved,
5209 Err(err) => {
5210 eprintln!("[server] FATAL: {err}");
5211 std::process::exit(1);
5212 }
5213 };
5214 if let Err(message) = validate_bind_security(&addr, api_auth.configured(), allow_open_bind) {
5219 eprintln!("[server] FATAL: {message}");
5220 std::process::exit(1);
5221 }
5222 if !bind_loopback && !api_auth.configured() {
5223 eprintln!(
5224 "[server] WARNING: MEMRA_ALLOW_OPEN_BIND=1 permits open completion routes on {addr}; \
5225 metrics remain bearer-protected"
5226 );
5227 }
5228 let metrics_token = match std::env::var("MEMRA_METRICS_TOKEN") {
5229 Ok(token) if token.is_empty() => {
5230 eprintln!("[server] FATAL: MEMRA_METRICS_TOKEN must not be empty");
5231 std::process::exit(1);
5232 }
5233 Ok(token) => Some(token),
5234 Err(std::env::VarError::NotPresent) => None,
5235 Err(std::env::VarError::NotUnicode(_)) => {
5236 eprintln!("[server] FATAL: MEMRA_METRICS_TOKEN must be valid UTF-8");
5237 std::process::exit(1);
5238 }
5239 };
5240 let metrics_auth = MetricsAuth::new(bind_loopback, api_auth.configured(), metrics_token);
5241
5242 let models = parse_models_config();
5243 let (openrouter_metadata, provider_metadata) = match load_openrouter_metadata(&models) {
5244 Ok(loaded) => loaded,
5245 Err(err) => {
5246 eprintln!("[server] FATAL: {err}");
5247 std::process::exit(1);
5248 }
5249 };
5250 let metering_obj: Option<Arc<dyn metering::Metering>> = {
5256 let factory = match wiring.metering {
5257 MeteringWiring::Stock => None,
5258 MeteringWiring::Custom(factory) => Some(factory),
5259 };
5260 for deployment_only in [
5261 "MEMRA_REQUEST_LEDGER",
5262 "MEMRA_TENANT_BUDGETS",
5263 "MEMRA_ADMIN_ADDR",
5264 "MEMRA_ADMIN_TOKEN_FILE",
5265 "MEMRA_CAPTURE_DIR",
5266 ] {
5267 if std::env::var_os(deployment_only).is_some()
5268 && !wiring.claimed_env.contains(&deployment_only)
5269 {
5270 eprintln!(
5271 "[server] FATAL: {deployment_only} is a deployment-binary surface; this \
5272 build ships no accounting/admin/capture. Wire a Metering implementation \
5273 through ServerWiring and claim the vars it consumes."
5274 );
5275 std::process::exit(1);
5276 }
5277 }
5278 match factory {
5279 None => None,
5280 Some(factory) => {
5281 let model_ids: Vec<String> =
5282 models.iter().map(|(name, _, _)| name.clone()).collect();
5283 match factory(&metering::MeteringInit { models: &model_ids }) {
5284 Ok(metering_obj) => metering_obj,
5285 Err(err) => {
5286 eprintln!("[server] FATAL: metering wiring: {err}");
5287 std::process::exit(1);
5288 }
5289 }
5290 }
5291 }
5292 };
5293 let budget_tokenizers = if metering_obj
5294 .as_ref()
5295 .is_some_and(|manager| manager.enforces_limits())
5296 {
5297 match load_budget_tokenizers(&models) {
5298 Ok(tokenizers) => Some(tokenizers),
5299 Err(err) => {
5300 eprintln!("[server] FATAL: prepaid reservation tokenizers: {err}");
5301 std::process::exit(1);
5302 }
5303 }
5304 } else {
5305 None
5306 };
5307 eprintln!("[server] starting; models config = {models:?}");
5308
5309 let health_state = health::WorkerHealth::new();
5314 health::spawn_gpu_watch(health_state.clone());
5318 health::spawn_sd_watchdog(health_state.clone());
5319
5320 let (cmd_tx, model_names, caps, metrics, worker_thread) =
5322 match worker::spawn(models, health_state.clone()) {
5323 Ok(v) => v,
5324 Err(err) => {
5325 eprintln!("[server] FATAL: worker init failed: {err}");
5326 health_state.mark_dead(format!("worker init failed: {err}"));
5327 health::sd_notify(&format!("STATUS=worker init failed: {err}"));
5328 std::process::exit(1);
5329 }
5330 };
5331 eprintln!("[server] worker ready; serving models: {model_names:?}");
5332
5333 let (drain_shutdown_tx, drain_shutdown_rx) = tokio::sync::watch::channel(false);
5340 if let Some(on_ready) = wiring.on_ready {
5341 on_ready(RuntimeHandles {
5342 trim: TrimHandle {
5343 cmd_tx: cmd_tx.clone(),
5344 },
5345 purge: PurgeHandle {
5346 cmd_tx: cmd_tx.clone(),
5347 },
5348 kv_handoff: HostHandoffHandle {
5349 cmd_tx: cmd_tx.clone(),
5350 },
5351 shutdown: drain_shutdown_rx.clone(),
5352 });
5353 }
5354
5355 let bg_handle = darklane::spawn_from_env(health_state.clone());
5358 let bg_state = bg_handle.as_ref().map(|h| {
5359 let mode = darklane::BgConfig::from_env()
5360 .map(|c| c.yield_mode.as_str())
5361 .unwrap_or("stop");
5362 (h.state.clone(), mode)
5363 });
5364
5365 let state = AppState {
5366 cmd_tx,
5367 models: model_names,
5368 caps,
5369 openrouter_metadata: Arc::new(openrouter_metadata),
5370 provider_metadata: Arc::new(provider_metadata),
5371 metering: metering_obj,
5372 budget_tokenizers,
5373 api_auth,
5374 metrics_auth,
5375 metrics,
5376 inflight: Arc::new(Default::default()),
5377 tenant_inflight: Arc::new(Default::default()),
5378 health: health_state.clone(),
5379 bg: bg_state,
5380 };
5381 let inflight_handle = state.inflight.clone();
5382 let drain_metering = state.metering.clone();
5385 worker::register_http_inflight(state.inflight.clone());
5390 let app = Router::new()
5391 .route("/health", get(health_live))
5396 .route("/livez", get(health_live))
5397 .route("/readyz", get(health_ready))
5398 .route("/models", get(list_models))
5399 .route("/v1/models", get(list_models_v1))
5400 .route("/v1/auth/check", get(auth_check))
5401 .route("/v1/completions", post(completions_admitted))
5402 .route("/v1/embeddings", post(embed_api::embeddings_admitted))
5403 .route("/v1/rerank", post(embed_api::rerank_admitted))
5404 .route("/v1/chat/completions", post(chat_completions_admitted))
5405 .route("/v1/messages", post(anthropic::messages_admitted))
5409 .route("/v1/responses", post(responses_api::responses_admitted))
5410 .route("/metrics", get(get_metrics))
5411 .route("/yield/metrics", get(yield_metrics))
5412 .with_state(state.clone());
5413 let app = apply_body_limit(app);
5416 let app = app.layer(middleware::from_fn_with_state(
5419 state,
5420 authenticate_inference_before_body,
5421 ));
5422 let app = if ttft::enabled() {
5423 app.layer(middleware::from_fn(ttft_request_start))
5424 } else {
5425 app
5426 };
5427
5428 let listener = tokio::net::TcpListener::bind(bind_addr).await?;
5429 eprintln!("[server] listening on http://{bind_addr}");
5430 drop(drain_shutdown_rx);
5431 health::sd_notify("READY=1\nSTATUS=serving");
5435 let inflight = inflight_handle;
5442 let signal_admin_shutdown = drain_shutdown_tx.clone();
5443 let serve_result = serve_bounded_http(listener, app, async move {
5444 let mut sigterm =
5445 match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
5446 Ok(s) => s,
5447 Err(err) => {
5448 eprintln!("[server] WARN: no SIGTERM handler ({err}); drain disabled");
5449 std::future::pending::<()>().await;
5450 unreachable!()
5451 }
5452 };
5453 sigterm.recv().await;
5454 DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
5455 let _ = signal_admin_shutdown.send(true);
5456 health::sd_notify(&format!(
5460 "STOPPING=1\nSTATUS=draining\nEXTEND_TIMEOUT_USEC={}",
5461 (drain_deadline_s() + 5) * 1_000_000
5462 ));
5463 let n: usize = inflight
5464 .iter()
5465 .map(|c| c.load(std::sync::atomic::Ordering::SeqCst))
5466 .sum();
5467 eprintln!(
5468 "[server] SIGTERM: draining ({n} in flight, deadline {}s)",
5469 drain_deadline_s()
5470 );
5471 let deadline = std::time::Duration::from_secs(drain_deadline_s());
5472 let t0 = std::time::Instant::now();
5473 loop {
5474 let n: usize = inflight
5475 .iter()
5476 .map(|c| c.load(std::sync::atomic::Ordering::SeqCst))
5477 .sum();
5478 if n == 0 {
5479 eprintln!(
5480 "[server] drain complete in {:.1}s; exiting",
5481 t0.elapsed().as_secs_f64()
5482 );
5483 break;
5484 }
5485 if t0.elapsed() >= deadline {
5486 eprintln!(
5487 "[server] drain deadline ({}s) hit with {n} in flight; exiting",
5488 drain_deadline_s()
5489 );
5490 if let Some(metering) = drain_metering.as_ref() {
5497 metering.drain_kill();
5498 }
5499 break;
5500 }
5501 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
5502 }
5503 })
5504 .await;
5505 let _ = drain_shutdown_tx.send(true);
5508 serve_result?;
5509 if let Some(h) = bg_handle {
5513 h.shutdown();
5514 }
5515 worker_thread.join().map_err(|_| {
5524 std::io::Error::other("GPU worker thread panicked during graceful shutdown")
5525 })?;
5526 eprintln!("[server] GPU worker shutdown complete");
5527 Ok(())
5528}
5529
5530fn validate_model_path(path: &str) -> Result<(), String> {
5536 let p = std::path::Path::new(path);
5537 if !p.exists() {
5538 return Err(format!("model path {path:?} does not exist"));
5539 }
5540 if p.is_file() {
5541 return Ok(()); }
5543 if p.join("manifest.json").exists() {
5544 return Ok(()); }
5546 let has_st =
5547 p.join("model.safetensors").exists() || p.join("model.safetensors.index.json").exists();
5548 if !has_st {
5549 return Err(format!(
5550 "model dir {path:?} is not a servable checkpoint: want model.safetensors or \
5551 model.safetensors.index.json + config.json (HF safetensors dir), or \
5552 manifest.json (memra repack dir)"
5553 ));
5554 }
5555 if !p.join("config.json").exists() {
5556 return Err(format!(
5557 "model dir {path:?} has safetensors weights but no config.json"
5558 ));
5559 }
5560 Ok(())
5561}
5562
5563fn parse_models_config() -> Vec<(String, String, Option<String>)> {
5571 if let Ok(spec) = std::env::var("MEMRA_MODELS") {
5572 let mut out = Vec::new();
5573 for entry in spec.split(',').filter(|s| !s.trim().is_empty()) {
5574 if let Some((name, path)) = entry.split_once('=') {
5575 let (mpath, dpath) = match path.trim().split_once('+') {
5578 Some((m, d)) => (m.trim(), Some(d.trim())),
5579 None => (path.trim(), None),
5580 };
5581 let resolve = |p: &str| {
5582 memra_gguf::hf::resolve_arg(p).unwrap_or_else(|err| {
5583 eprintln!("[server] FATAL: model {name:?}: {err}");
5584 std::process::exit(1);
5585 })
5586 };
5587 let mpath = resolve(mpath);
5588 if let Err(err) = validate_model_path(&mpath) {
5589 eprintln!("[server] FATAL: model {name:?}: {err}");
5590 std::process::exit(1);
5591 }
5592 let dpath = dpath.map(|d| {
5602 let d = resolve(d);
5603 let p = std::path::Path::new(&d);
5604 if !p.exists() {
5605 eprintln!(
5606 "[server] FATAL: model {name:?}: drafter path {d:?} does not \
5607 exist (MEMRA_MODELS '+draft' attach). Refusing to start \
5608 rather than serving plain decode under a config that asked \
5609 for speculative decoding."
5610 );
5611 std::process::exit(1);
5612 }
5613 if !p.is_file() {
5614 eprintln!(
5615 "[server] FATAL: model {name:?}: drafter path {d:?} is not a \
5616 file — a '+draft' attach must be a NextN/MTP GGUF file."
5617 );
5618 std::process::exit(1);
5619 }
5620 d
5621 });
5622 out.push((name.trim().to_string(), mpath, dpath));
5623 } else {
5624 eprintln!(
5625 "[server] WARN: bad MEMRA_MODELS entry {entry:?} (want name=/path[+/draft]); skipping"
5626 );
5627 }
5628 }
5629 if !out.is_empty() {
5630 return out;
5631 }
5632 }
5633 vec![
5635 (
5636 "main".into(),
5637 "/data/ai-ml/hf-models/qwen36-27b-nvfp4-mtp/Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf".into(),
5638 None,
5639 ),
5640 (
5641 "judge".into(),
5642 "/data/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf".into(),
5643 None,
5644 ),
5645 ]
5646}
5647
5648fn load_budget_tokenizers(
5649 models: &[(String, String, Option<String>)],
5650) -> Result<Arc<HashMap<String, Arc<Tokenizer>>>, String> {
5651 let mut tokenizers = HashMap::new();
5652 for (alias, path, _) in models {
5653 let path = std::path::Path::new(path);
5654 let tokenizer = if path.is_dir() {
5655 let tokenizer_dir = if path.join("manifest.json").exists() {
5656 let repack = memra_gguf::source::Hy3RepackSource::open(path).map_err(|err| {
5657 format!("model {alias:?}: open repack tokenizer source: {err}")
5658 })?;
5659 repack
5660 .source_dir()
5661 .filter(|source| source.join("tokenizer.json").exists())
5662 .unwrap_or(path)
5663 .to_path_buf()
5664 } else {
5665 path.to_path_buf()
5666 };
5667 Tokenizer::from_hf_dir(&tokenizer_dir)
5668 .map_err(|err| format!("model {alias:?}: reservation tokenizer: {err}"))?
5669 } else {
5670 let gguf = memra_gguf::GgufFile::open(path)
5671 .map_err(|err| format!("model {alias:?}: open reservation tokenizer: {err}"))?;
5672 Tokenizer::from_gguf(&gguf)
5673 .map_err(|err| format!("model {alias:?}: reservation tokenizer: {err}"))?
5674 };
5675 tokenizers.insert(alias.clone(), Arc::new(tokenizer));
5676 }
5677 Ok(Arc::new(tokenizers))
5678}
5679
5680fn health_payload(st: &AppState, status: &str, detail: Option<&str>) -> serde_json::Value {
5682 let s = st.health.snapshot();
5683 let mut v = json!({
5684 "status": status,
5685 "models": *st.models,
5686 "worker": {
5687 "phase": health::phase_name(s.phase),
5688 "beat_age_ms": s.beat_age_ms,
5689 "tick_max_ms": s.tick_max_ms,
5690 "stall_threshold_ms": s.stall_threshold_ms,
5691 "generation": s.generation,
5692 "xid_warnings": s.xid_warns,
5693 },
5694 });
5695 if let Some(d) = detail {
5696 v["detail"] = json!(d);
5697 }
5698 v
5699}
5700
5701fn readiness_payload(st: &AppState, status: &str, detail: Option<&str>) -> serde_json::Value {
5705 let mut v = health_payload(st, status, detail);
5706 v["peer_probe_integrity"] = json!(st.health.peer_probe_integrity().detail());
5707 v
5708}
5709
5710async fn auth_check() -> impl IntoResponse {
5714 StatusCode::NO_CONTENT
5715}
5716
5717async fn health_live(State(st): State<AppState>) -> impl IntoResponse {
5733 if draining() {
5734 return (StatusCode::OK, Json(health_payload(&st, "draining", None))).into_response();
5737 }
5738 match st.health.live() {
5739 Ok(()) => (StatusCode::OK, Json(health_payload(&st, "ok", None))).into_response(),
5740 Err(why) => retry_contract_response(
5741 (
5742 StatusCode::SERVICE_UNAVAILABLE,
5743 Json(health_payload(&st, "unhealthy", Some(&why))),
5744 )
5745 .into_response(),
5746 Some(worker::WORKER_RESPAWN_BACKOFF_BASE_S),
5747 ),
5748 }
5749}
5750
5751async fn health_ready(State(st): State<AppState>) -> impl IntoResponse {
5762 let is_draining = draining();
5763 match st.health.ready(is_draining) {
5764 Ok(()) => (StatusCode::OK, Json(readiness_payload(&st, "ready", None))).into_response(),
5765 Err(why) => retry_contract_response(
5766 (
5767 StatusCode::SERVICE_UNAVAILABLE,
5768 Json(readiness_payload(&st, "not_ready", Some(&why))),
5769 )
5770 .into_response(),
5771 Some(if is_draining {
5772 drain_deadline_s()
5773 } else {
5774 worker::WORKER_RESPAWN_BACKOFF_BASE_S
5775 }),
5776 ),
5777 }
5778}
5779
5780#[derive(Clone, Copy)]
5781struct DualPpMetricsSnapshot {
5782 stage_ns: [u64; 4],
5783 stage_samples: [usize; 4],
5784 dropped_timing_samples: usize,
5785 overlaps: usize,
5786 slot_pairs: usize,
5787 slot_uses: [usize; 2],
5788 slot_collisions: usize,
5789}
5790
5791impl DualPpMetricsSnapshot {
5792 fn current() -> Self {
5793 let (stage_ns, stage_samples) = memra_engine::pp::dual_pp_timing_snapshot();
5794 let (slot_pairs, slot_uses, slot_collisions) = memra_engine::pp::dual_pp_slot_snapshot();
5795 Self {
5796 stage_ns,
5797 stage_samples,
5798 dropped_timing_samples: memra_engine::pp::dual_pp_timing_dropped(),
5799 overlaps: memra_engine::pp::dual_pp_overlaps(),
5800 slot_pairs,
5801 slot_uses,
5802 slot_collisions,
5803 }
5804 }
5805
5806 fn populated(self) -> bool {
5807 self.stage_samples.iter().any(|&n| n > 0)
5808 || self.dropped_timing_samples > 0
5809 || self.slot_pairs > 0
5810 || self.slot_collisions > 0
5811 }
5812}
5813
5814fn insert_dual_pp_metrics(
5815 body: &mut serde_json::Value,
5816 metrics_scope: &MetricsScope,
5817 snapshot: impl FnOnce() -> DualPpMetricsSnapshot,
5818) {
5819 if !metrics_scope.operator() {
5822 return;
5823 }
5824 let snapshot = snapshot();
5825 if !snapshot.populated() {
5826 return;
5827 }
5828 let timings: serde_json::Map<String, serde_json::Value> = memra_engine::pp::DUAL_PP_STAGE_NAMES
5829 .iter()
5830 .enumerate()
5831 .map(|(i, name)| {
5832 let total_ms = snapshot.stage_ns[i] as f64 / 1_000_000.0;
5833 (
5834 name.to_string(),
5835 json!({
5836 "samples": snapshot.stage_samples[i],
5837 "total_ms": total_ms,
5838 "mean_ms": if snapshot.stage_samples[i] > 0 {
5839 total_ms / snapshot.stage_samples[i] as f64
5840 } else { 0.0 },
5841 }),
5842 )
5843 })
5844 .collect();
5845 body["dual_pp"] = json!({
5846 "overlaps": snapshot.overlaps,
5847 "slot_pairs": snapshot.slot_pairs,
5848 "slot_uses": snapshot.slot_uses,
5849 "slot_collisions": snapshot.slot_collisions,
5850 "cuda_event_spans": timings,
5851 "dropped_timing_samples": snapshot.dropped_timing_samples,
5852 });
5853}
5854
5855#[derive(Clone, Copy)]
5856struct PpWaveMetricsSnapshot {
5857 ticks: usize,
5858 cells: usize,
5859 overlaps: usize,
5860}
5861
5862impl PpWaveMetricsSnapshot {
5863 fn current() -> Self {
5864 let (ticks, cells, overlaps) = memra_engine::pp::pp_wave_snapshot();
5865 Self {
5866 ticks,
5867 cells,
5868 overlaps,
5869 }
5870 }
5871}
5872
5873fn insert_pp_wave_metrics(
5874 body: &mut serde_json::Value,
5875 metrics_scope: &MetricsScope,
5876 snapshot: impl FnOnce() -> PpWaveMetricsSnapshot,
5877) {
5878 if !metrics_scope.operator() {
5879 return;
5880 }
5881 let snapshot = snapshot();
5882 if snapshot.ticks == 0 && snapshot.cells == 0 {
5883 return;
5884 }
5885 body["pp_wave"] = json!({
5886 "ticks": snapshot.ticks,
5887 "cells": snapshot.cells,
5888 "overlaps": snapshot.overlaps,
5889 });
5890}
5891
5892fn insert_spec_acceptance_metrics(
5893 body: &mut serde_json::Value,
5894 metrics_scope: &MetricsScope,
5895 snapshot: impl FnOnce() -> HashMap<String, memra_engine::spec::SpecTelemetry>,
5896) {
5897 if !metrics_scope.operator() {
5900 return;
5901 }
5902 let snapshot = snapshot();
5903 if snapshot.is_empty() {
5904 return;
5905 }
5906
5907 let mut tau = serde_json::Map::new();
5908 let mut by_position = serde_json::Map::new();
5909 for (model, telemetry) in snapshot {
5910 if telemetry.rounds == 0 {
5911 continue;
5912 }
5913 let n_pos = telemetry
5914 .pos_drafted
5915 .iter()
5916 .rposition(|&n| n > 0)
5917 .map_or(0, |position| position + 1);
5918 tau.insert(model.clone(), json!(telemetry.tau()));
5919 by_position.insert(
5920 model,
5921 json!({
5922 "window_seconds": worker::SPEC_METRICS_WINDOW_S,
5923 "rounds": telemetry.rounds,
5924 "offered": telemetry.pos_drafted[..n_pos].to_vec(),
5925 "accepted": telemetry.pos_accepted[..n_pos].to_vec(),
5926 "accept_rate": (0..n_pos).map(|position| {
5927 let offered = telemetry.pos_drafted[position];
5928 if offered > 0 {
5929 telemetry.pos_accepted[position] as f64 / offered as f64
5930 } else {
5931 0.0
5932 }
5933 }).collect::<Vec<f64>>(),
5934 }),
5935 );
5936 }
5937 if !tau.is_empty() {
5938 body["spec_tau"] = serde_json::Value::Object(tau);
5939 body["spec_accept_by_position"] = serde_json::Value::Object(by_position);
5940 }
5941}
5942
5943fn insert_peer_probe_metrics(
5944 body: &mut serde_json::Value,
5945 metrics_scope: &MetricsScope,
5946 snapshot: impl FnOnce() -> memra_engine::pp::PeerProbeMetrics,
5947) {
5948 if !metrics_scope.operator() {
5951 return;
5952 }
5953 let snapshot = snapshot();
5954 body["peer_probe_bypassed"] = json!(snapshot.bypassed);
5955 body["peer_probe_boundary_copies"] = json!(snapshot.boundary_copies);
5956 body["peer_probe_runtime_reprobes"] = json!(snapshot.runtime_probes);
5957 body["peer_probe_runtime_failures"] = json!(snapshot.runtime_failures);
5958 body["peer_probe_deferred_total"] = json!(snapshot.deferred_total);
5959 body["peer_probe_integrity_degraded"] = json!(snapshot.integrity_degraded);
5960 body["peer_probe_degraded_to_host_bounce"] = json!(snapshot.degraded_to_host_bounce);
5961}
5962
5963async fn get_metrics(State(st): State<AppState>, headers: HeaderMap) -> Response {
5965 let metrics_scope = match authorize_metrics(&st.api_auth, &st.metrics_auth, &headers) {
5966 Ok(scope) => scope,
5967 Err(response) => return response,
5968 };
5969 let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
5970 let mut body = if metrics_scope.process_wide() {
5974 json!({
5975 "admitted": m.admitted,
5976 "completed": m.completed,
5977 "tokens_out": m.tokens_out,
5978 "step_p50_ms": m.step_p50_ms,
5979 "step_p99_ms": m.step_p99_ms,
5980 "prompt_tokens_in": m.prompt_tokens_in,
5982 "cached_tokens_in": m.cached_tokens_in,
5983 "computed_tokens_in": m.prompt_tokens_in.saturating_sub(m.cached_tokens_in),
5986 "admission_session_defers": m.admission_session_defers,
5989 "admission_vram_defers": m.admission_vram_defers,
5990 "step_oom_parks": m.step_oom_parks,
5991 "continuation_pool_hits": m.continuation_pool_hits,
5992 "continuation_pool_evictions": m.continuation_pool_evictions,
5993 "plain_affinity_rewinds": m.plain_affinity_rewinds,
5994 "served_dspark": m.served_dspark,
5995 "served_spec": m.served_spec,
5996 "served_plain": m.served_plain,
5997 "spec_pool_hits": m.spec_pool_hits,
5998 "spec_pool_misses": m.spec_pool_misses,
5999 "spec_pool_affinity_rewinds": m.spec_pool_affinity_rewinds,
6000 "spec_pool_evictions": m.spec_pool_evictions,
6001 "spec_pool_sampler_refusals": m.spec_pool_sampler_refusals,
6004 })
6005 } else {
6006 json!({})
6007 };
6008 if metrics_scope.operator() {
6012 if let Some(budget_health) = st.metering.as_ref().and_then(|m| m.limits_health()) {
6013 body["budget_source_reload_failed"] = json!(budget_health.source_reload_failed);
6014 body["budget_source_reload_consecutive"] =
6015 json!(budget_health.source_reload_consecutive);
6016 body["budget_source_available"] = json!(budget_health.source_available);
6017 }
6018 body["cache_hit_token_ratio"] = json!(if m.prompt_tokens_in > 0 {
6020 m.cached_tokens_in as f64 / m.prompt_tokens_in as f64
6021 } else {
6022 0.0
6023 });
6024 body["prefix_cache_hits"] = json!(m.prefix_hits);
6025 body["prefix_cache_misses"] = json!(m.prefix_misses);
6026 body["prefix_cache_inserts"] = json!(m.prefix_inserts);
6027 body["prefix_cache_evictions"] = json!(m.prefix_evictions);
6028 body["prefix_cache_skips_budget"] = json!(m.prefix_skips_budget);
6029 body["prefix_cache_skips_pinned"] = json!(m.prefix_skips_pinned);
6030 body["prefix_cache_hit_tokens"] = json!(m.prefix_hit_tokens);
6031 body["prefix_host_entries"] = json!(m.prefix_host_entries);
6035 body["prefix_host_bytes"] = json!(m.prefix_host_bytes);
6036 body["prefix_host_demotions"] = json!(m.prefix_host_demotions);
6037 body["prefix_host_promotions"] = json!(m.prefix_host_promotions);
6038 body["prefix_host_demote_ms"] = json!(m.prefix_host_demote_ms);
6039 body["prefix_host_promote_ms"] = json!(m.prefix_host_promote_ms);
6040 body["prefix_host_rejected_allocs"] = json!(m.prefix_host_rejected_allocs);
6041 body["prefix_host_purges"] = json!(m.prefix_host_purges);
6042 body["prefix_host_purged_entries"] = json!(m.prefix_host_purged_entries);
6043 body["prefix_host_purged_bytes"] = json!(m.prefix_host_purged_bytes);
6044 body["prefix_host_tenant_rejects"] = json!(m.prefix_host_tenant_rejects);
6045 body["prefix_host_pause_demotes"] = json!(m.prefix_host_pause_demotes);
6049 body["prefix_host_pause_cancels"] = json!(m.prefix_host_pause_cancels);
6050 body["prefix_host_handoff_exports"] = json!(m.prefix_host_handoff_exports);
6051 body["prefix_host_handoff_imported_entries"] =
6052 json!(m.prefix_host_handoff_imported_entries);
6053 body["prefix_host_handoff_imported_bytes"] = json!(m.prefix_host_handoff_imported_bytes);
6054 body["prefix_host_handoff_skips"] = json!(m.prefix_host_handoff_skips);
6055 body["kv_flex_borrowed_bytes"] = json!(m.kv_flex_borrowed_bytes);
6060 body["kv_flex_sheds"] = json!(m.kv_flex_sheds);
6061 body["kv_flex_shed_ms"] = json!(m.kv_flex_shed_ms);
6062 body["lcp_histogram"] = json!({
6065 "edges": worker::LCP_HIST_EDGES.to_vec(),
6066 "counts": m.lcp_hist.to_vec(),
6067 });
6068 let idle_s = darklane::ValleySignal::new(st.health.clone()).idle_seconds();
6072 body["prefix_cache_entries"] = json!(m.prefix_entries);
6073 body["prefix_cache_bytes"] = json!(m.prefix_bytes);
6074 body["active_sessions"] = json!(m.active_sessions);
6075 body["queued_requests"] = json!(m.queued_requests);
6076 body["admission_inflight"] = json!(m.admission_inflight);
6080 body["admission_booked_bytes"] = json!(m.admission_booked_bytes);
6081 body["continuation_pool_entries"] = json!(m.continuation_pool_entries);
6082 body["spec_pool_entries"] = json!(m.spec_pool_entries);
6083 body["cuda_driver_free_bytes"] = json!(m.cuda_driver_free_bytes);
6084 body["cuda_pool_reserved_bytes"] = json!(m.cuda_pool_reserved_bytes);
6085 body["cuda_pool_used_bytes"] = json!(m.cuda_pool_used_bytes);
6086 body["cuda_pool_cached_bytes"] = json!(m.cuda_pool_cached_bytes);
6087 if !m.constraint_compiler_fail_closed.is_empty() {
6088 body["constraint_compiler_fail_closed"] = serde_json::Value::Object(
6089 m.constraint_compiler_fail_closed
6090 .iter()
6091 .map(|(model, gauge)| {
6092 let value = u8::from(gauge.load(std::sync::atomic::Ordering::Acquire));
6093 (model.clone(), json!(value))
6094 })
6095 .collect(),
6096 );
6097 }
6098 body["serve_idle_seconds"] = json!((idle_s * 1000.0).round() / 1000.0);
6099 }
6100 if !m.ns_tokens.is_empty() {
6105 let tenants: serde_json::Map<String, serde_json::Value> = m
6106 .ns_tokens
6107 .iter()
6108 .filter(|(ns, _)| metrics_scope.includes(ns))
6109 .map(|(ns, [p, c])| {
6110 (
6111 ns.clone(),
6112 json!({
6113 "prompt_tokens_in": p,
6114 "cached_tokens_in": c,
6115 "cache_hit_token_ratio": if *p > 0 { *c as f64 / *p as f64 } else { 0.0 },
6116 }),
6117 )
6118 })
6119 .collect();
6120 if !tenants.is_empty() {
6121 body["tenants"] = serde_json::Value::Object(tenants);
6122 }
6123 }
6124 let adsd_suspect_total: serde_json::Map<String, serde_json::Value> = m
6125 .adsd_suspect_total
6126 .iter()
6127 .filter(|(tenant, _)| metrics_scope.includes(tenant))
6128 .map(|(tenant, total)| (tenant.clone(), json!(total)))
6129 .collect();
6130 if !adsd_suspect_total.is_empty() {
6131 body["adsd_suspect_total"] = serde_json::Value::Object(adsd_suspect_total);
6132 }
6133 if metrics_scope.operator()
6135 && let Some((bg, mode)) = &st.bg
6136 {
6137 body["bg"] = bg.to_json(mode);
6138 }
6139 if metrics_scope.operator() {
6146 let spec: serde_json::Map<String, serde_json::Value> = m
6147 .spec
6148 .iter()
6149 .map(|(model, t)| {
6150 let n_pos = t
6151 .pos_drafted
6152 .iter()
6153 .rposition(|&d| d > 0)
6154 .map_or(0, |p| p + 1);
6155 (
6156 model.clone(),
6157 json!({
6158 "rounds": t.rounds,
6159 "drafted": t.drafted,
6160 "accepted": t.accepted,
6161 "acceptance_rate": if t.drafted > 0 {
6162 t.accepted as f64 / t.drafted as f64 } else { 0.0 },
6163 "tokens_per_round": if t.rounds > 0 {
6164 (t.accepted + t.rounds) as f64 / t.rounds as f64 } else { 0.0 },
6165 "pos_drafted": t.pos_drafted[..n_pos].to_vec(),
6166 "pos_accepted": t.pos_accepted[..n_pos].to_vec(),
6167 "accept_rate_per_pos": (0..n_pos).map(|j| if t.pos_drafted[j] > 0 {
6168 t.pos_accepted[j] as f64 / t.pos_drafted[j] as f64 } else { 0.0 })
6169 .collect::<Vec<f64>>(),
6170 }),
6171 )
6172 })
6173 .collect();
6174 if !spec.is_empty() {
6175 body["spec"] = serde_json::Value::Object(spec);
6176 }
6177 }
6178 insert_spec_acceptance_metrics(&mut body, &metrics_scope, || m.spec_window.clone());
6179 insert_dual_pp_metrics(&mut body, &metrics_scope, DualPpMetricsSnapshot::current);
6180 insert_pp_wave_metrics(&mut body, &metrics_scope, PpWaveMetricsSnapshot::current);
6181 insert_peer_probe_metrics(
6182 &mut body,
6183 &metrics_scope,
6184 memra_engine::pp::peer_probe_metrics,
6185 );
6186 Json(body).into_response()
6187}
6188
6189#[derive(Debug, Default, Deserialize)]
6190struct ModelsQuery {
6191 #[serde(default)]
6192 schema: Option<String>,
6193}
6194
6195fn models_openai_body(models: &[String]) -> serde_json::Value {
6196 let data: Vec<_> = models
6197 .iter()
6198 .map(|m| json!({ "id": m, "object": "model" }))
6199 .collect();
6200 json!({ "object": "list", "data": data })
6201}
6202
6203fn declared_surface(metadata: Option<&OpenRouterModelMetadata>) -> &'static str {
6208 match metadata.and_then(|m| m.surface.as_deref()) {
6209 Some("embedding") => "embedding",
6210 Some("rerank") => "rerank",
6211 _ => "chat",
6212 }
6213}
6214
6215fn openrouter_supported_parameters(
6216 caps: Option<&ModelCaps>,
6217 max_output_length: Option<u64>,
6218 is_chat: bool,
6219) -> serde_json::Value {
6220 let mut parameters = serde_json::Map::new();
6221 if !is_chat {
6229 return serde_json::Value::Object(parameters);
6230 }
6231 for name in [
6232 "temperature",
6233 "top_p",
6234 "min_p",
6235 "frequency_penalty",
6236 "presence_penalty",
6237 "repetition_penalty",
6238 "stop",
6239 ] {
6240 parameters.insert(name.into(), json!({ "type": "unknown" }));
6241 }
6242 parameters.insert("top_k".into(), json!({ "type": "integer", "min": 0 }));
6243 parameters.insert(
6244 "seed".into(),
6245 json!({ "type": "integer", "min": 0, "max": JSON_SAFE_INTEGER_MAX }),
6246 );
6247 let mut max_tokens = json!({ "type": "integer", "min": 1, "unit": "token" });
6248 if let Some(max) = max_output_length {
6249 max_tokens["max"] = json!(max);
6250 }
6251 parameters.insert("max_tokens".into(), max_tokens);
6252 if is_chat
6265 && caps.is_some_and(|c| {
6266 !c.dsv4 && !(c.qwen_think && !c.think_switch && c.think_close.is_empty())
6267 })
6268 {
6269 parameters.insert("json_mode".into(), json!({ "type": "boolean" }));
6270 parameters.insert("structured_outputs".into(), json!({ "type": "boolean" }));
6271 }
6272 if is_chat && caps.is_some_and(|c| c.tools_branch) {
6273 parameters.insert("tools".into(), json!({ "type": "boolean" }));
6274 parameters.insert(
6275 "tool_choice".into(),
6276 json!({ "type": "enum", "values": ["auto", "none"] }),
6277 );
6278 }
6279 if is_chat && caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think) {
6280 parameters.insert("reasoning".into(), json!({ "type": "boolean" }));
6281 }
6282 serde_json::Value::Object(parameters)
6283}
6284
6285fn published_context_length(
6301 caps: Option<&ModelCaps>,
6302 metadata: Option<&OpenRouterModelMetadata>,
6303) -> Option<u64> {
6304 let trained = caps
6305 .map(|c| c.context_length as u64)
6306 .filter(|&value| value > 0)?;
6307 let envelope = metadata.and_then(|m| {
6308 let prompt = m.max_prompt_length?;
6309 let output = m.max_output_length?;
6310 prompt.checked_add(output)
6311 });
6312 Some(envelope.map_or(trained, |envelope| trained.min(envelope)))
6313}
6314
6315fn model_entry_openrouter(
6316 name: &str,
6317 caps: Option<&ModelCaps>,
6318 metadata: Option<&OpenRouterModelMetadata>,
6319) -> serde_json::Value {
6320 let empty = OpenRouterModelMetadata::default();
6321 let metadata = metadata.unwrap_or(&empty);
6322 let context_length =
6323 published_context_length(caps, Some(metadata)).filter(|&v| v <= JSON_SAFE_INTEGER_MAX);
6324 let tokenizer = caps
6325 .map(|c| c.tokenizer.as_str())
6326 .filter(|tokenizer| !tokenizer.is_empty());
6327
6328 let mut input = serde_json::Map::new();
6329 input.insert("type".into(), json!("text"));
6330 let mut supported_inputs = serde_json::Map::new();
6331 if let Some(value) = context_length {
6332 supported_inputs.insert(
6333 "max_context_length".into(),
6334 json!({ "value": value, "unit": "token" }),
6335 );
6336 }
6337 if let Some(value) = metadata.max_prompt_length {
6338 supported_inputs.insert(
6339 "max_prompt_length".into(),
6340 json!({ "value": value, "unit": "token" }),
6341 );
6342 }
6343 if !supported_inputs.is_empty() {
6344 input.insert(
6345 "supported_inputs".into(),
6346 serde_json::Value::Object(supported_inputs),
6347 );
6348 }
6349 let mut input_pricing = Vec::new();
6350 for (kind, cost) in [
6351 ("prompt", metadata.pricing.prompt.as_deref()),
6352 ("cached_prompt", metadata.pricing.cached_prompt.as_deref()),
6353 ("cache_write", metadata.pricing.cache_write.as_deref()),
6354 ] {
6355 if let Some(cost) = cost {
6356 input_pricing.push(json!({
6357 "type": kind,
6358 "unit": "token",
6359 "cost_usd": cost,
6360 }));
6361 }
6362 }
6363 if !input_pricing.is_empty() {
6364 input.insert("pricing".into(), serde_json::Value::Array(input_pricing));
6365 }
6366 let mut input_capacity = Vec::new();
6367 for (kind, value) in [
6368 ("prompt", metadata.capacity.prompt_tpm),
6369 ("cached_prompt", metadata.capacity.cached_prompt_tpm),
6370 ] {
6371 if let Some(value) = value {
6372 input_capacity.push(json!({
6373 "type": kind,
6374 "unit": "token",
6375 "per": "minute",
6376 "value": value,
6377 }));
6378 }
6379 }
6380 if !input_capacity.is_empty() {
6381 input.insert("capacity".into(), serde_json::Value::Array(input_capacity));
6382 }
6383
6384 let or_surface = declared_surface(Some(metadata));
6385 let or_is_chat = or_surface == "chat";
6386 let mut output = serde_json::Map::new();
6387 output.insert(
6395 "type".into(),
6396 json!(match or_surface {
6397 "embedding" => "embeddings",
6398 "rerank" => "rerank",
6399 _ => "text",
6400 }),
6401 );
6402 output.insert(
6403 "supported_parameters".into(),
6404 openrouter_supported_parameters(caps, metadata.max_output_length, or_is_chat),
6405 );
6406 if or_is_chat {
6410 output.insert("streaming".into(), json!(true));
6411 }
6412 if let Some(value) = metadata.max_output_length
6415 && or_is_chat
6416 {
6417 output.insert(
6418 "max_length".into(),
6419 json!({ "value": value, "unit": "token" }),
6420 );
6421 }
6422 let mut output_pricing = Vec::new();
6423 for (kind, cost) in [
6424 ("completion", metadata.pricing.completion.as_deref()),
6425 (
6426 "internal_reasoning",
6427 metadata.pricing.internal_reasoning.as_deref(),
6428 ),
6429 ] {
6430 if let Some(cost) = cost {
6431 output_pricing.push(json!({
6432 "type": kind,
6433 "unit": "token",
6434 "cost_usd": cost,
6435 }));
6436 }
6437 }
6438 if !output_pricing.is_empty() {
6439 output.insert("pricing".into(), serde_json::Value::Array(output_pricing));
6440 }
6441 let mut output_capacity = Vec::new();
6442 if let Some(value) = metadata.capacity.completion_tpm {
6443 output_capacity.push(json!({
6444 "type": "completion",
6445 "unit": "token",
6446 "per": "minute",
6447 "value": value,
6448 }));
6449 }
6450 if let Some(value) = metadata.capacity.concurrency {
6451 output_capacity.push(json!({
6452 "type": "concurrency",
6453 "unit": "request",
6454 "value": value,
6455 }));
6456 }
6457 if !output_capacity.is_empty() {
6458 output.insert("capacity".into(), serde_json::Value::Array(output_capacity));
6459 }
6460
6461 let mut entry = serde_json::Map::new();
6462 entry.insert("schema_version".into(), json!(OPENROUTER_SCHEMA_VERSION));
6463 entry.insert("id".into(), json!(name));
6464 entry.insert("name".into(), json!(name));
6465 if let Some(value) = metadata.hugging_face_id.as_deref() {
6466 entry.insert("hugging_face_id".into(), json!(value));
6467 }
6468 if let Some(value) = metadata.created {
6469 entry.insert("created".into(), json!(value));
6470 }
6471 if let Some(value) = metadata.quantization.as_deref() {
6472 entry.insert("quantization".into(), json!(value));
6473 }
6474 if let Some(value) = tokenizer {
6475 entry.insert("tokenizer".into(), json!(value));
6476 }
6477 if let Some(value) = metadata.description.as_deref() {
6478 entry.insert("description".into(), json!(value));
6479 }
6480 let mut input_modalities = vec![serde_json::Value::Object(input)];
6481 for m in &metadata.input_modalities {
6482 let mut extra = serde_json::Map::new();
6483 extra.insert("type".into(), json!(m));
6484 if let Some(cost) = metadata.pricing.prompt.as_deref() {
6485 extra.insert(
6487 "pricing".into(),
6488 json!([{ "type": "prompt", "unit": "token", "cost_usd": cost }]),
6489 );
6490 }
6491 input_modalities.push(serde_json::Value::Object(extra));
6492 }
6493 entry.insert(
6494 "input_modalities".into(),
6495 serde_json::Value::Array(input_modalities),
6496 );
6497 entry.insert(
6498 "output_modalities".into(),
6499 serde_json::Value::Array(vec![serde_json::Value::Object(output)]),
6500 );
6501 if let Some(cost) = metadata.pricing.request.as_deref() {
6502 entry.insert(
6503 "pricing".into(),
6504 json!([{ "type": "request", "unit": "request", "cost_usd": cost }]),
6505 );
6506 }
6507 if let Some(value) = metadata.capacity.request_rpm {
6508 entry.insert(
6509 "capacity".into(),
6510 json!([{
6511 "type": "request",
6512 "unit": "request",
6513 "per": "minute",
6514 "value": value,
6515 }]),
6516 );
6517 }
6518 if let Some(value) = metadata.is_ready {
6519 entry.insert("is_ready".into(), json!(value));
6520 }
6521 if let Some(value) = metadata.is_free {
6522 entry.insert("is_free".into(), json!(value));
6523 }
6524 if let Some(value) = metadata.discount_to_user {
6525 entry.insert("discount_to_user".into(), json!(value));
6526 }
6527 if let Some(value) = metadata.openrouter_slug.as_deref() {
6528 entry.insert("openrouter".into(), json!({ "slug": value }));
6529 }
6530 if !metadata.datacenters.is_empty() {
6531 entry.insert("datacenters".into(), json!(metadata.datacenters));
6532 }
6533 let mut compliance = serde_json::Map::new();
6534 if let Some(value) = metadata.zdr {
6535 compliance.insert("zdr".into(), json!(value));
6536 }
6537 if let Some(value) = metadata.hipaa {
6538 compliance.insert("hipaa".into(), json!(value));
6539 }
6540 if !compliance.is_empty() {
6541 entry.insert("compliance".into(), serde_json::Value::Object(compliance));
6542 }
6543 serde_json::Value::Object(entry)
6544}
6545
6546fn models_openrouter_body(st: &AppState) -> serde_json::Value {
6547 let data: Vec<_> = st
6548 .models
6549 .iter()
6550 .map(|model| {
6551 model_entry_openrouter(model, st.caps.get(model), st.openrouter_metadata.get(model))
6552 })
6553 .collect();
6554 json!({ "data": data })
6555}
6556
6557fn model_entry_openmodels(
6558 name: &str,
6559 caps: Option<&ModelCaps>,
6560 metadata: Option<&OpenRouterModelMetadata>,
6561) -> Result<serde_json::Value, String> {
6562 let metadata = metadata.ok_or_else(|| {
6563 format!("OpenModels feed requires MEMRA_MODEL_METADATA for model {name:?}")
6564 })?;
6565 let context_length = published_context_length(caps, Some(metadata))
6566 .filter(|&value| value <= JSON_SAFE_INTEGER_MAX)
6567 .ok_or_else(|| format!("OpenModels feed requires context_length for model {name:?}"))?;
6568 let created = metadata
6569 .created
6570 .ok_or_else(|| format!("OpenModels feed requires created for model {name:?}"))?;
6571 let max_output_length = metadata
6572 .max_output_length
6573 .ok_or_else(|| format!("OpenModels feed requires max_output_length for model {name:?}"))?;
6574 let prompt = metadata
6575 .pricing
6576 .prompt
6577 .as_deref()
6578 .ok_or_else(|| format!("OpenModels feed requires pricing.prompt for model {name:?}"))?;
6579 let completion =
6580 metadata.pricing.completion.as_deref().ok_or_else(|| {
6581 format!("OpenModels feed requires pricing.completion for model {name:?}")
6582 })?;
6583 let input_cache_read = metadata.pricing.cached_prompt.as_deref().ok_or_else(|| {
6584 format!("OpenModels feed requires pricing.cached_prompt for model {name:?}")
6585 })?;
6586 let is_ready = metadata
6587 .is_ready
6588 .ok_or_else(|| format!("OpenModels feed requires is_ready for model {name:?}"))?;
6589 let is_free = metadata
6590 .is_free
6591 .ok_or_else(|| format!("OpenModels feed requires is_free for model {name:?}"))?;
6592 let discount_to_user = metadata
6593 .discount_to_user
6594 .ok_or_else(|| format!("OpenModels feed requires discount_to_user for model {name:?}"))?;
6595
6596 let mut pricing = serde_json::Map::new();
6597 pricing.insert("prompt".into(), json!(prompt));
6598 pricing.insert("completion".into(), json!(completion));
6599 pricing.insert("input_cache_read".into(), json!(input_cache_read));
6600 if let Some(value) = metadata.pricing.request.as_deref() {
6601 pricing.insert("request".into(), json!(value));
6602 }
6603
6604 let om_surface = declared_surface(Some(metadata));
6605 let om_is_chat = om_surface == "chat";
6606 let mut supported_features = Vec::new();
6607 if om_is_chat && caps.is_some_and(|c| c.tools_branch) {
6608 supported_features.push("tool_calling");
6609 }
6610 if om_is_chat && caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think) {
6611 supported_features.push("reasoning");
6612 }
6613
6614 let mut entry = serde_json::Map::new();
6615 entry.insert("id".into(), json!(name));
6616 entry.insert("name".into(), json!(name));
6617 entry.insert("created".into(), json!(created));
6618 entry.insert("input_modalities".into(), json!(["text"]));
6619 entry.insert(
6620 "output_modalities".into(),
6621 json!(match om_surface {
6622 "embedding" => ["embeddings"],
6623 "rerank" => ["rerank"],
6624 _ => ["text"],
6625 }),
6626 );
6627 entry.insert("context_length".into(), json!(context_length));
6628 entry.insert("max_output_length".into(), json!(max_output_length));
6629 entry.insert("currency".into(), json!("USD"));
6632 entry.insert("pricing".into(), serde_json::Value::Object(pricing));
6633 entry.insert("supported_features".into(), json!(supported_features));
6634 entry.insert("is_ready".into(), json!(is_ready));
6635 entry.insert("is_free".into(), json!(is_free));
6636 entry.insert("discount_to_user".into(), json!(discount_to_user));
6637 Ok(serde_json::Value::Object(entry))
6638}
6639
6640fn models_openmodels_body(st: &AppState) -> Result<serde_json::Value, String> {
6641 let data: Result<Vec<_>, _> = st
6642 .models
6643 .iter()
6644 .map(|model| {
6645 model_entry_openmodels(model, st.caps.get(model), st.openrouter_metadata.get(model))
6646 })
6647 .collect();
6648 Ok(json!({ "data": data? }))
6649}
6650
6651async fn list_models(State(st): State<AppState>, Query(query): Query<ModelsQuery>) -> Response {
6652 match query.schema.as_deref() {
6653 None | Some("openai") => Json(models_openai_body(st.models.as_ref())).into_response(),
6654 Some("openrouter") => Json(models_openrouter_body(&st)).into_response(),
6655 Some("openmodels") => match models_openmodels_body(&st) {
6656 Ok(body) => Json(body).into_response(),
6657 Err(error) => bad_request(&error, Some("schema")),
6658 },
6659 Some(schema) => bad_request(
6660 &format!(
6661 "unsupported models schema {schema:?}; expected openai, openrouter, or openmodels"
6662 ),
6663 Some("schema"),
6664 ),
6665 }
6666}
6667
6668fn model_entry_v1(
6676 name: &str,
6677 caps: Option<&ModelCaps>,
6678 metadata: Option<&OpenRouterModelMetadata>,
6679) -> serde_json::Value {
6680 let ctx = published_context_length(caps, metadata);
6681 let thinking = caps.is_some_and(|c| c.qwen_think || c.effort_levels || c.gemma_think || c.dsv4);
6685 let is_dsv4 = caps.is_some_and(|c| c.dsv4);
6688 let per_1m = |v: Option<&str>| match v.and_then(per_million_price) {
6689 Some(p) => json!(p),
6690 None => serde_json::Value::Null,
6691 };
6692 let owned_by = metadata
6693 .and_then(|m| m.owned_by.as_deref())
6694 .unwrap_or_else(|| name.split('/').next().unwrap_or(name));
6695 let mut input_modalities = vec!["text"];
6696 if let Some(meta) = metadata {
6697 input_modalities.extend(meta.input_modalities.iter().map(String::as_str));
6698 }
6699 let lifecycle = metadata.and_then(|m| m.lifecycle.as_ref());
6700 let reliability = metadata.and_then(|m| m.reliability.as_ref());
6701 let surface = declared_surface(metadata);
6707 let (model_type, endpoints, output_modalities) = match surface {
6708 "embedding" => ("embedding", vec!["embeddings"], vec!["embeddings"]),
6712 "rerank" => ("rerank", vec!["rerank"], vec!["rerank"]),
6713 _ => ("chat", vec!["chat/completions"], vec!["text"]),
6714 };
6715 let is_chat = surface == "chat";
6716 json!({
6717 "id": name,
6718 "name": name,
6719 "object": "model",
6720 "owned_by": owned_by,
6721 "type": model_type,
6722 "context_length": ctx,
6723 "max_output_tokens": if is_chat { metadata.and_then(|m| m.max_output_length) } else { None },
6726 "endpoints": endpoints,
6727 "input_modalities": input_modalities,
6728 "output_modalities": output_modalities,
6729 "capabilities": {
6730 "streaming": is_chat,
6733 "tools": is_chat && caps.is_some_and(|c| c.tools_branch),
6734 "structured_output": is_chat
6742 && !is_dsv4
6743 && !caps.is_some_and(|c| c.qwen_think && !c.think_switch && c.think_close.is_empty()),
6744 "reasoning": is_chat && thinking,
6745 "prompt_caching": is_chat && !is_dsv4,
6746 },
6747 "pricing": {
6748 "currency": "USD",
6749 "unit": "per_1m_tokens",
6750 "input": per_1m(metadata.and_then(|m| m.pricing.prompt.as_deref())),
6751 "output": per_1m(metadata.and_then(|m| m.pricing.completion.as_deref())),
6752 "cached_input": per_1m(metadata.and_then(|m| m.pricing.cached_prompt.as_deref())),
6753 "cache_write": per_1m(metadata.and_then(|m| m.pricing.cache_write.as_deref())),
6754 "minimum_request": metadata
6756 .and_then(|m| m.pricing.request.as_deref())
6757 .unwrap_or("0"),
6758 },
6759 "lifecycle": {
6760 "status": lifecycle.and_then(|l| l.status.as_deref()).unwrap_or("active"),
6761 "deprecation_at": lifecycle.and_then(|l| l.deprecation_at.as_deref()),
6762 "retirement_at": lifecycle.and_then(|l| l.retirement_at.as_deref()),
6763 "replacement_model_id": lifecycle.and_then(|l| l.replacement_model_id.as_deref()),
6764 },
6765 "reliability": {
6766 "first_token_timeout_seconds":
6767 reliability.and_then(|r| r.first_token_timeout_seconds).unwrap_or(120),
6768 "completion_timeout_seconds":
6769 reliability.and_then(|r| r.completion_timeout_seconds).unwrap_or(900),
6770 "stream_idle_timeout_seconds":
6771 reliability.and_then(|r| r.stream_idle_timeout_seconds).unwrap_or(60),
6772 "capacity_scope":
6773 reliability.and_then(|r| r.capacity_scope.as_deref()).unwrap_or("model_region"),
6774 },
6775 })
6776}
6777
6778async fn list_models_v1(State(st): State<AppState>) -> impl IntoResponse {
6781 let data: Vec<_> = st
6782 .models
6783 .iter()
6784 .map(|m| model_entry_v1(m, st.caps.get(m), st.openrouter_metadata.get(m)))
6785 .collect();
6786 let mut body = json!({
6787 "object": "list",
6788 "contract_version": "2.0",
6789 "data": data,
6790 });
6791 if let Some(provider) = st.provider_metadata.as_ref() {
6796 body["provider"] = json!({
6797 "id": provider.id,
6798 "status_url": provider.status_url,
6799 "support_contact": provider.support_contact,
6800 "incident_contact": provider.incident_contact,
6801 "regions": provider.regions,
6802 "request_id_header": "x-request-id",
6803 "error_contract": {
6804 "rate_limit_status": 429,
6805 "overload_status": 503,
6806 "retry_after_header": "Retry-After",
6807 "account_quota_error_codes": ["insufficient_balance"],
6808 },
6809 });
6810 }
6811 Json(body)
6812}
6813
6814async fn yield_metrics(State(st): State<AppState>, headers: HeaderMap) -> Response {
6817 let metrics_scope = match authorize_metrics(&st.api_auth, &st.metrics_auth, &headers) {
6818 Ok(scope) => scope,
6819 Err(response) => return response,
6820 };
6821 if !metrics_scope.process_wide() {
6822 return error_response(
6823 StatusCode::FORBIDDEN,
6824 "completion api keys do not authorize process-wide yield metrics; configure \
6825 MEMRA_METRICS_TOKEN",
6826 "authentication_error",
6827 None,
6828 );
6829 }
6830 let m = st.metrics.lock().map(|m| m.clone()).unwrap_or_default();
6831 let lane = |i: usize| {
6832 json!({
6833 "admitted": m.lane_admitted[i], "shed": m.lane_shed[i],
6834 "completed": m.lane_completed[i], "tokens_out": m.lane_tokens[i],
6835 })
6836 };
6837 let mut body = json!({
6838 "lanes": {
6839 "interactive": lane(0), "judge": lane(1), "harvest": lane(2),
6840 },
6841 "interactive_step_ms": { "p50": m.step_p50_ms, "p99": m.step_p99_ms },
6842 });
6843 if metrics_scope.operator() {
6844 body["batch_size_last"] = json!(m.batch_size_last);
6845 }
6846 Json(body).into_response()
6847}
6848
6849async fn peek_admission(
6862 mut rx: worker::EventReceiver,
6863) -> Result<worker::EventReceiver, (Response, &'static str)> {
6864 match rx.recv().await {
6865 Some(Event::Error(e)) => {
6870 let error_code = engine_error_code(e.class);
6871 Err((engine_error_response(&e), error_code))
6872 }
6873 first => {
6874 let (tx2, rx2) = worker::event_channel();
6875 if let Some(ev) = first {
6876 let _ = tx2.send(ev);
6877 }
6878 tokio::spawn(forward_events(rx, tx2));
6879 Ok(rx2)
6880 }
6881 }
6882}
6883
6884async fn forward_events(mut rx: worker::EventReceiver, tx2: worker::EventSender) {
6895 loop {
6896 tokio::select! {
6897 biased;
6898 () = tx2.closed() => break,
6899 ev = rx.recv() => match ev {
6900 Some(ev) => {
6901 if tx2.send(ev).is_err() {
6902 break;
6903 }
6904 }
6905 None => break,
6906 },
6907 }
6908 }
6909}
6910
6911async fn peek_first_token(
6926 mut rx: worker::EventReceiver,
6927 deadline: RequestDeadline,
6928) -> Result<worker::EventReceiver, ()> {
6929 let mut buffered: Vec<Event> = Vec::new();
6930 loop {
6931 match tokio::time::timeout_at(deadline.at, rx.recv()).await {
6932 Err(_) => return Err(()), Ok(None) => break, Ok(Some(ev)) => {
6935 let first_delivery = matches!(
6936 ev,
6937 Event::Token { .. } | Event::Done { .. } | Event::Error(_)
6938 );
6939 buffered.push(ev);
6940 if first_delivery {
6941 break;
6942 }
6943 }
6944 }
6945 }
6946 let (tx2, rx2) = worker::event_channel();
6947 for ev in buffered {
6948 let _ = tx2.send(ev);
6949 }
6950 tokio::spawn(forward_events(rx, tx2));
6951 Ok(rx2)
6952}
6953
6954#[cfg(test)]
6956fn build_request(
6960 req: &CompletionReq,
6961 tx: worker::EventSender,
6962 lane: lanes::Lane,
6963 affinity: Option<String>,
6964) -> Request {
6965 build_request_with_trace(req, tx, lane, affinity, None, &SamplingDefaults::default())
6966}
6967
6968fn build_request_with_trace(
6969 req: &CompletionReq,
6970 tx: worker::EventSender,
6971 lane: lanes::Lane,
6972 affinity: Option<String>,
6973 ttft: Option<Arc<ttft::Trace>>,
6974 sampling_defaults: &SamplingDefaults,
6975) -> Request {
6976 let params = GenParams {
6977 max_new: req.max_tokens.unwrap_or(worker::MAX_NEW_CTX_BOUNDED),
6978 max_ctx: req.max_ctx,
6979 eos: Vec::new(), };
6981 let sampler_cfg = resolve_sampler_config(req.into(), sampling_defaults);
6986 Request {
6987 model: req.model.clone(),
6988 prompt_ids: req.prompt_ids.clone(),
6989 prompt_text: req.prompt.clone(),
6990 chat: req.chat,
6991 chat_turns: Vec::new(),
6992 tools_json: Vec::new(),
6993 tools_struct: Vec::new(),
6994 think: ThinkMode::Default,
6995 reasoning_effort: None, params,
6997 sampler_cfg,
6998 stop_strings: req.stop.clone().into_vec(),
6999 trace_id: req.trace_id.clone(),
7000 request_id: String::new(),
7003 admit_predict_logged: false,
7004 max_prompt_tokens: None,
7005 cache_ns: cache_namespace(&req.cache_salt),
7006 affinity,
7007 lane,
7008 grammar: None, prepared_constraint: None,
7010 constraint_ready: None,
7011 oom_retries: 0, spec_k_replay: None,
7013 prepared_prompt: None,
7014 capture: None, images: Vec::new(), gemma_images: Vec::new(),
7017 glm5_images: Vec::new(),
7018 step_images: Vec::new(),
7019 vision_memory: None,
7020 wire_deadline: None, ttft,
7022 tx,
7023 }
7024}
7025
7026struct ChatPlan {
7029 request: Request,
7030 parser: Option<ToolStreamParser>,
7033 pending_images: Vec<PendingVisionUnit>,
7036 pending_gemma: Vec<PendingGemmaImage>,
7037 pending_glm5: Vec<PendingGlm5Image>,
7038 pending_step: Vec<PendingStepImage>,
7039 vision_memory: Option<VisionMemoryPermit>,
7043}
7044
7045pub(crate) fn request_has_vision(req: &ChatCompletionReq) -> bool {
7046 req.messages.iter().any(|message| {
7047 message.content.as_array().is_some_and(|parts| {
7048 parts.iter().any(|part| {
7049 matches!(
7050 part.get("type").and_then(serde_json::Value::as_str),
7051 Some("image_url" | "video_url")
7052 )
7053 })
7054 })
7055 })
7056}
7057
7058fn planned_vision_bytes(plan: &ChatPlan) -> Result<usize, String> {
7059 let mut total = 0usize;
7060 let mut add = |bytes: usize| {
7061 total = total.checked_add(bytes).ok_or_else(|| {
7062 "vision patch memory reservation overflowed while planning".to_string()
7063 })?;
7064 Ok::<(), String>(())
7065 };
7066 for unit in &plan.pending_images {
7067 let bytes = match unit {
7068 PendingVisionUnit::Still { gh, gw, .. } => gh
7069 .checked_mul(*gw)
7070 .and_then(|n| n.checked_mul(memra_engine::vision::V_PATCH_IN))
7071 .and_then(|n| n.checked_mul(std::mem::size_of::<f32>()))
7072 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?,
7073 PendingVisionUnit::Video { groups, .. } => {
7074 groups.iter().try_fold(0usize, |total, group| {
7075 let bytes = group
7076 .gh
7077 .checked_mul(group.gw)
7078 .and_then(|n| n.checked_mul(memra_engine::vision::V_PATCH_IN))
7079 .and_then(|n| n.checked_mul(std::mem::size_of::<f32>()))
7080 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?;
7081 total.checked_add(bytes).ok_or_else(|| {
7082 "vision patch memory reservation overflowed while planning".to_string()
7083 })
7084 })?
7085 }
7086 };
7087 add(bytes)?;
7088 }
7089 for unit in &plan.pending_gemma {
7090 let bytes = unit
7091 .gw
7092 .checked_mul(unit.gh)
7093 .and_then(|n| n.checked_mul(memra_engine::vision_gemma::GV_PATCH_IN))
7094 .and_then(|n| n.checked_mul(std::mem::size_of::<f32>()))
7095 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?;
7096 add(bytes)?;
7097 }
7098 for unit in &plan.pending_glm5 {
7099 let bytes = unit
7100 .gh
7101 .checked_mul(unit.gw)
7102 .and_then(|n| n.checked_mul(memra_engine::vision_glm5::G5V_PATCH_IN))
7103 .and_then(|n| n.checked_mul(std::mem::size_of::<f32>()))
7104 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?;
7105 add(bytes)?;
7106 }
7107 for unit in &plan.pending_step {
7108 use memra_engine::vision_step::{SV_GRID_MAIN, SV_GRID_TILE, SV_PATCH_IN};
7109 let patches = unit
7111 .plan
7112 .n_tiles
7113 .checked_mul(SV_GRID_TILE * SV_GRID_TILE)
7114 .and_then(|n| n.checked_add(SV_GRID_MAIN * SV_GRID_MAIN))
7115 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?;
7116 let bytes = patches
7117 .checked_mul(SV_PATCH_IN)
7118 .and_then(|n| n.checked_mul(std::mem::size_of::<f32>()))
7119 .ok_or_else(|| "vision patch memory reservation overflowed".to_string())?;
7120 add(bytes)?;
7121 }
7122 Ok(total)
7123}
7124
7125pub(crate) fn reserve_vision_memory(
7126 plan: &ChatPlan,
7127) -> Result<Option<VisionMemoryPermit>, VisionMemoryError> {
7128 let bytes = planned_vision_bytes(plan).map_err(VisionMemoryError::Request)?;
7129 try_reserve_vision_memory(bytes)
7130}
7131
7132#[cfg(test)]
7133fn build_chat_request(
7134 req: ChatCompletionReq,
7135 caps: Option<&ModelCaps>,
7136 tx: worker::EventSender,
7137 lane: lanes::Lane,
7138 affinity: Option<String>,
7139) -> Result<ChatPlan, String> {
7140 let defaults = ModelSamplingDefaults::resolve(None, caps);
7143 build_chat_request_with_trace(req, caps, tx, lane, affinity, None, None, &defaults)
7144}
7145
7146#[allow(clippy::too_many_arguments)]
7160fn build_chat_request_with_trace(
7161 req: ChatCompletionReq,
7162 caps: Option<&ModelCaps>,
7163 tx: worker::EventSender,
7164 lane: lanes::Lane,
7165 affinity: Option<String>,
7166 ttft: Option<Arc<ttft::Trace>>,
7167 default_effort: Option<&str>,
7168 sampling_defaults: &ModelSamplingDefaults,
7169) -> Result<ChatPlan, String> {
7170 req.stop.validate()?;
7171 let client_sampling: ClientSampling = (&req).into();
7174 let tool_choice = parse_tool_choice(&req.tool_choice)?;
7175 if let Some(c) = caps
7180 && !c.chat_ok
7181 {
7182 return Err(format!(
7183 "model {:?} has no chat template (checkpoint carries neither \
7184 tokenizer_config.json chat_template nor chat_template.jinja) — \
7185 /v1/chat/completions unavailable; use /v1/completions with a raw prompt",
7186 req.model
7187 ));
7188 }
7189 let vllm_switch = resolve_vllm_think_switch(req.enable_thinking, &req.chat_template_kwargs)?;
7190 let (mut think, effort_level, think_client_explicit) = parse_think(
7191 &req.reasoning_effort,
7192 &req.reasoning,
7193 vllm_switch,
7194 req.include_reasoning,
7195 default_effort,
7196 caps.is_some_and(|c| c.dsv4 || c.glm5),
7200 )?;
7201 let level_template = caps
7205 .map(|c| c.effort_levels || c.dsv4 || c.qwen_effort || c.glm5)
7206 .unwrap_or(false);
7207 if think_client_explicit
7228 && think == ThinkMode::NoThink
7229 && let Some(c) = caps
7230 && c.qwen_think
7231 && !c.think_switch
7232 && !c.dsv4
7233 {
7234 return Err(format!(
7235 "model {:?} cannot disable reasoning: its chat template opens a think \
7236 tail unconditionally and carries no enable_thinking switch, so \
7237 reasoning_effort/enable_thinking cannot turn it off on this model",
7238 req.model
7239 ));
7240 }
7241 let reasoning_effort = if level_template { effort_level } else { None };
7274 let grammar = constrained::parse_response_format(req.response_format.as_ref())?;
7277 if grammar.is_some()
7291 && let Some(c) = caps
7292 && c.qwen_think
7293 && think != ThinkMode::NoThink
7294 {
7295 if c.think_switch {
7296 think = ThinkMode::NoThink;
7297 } else if c.think_close.is_empty() {
7298 return Err(
7299 "response_format requires the model's think channel to close \
7300 before the grammar can engage, but this chat template has \
7301 neither an enable_thinking switch nor a recognizable \
7302 think-close token sequence"
7303 .into(),
7304 );
7305 }
7306 }
7310
7311 let sampler_cfg = resolve_sampler_config(client_sampling, sampling_defaults.for_mode(think));
7318
7319 let (tools_json, tools_struct, schemas) =
7322 if !req.tools.is_empty() && tool_choice == ToolChoice::Auto {
7323 prepare_tools(&req.tools)?
7324 } else {
7325 (Vec::new(), Vec::new(), HashMap::new())
7326 };
7327
7328 let mut turns: Vec<TmplTurn> = Vec::with_capacity(req.messages.len());
7329 let mut images: Vec<PendingVisionUnit> = Vec::new();
7330 let mut gemma_images: Vec<PendingGemmaImage> = Vec::new();
7331 let mut glm5_images: Vec<PendingGlm5Image> = Vec::new();
7332 let mut step_images: Vec<PendingStepImage> = Vec::new();
7333 let mut next_video = 0usize;
7334 for msg in &req.messages {
7335 let content = content_to_text_vision(
7336 &msg.content,
7337 &mut images,
7338 &mut gemma_images,
7339 &mut glm5_images,
7340 &mut step_images,
7341 &mut next_video,
7342 )
7343 .map_err(|e| format!("{} message: {e}", msg.role))?;
7344 let tool_calls = msg
7345 .tool_calls
7346 .iter()
7347 .map(render_req_tool_call)
7348 .collect::<Result<Vec<_>, _>>()?;
7349 if !tool_calls.is_empty() && msg.role != "assistant" {
7350 return Err("tool_calls are only valid on assistant messages".into());
7351 }
7352 let role = if msg.role == "developer" {
7355 "system".to_string()
7356 } else {
7357 msg.role.clone()
7358 };
7359 turns.push(TmplTurn {
7360 role,
7361 content,
7362 tool_calls,
7363 reasoning: msg.reasoning.clone().filter(|r| !r.is_empty()),
7365 tool_call_id: msg.tool_call_id.clone(),
7366 tool_name: msg.name.clone(),
7367 tool_responses: Vec::new(),
7368 task: None,
7372 tools: Vec::new(),
7373 });
7374 }
7375
7376 let has_tool_features = !tools_json.is_empty()
7379 || turns
7380 .iter()
7381 .any(|t| t.role == "tool" || !t.tool_calls.is_empty());
7382 if has_tool_features && !caps.map(|c| c.tools_branch).unwrap_or(false) {
7383 return Err(format!(
7384 "model {:?} chat template has no tools branch",
7385 req.model
7386 ));
7387 }
7388
7389 let think_open = caps
7392 .map(|c| c.qwen_think && !(think == ThinkMode::NoThink && c.think_switch))
7393 .unwrap_or(false);
7394 let gemma_tools = !tools_json.is_empty() && caps.map(|c| c.gemma_think).unwrap_or(false);
7412 let is_dsv4 = caps.map(|c| c.dsv4).unwrap_or(false);
7418 let dsv4_think_open = is_dsv4 && think != ThinkMode::NoThink;
7419 let dsv4_tools = is_dsv4 && !tools_struct.is_empty();
7420 let glm5 = caps.map(|c| c.glm5).unwrap_or(false);
7426 let is_hy3 = caps.map(|c| c.hy3).unwrap_or(false);
7429 let hy3_think_open = is_hy3 && think == ThinkMode::Think;
7430 let hy3_tools = is_hy3 && !tools_json.is_empty();
7431 let parser = if glm5 {
7432 Some(ToolStreamParser::glm5(think_open, schemas))
7433 } else if is_hy3 && (hy3_tools || hy3_think_open) {
7434 Some(ToolStreamParser::hy3(schemas, hy3_think_open))
7435 } else if is_dsv4 && (dsv4_tools || dsv4_think_open) {
7436 Some(ToolStreamParser::dsv4(dsv4_think_open))
7437 } else if gemma_tools {
7438 Some(ToolStreamParser::gemma_tools())
7439 } else if !tools_json.is_empty() {
7440 Some(ToolStreamParser::new(schemas, think_open))
7441 } else if think_open {
7442 Some(ToolStreamParser::reasoning_only())
7443 } else if caps.map(|c| c.gemma_think).unwrap_or(false) {
7444 Some(ToolStreamParser::gemma_thought())
7452 } else {
7453 None
7454 };
7455
7456 Ok(ChatPlan {
7457 request: Request {
7458 model: req.model,
7459 prompt_ids: Vec::new(),
7460 prompt_text: String::new(),
7461 chat: false,
7462 chat_turns: turns,
7463 tools_json,
7464 tools_struct,
7465 think,
7466 reasoning_effort,
7467 params: GenParams {
7468 max_new: req.max_tokens.unwrap_or(worker::MAX_NEW_CTX_BOUNDED),
7469 max_ctx: req.max_ctx,
7470 eos: Vec::new(),
7471 },
7472 sampler_cfg,
7473 stop_strings: {
7474 let mut stops = req.stop.into_vec();
7479 if gemma_tools {
7480 stops.push("<tool_call|>".to_string());
7481 }
7482 if dsv4_tools {
7487 stops.push("</\u{ff5c}DSML\u{ff5c}tool_calls>".to_string());
7488 }
7489 if hy3_tools {
7492 stops.push("</tool_calls:opensource>".to_string());
7493 }
7494 stops
7495 },
7496 trace_id: None,
7497 request_id: String::new(),
7500 admit_predict_logged: false,
7501 max_prompt_tokens: None,
7502 cache_ns: cache_namespace(&req.cache_salt),
7503 affinity,
7504 lane,
7505 grammar,
7506 prepared_constraint: None,
7507 constraint_ready: None,
7508 oom_retries: 0, spec_k_replay: None,
7510 prepared_prompt: None,
7511 images: Vec::new(),
7516 gemma_images: Vec::new(),
7517 glm5_images: Vec::new(),
7518 step_images: Vec::new(),
7519 capture: None, vision_memory: None,
7521 wire_deadline: None, ttft,
7523 tx,
7524 },
7525 parser,
7526 pending_images: images,
7527 pending_gemma: gemma_images,
7528 pending_glm5: glm5_images,
7529 pending_step: step_images,
7530 vision_memory: None,
7531 })
7532}
7533
7534fn decode_pending_vision(plan: &mut ChatPlan) -> Result<(), String> {
7540 for (i, unit) in plan.pending_images.drain(..).enumerate() {
7541 match unit {
7542 PendingVisionUnit::Still { bytes, gh, gw } => {
7543 let prep = memra_engine::vision_pre::prep_image_bytes(&bytes)
7544 .map_err(|e| format!("image {}: {e}", i + 1))?;
7545 if (prep.gh, prep.gw) != (gh, gw) {
7546 return Err(format!(
7547 "image {}: decoded grid {}x{} differs from its header-planned grid {gh}x{gw} — refusing (pad runs already rendered)",
7548 i + 1,
7549 prep.gh,
7550 prep.gw
7551 ));
7552 }
7553 plan.request
7554 .images
7555 .push(memra_engine::vision_pre::VisionUnit { prep, video: None });
7556 }
7557 PendingVisionUnit::Video {
7558 bytes,
7559 groups,
7560 video,
7561 } => {
7562 let prepared = memra_engine::vision_pre::prep_video_gif(&bytes)
7563 .map_err(|e| format!("video {}: {e}", i + 1))?;
7564 if prepared.groups.len() != groups.len() {
7565 return Err(format!(
7566 "video {}: decoded {} groups differ from its header-planned {} groups",
7567 i + 1,
7568 prepared.groups.len(),
7569 groups.len()
7570 ));
7571 }
7572 for ((group, prep), timestamp) in
7573 groups.iter().zip(prepared.groups).zip(prepared.timestamps)
7574 {
7575 if (prep.gh, prep.gw) != (group.gh, group.gw) {
7576 return Err(format!(
7577 "video {}: decoded grid {}x{} differs from its header-planned grid {}x{}",
7578 i + 1,
7579 prep.gh,
7580 prep.gw,
7581 group.gh,
7582 group.gw
7583 ));
7584 }
7585 if (timestamp - group.timestamp).abs() > 0.001 {
7586 return Err(format!(
7587 "video {}: decoded timestamp {timestamp:.3} differs from its header-planned timestamp {:.3}",
7588 i + 1,
7589 group.timestamp
7590 ));
7591 }
7592 plan.request
7593 .images
7594 .push(memra_engine::vision_pre::VisionUnit {
7595 prep,
7596 video: Some(video),
7597 });
7598 }
7599 }
7600 }
7601 }
7602 for (i, unit) in plan.pending_gemma.drain(..).enumerate() {
7603 let (patches, gw, gh) = memra_engine::vision_gemma::gemma_prep_image(&unit.bytes)
7604 .map_err(|e| format!("image {}: {e}", i + 1))?;
7605 if (gw, gh) != (unit.gw, unit.gh) {
7606 return Err(format!(
7607 "image {}: decoded grid {gw}x{gh} differs from its header-planned grid {}x{} — refusing (pad runs already rendered)",
7608 i + 1,
7609 unit.gw,
7610 unit.gh
7611 ));
7612 }
7613 plan.request
7614 .gemma_images
7615 .push(memra_engine::vision_gemma::GemmaVisionUnit { patches, gw, gh });
7616 }
7617 for (i, unit) in plan.pending_glm5.drain(..).enumerate() {
7618 let (patches, gh, gw) = memra_engine::vision_glm5::glm5_prep_image(&unit.bytes)
7619 .map_err(|e| format!("image {}: {e}", i + 1))?;
7620 if (gh, gw) != (unit.gh, unit.gw) {
7621 return Err(format!(
7622 "image {}: decoded grid {gh}x{gw} differs from its header-planned grid {}x{} — refusing (placeholder runs already rendered)",
7623 i + 1,
7624 unit.gh,
7625 unit.gw
7626 ));
7627 }
7628 plan.request
7629 .glm5_images
7630 .push(memra_engine::vision_glm5::Glm5VisionUnit { patches, gh, gw });
7631 }
7632 for (i, unit) in plan.pending_step.drain(..).enumerate() {
7633 let prepped = memra_engine::vision_step::step_prep_image(&unit.bytes)
7634 .map_err(|e| format!("image {}: {e}", i + 1))?;
7635 if prepped.tiles.len() != unit.plan.n_tiles
7636 || prepped.newline_mask != unit.plan.newline_mask
7637 {
7638 return Err(format!(
7639 "image {}: decoded tiling ({} tiles) differs from its header-planned tiling \
7640 ({} tiles) — refusing (pad runs already rendered)",
7641 i + 1,
7642 prepped.tiles.len(),
7643 unit.plan.n_tiles
7644 ));
7645 }
7646 plan.request.step_images.push(prepped);
7647 }
7648 Ok(())
7649}
7650
7651fn bearer_token(headers: &HeaderMap) -> Option<&str> {
7659 headers
7660 .get("authorization")
7661 .and_then(|value| value.to_str().ok())
7662 .and_then(|value| value.strip_prefix("Bearer "))
7663}
7664
7665fn authentication_error(why: auth::AuthDenied) -> Response {
7666 match why {
7667 auth::AuthDenied::Unknown => error_response(
7668 StatusCode::UNAUTHORIZED,
7669 "invalid api key",
7670 "authentication_error",
7671 None,
7672 ),
7673 auth::AuthDenied::Disabled => error_response(
7674 StatusCode::FORBIDDEN,
7675 "api key is disabled",
7676 "authentication_error",
7677 None,
7678 ),
7679 }
7680}
7681
7682#[allow(clippy::result_large_err)] fn authenticate(api_auth: &ApiAuth, headers: &HeaderMap) -> Result<auth::TenantCtx, Response> {
7684 auth::authenticate_with(
7685 api_auth.keyring,
7686 api_auth.single_key.as_deref(),
7687 bearer_token(headers),
7688 )
7689 .map_err(authentication_error)
7690}
7691
7692#[derive(Debug, Clone, PartialEq, Eq)]
7693enum MetricsScope {
7694 All,
7695 CompletionDomain,
7696 Tenant(String),
7697}
7698
7699impl MetricsScope {
7700 fn operator(&self) -> bool {
7701 matches!(self, MetricsScope::All)
7702 }
7703
7704 fn process_wide(&self) -> bool {
7705 matches!(self, MetricsScope::All | MetricsScope::CompletionDomain)
7706 }
7707
7708 fn includes(&self, tenant_row: &str) -> bool {
7709 match self {
7710 MetricsScope::All | MetricsScope::CompletionDomain => true,
7711 MetricsScope::Tenant(tenant) => tenant == tenant_row,
7712 }
7713 }
7714}
7715
7716#[allow(clippy::result_large_err)] fn authorize_metrics(
7718 api_auth: &ApiAuth,
7719 metrics_auth: &MetricsAuth,
7720 headers: &HeaderMap,
7721) -> Result<MetricsScope, Response> {
7722 if !metrics_auth.required {
7723 return Ok(MetricsScope::All);
7724 }
7725 let Some(candidate) = bearer_token(headers) else {
7726 return Err(authentication_error(auth::AuthDenied::Unknown));
7727 };
7728 if let Some(token) = metrics_auth.token.as_deref() {
7729 if auth::constant_time_secret_eq(token, candidate) {
7730 return Ok(MetricsScope::All);
7731 }
7732 if api_auth.configured() {
7733 return match auth::authenticate_with(
7734 api_auth.keyring,
7735 api_auth.single_key.as_deref(),
7736 Some(candidate),
7737 ) {
7738 Ok(_) => Err(error_response(
7739 StatusCode::FORBIDDEN,
7740 "completion api keys do not authorize metrics while \
7741 MEMRA_METRICS_TOKEN is configured",
7742 "authentication_error",
7743 None,
7744 )),
7745 Err(why) => Err(authentication_error(why)),
7746 };
7747 }
7748 return Err(authentication_error(auth::AuthDenied::Unknown));
7749 }
7750 if api_auth.configured() {
7751 let tenant = authenticate(api_auth, headers)?;
7752 return Ok(if api_auth.keyring.is_some() {
7753 MetricsScope::Tenant(format!("t:{}", tenant.tenant))
7754 } else {
7755 MetricsScope::CompletionDomain
7759 });
7760 }
7761 Err(authentication_error(auth::AuthDenied::Unknown))
7762}
7763
7764#[allow(clippy::result_large_err)] fn lane_for_tenant(
7771 headers: &axum::http::HeaderMap,
7772 tenant: &auth::TenantCtx,
7773) -> Result<lanes::Lane, Response> {
7774 let requested = match headers.get("x-lane").map(|v| v.to_str().unwrap_or("?")) {
7775 None => None,
7776 Some(v) => Some(lanes::Lane::parse(v).ok_or_else(|| {
7781 error_response_coded(
7782 StatusCode::BAD_REQUEST,
7783 &format!("unknown x-lane {v:?}; expected one of interactive, judge, harvest"),
7784 "invalid_request_error",
7785 Some("x-lane"),
7786 Some("invalid_lane"),
7787 )
7788 })?),
7789 };
7790 match tenant.lane_class {
7791 auth::LaneClass::Interactive => Ok(requested.unwrap_or(lanes::Lane::Interactive)),
7792 auth::LaneClass::Batch => match requested {
7793 None => Ok(lanes::Lane::Harvest),
7794 Some(lanes::Lane::Interactive) => Err(error_response(
7795 StatusCode::FORBIDDEN,
7796 "this api key is batch-class: x-lane interactive is not permitted \
7797 (use judge or harvest)",
7798 "authentication_error",
7799 Some("x-lane"),
7800 )),
7801 Some(l) => Ok(l),
7802 },
7803 }
7804}
7805
7806fn tenant_namespace(
7810 tenant: &auth::TenantCtx,
7811 cache_salt: &Option<String>,
7812) -> Result<String, &'static str> {
7813 let keyring_configured = auth::global().is_some();
7814 let raw = validate_cache_namespace(cache_salt, keyring_configured)?;
7815 if keyring_configured {
7816 Ok(auth::scope_namespace(&tenant.tenant, &raw))
7817 } else {
7818 Ok(raw)
7819 }
7820}
7821
7822fn meter_admit(env: &Envelope, tenant: &auth::TenantCtx, model: &str, lane: lanes::Lane) {
7827 eprintln!(
7828 "[meter] admit id={} tenant={} lane={} model={:?}",
7829 env.id,
7830 tenant.tenant,
7831 lane.as_str(),
7832 model
7833 );
7834}
7835
7836fn apply_model_request_limits(
7837 request: &mut Request,
7838 metadata: Option<&OpenRouterModelMetadata>,
7839 caps: Option<&ModelCaps>,
7840) -> Result<(), (String, &'static str)> {
7841 let Some(metadata) = metadata else {
7842 return Ok(());
7843 };
7844 let max_prompt = metadata
7845 .max_prompt_length
7846 .map(usize::try_from)
7847 .transpose()
7848 .map_err(|_| {
7849 (
7850 "configured model prompt limit does not fit this platform".into(),
7851 "model",
7852 )
7853 })?;
7854 let max_output = metadata
7855 .max_output_length
7856 .map(usize::try_from)
7857 .transpose()
7858 .map_err(|_| {
7859 (
7860 "configured model output limit does not fit this platform".into(),
7861 "model",
7862 )
7863 })?;
7864
7865 request.max_prompt_tokens = max_prompt;
7866 if let Some(max_output) = max_output {
7867 if request.params.max_new == worker::MAX_NEW_CTX_BOUNDED {
7868 request.params.max_new = metadata
7869 .default_output_length
7870 .map(usize::try_from)
7871 .transpose()
7872 .map_err(|_| {
7873 (
7874 "configured default output length does not fit this platform".into(),
7875 "model",
7876 )
7877 })?
7878 .unwrap_or(max_output);
7879 } else if request.params.max_new > max_output {
7880 return Err((
7881 format!(
7882 "max_tokens {} exceeds configured model maximum {max_output}",
7883 request.params.max_new
7884 ),
7885 "max_tokens",
7886 ));
7887 }
7888 }
7889
7890 if let (Some(max_prompt), Some(max_output), Some(requested_ctx)) =
7894 (max_prompt, max_output, request.params.max_ctx)
7895 {
7896 let operational_ctx = max_prompt
7897 .checked_add(max_output)
7898 .and_then(|value| value.checked_add(8))
7899 .ok_or_else(|| {
7900 (
7901 "configured model context envelope overflowed".into(),
7902 "model",
7903 )
7904 })?;
7905 let operational_ctx = caps
7906 .map(|caps| caps.context_length)
7907 .filter(|&context| context > 0)
7908 .map_or(operational_ctx, |context| operational_ctx.min(context));
7909 if requested_ctx > operational_ctx {
7910 return Err((
7911 format!(
7912 "max_ctx {requested_ctx} exceeds configured model envelope {operational_ctx}"
7913 ),
7914 "max_ctx",
7915 ));
7916 }
7917 }
7918 Ok(())
7919}
7920
7921fn effective_max_tokens(request: &worker::Request) -> Option<u64> {
7925 (request.params.max_new != worker::MAX_NEW_CTX_BOUNDED).then_some(request.params.max_new as u64)
7926}
7927
7928#[allow(clippy::too_many_arguments)]
7929fn start_request_receipt(
7930 st: &AppState,
7931 env: &Envelope,
7932 tenant: &auth::TenantCtx,
7933 model: &str,
7934 route: &'static str,
7935 lane: lanes::Lane,
7936 stream: bool,
7937 max_tokens: Option<u64>,
7938 reserved_ctx: Option<u64>,
7939 budget_permit: Option<metering::Permit>,
7940) -> Option<Box<dyn metering::Receipt>> {
7941 st.metering.as_ref().map(|accounting| {
7942 accounting.open(
7943 &metering::RequestMeta {
7944 request_id: &env.id,
7945 tenant: &tenant.tenant,
7946 principal: tenant.key_prefix.as_deref(),
7947 model,
7948 route,
7949 lane: lane.as_str(),
7950 stream,
7951 max_tokens,
7952 reserved_ctx,
7953 },
7954 budget_permit,
7955 )
7956 })
7957}
7958
7959fn arm_capture(
7965 mut receipt: Option<Box<dyn metering::Receipt>>,
7966 prompt: impl FnOnce() -> serde_json::Value,
7967) -> Option<Box<dyn metering::Receipt>> {
7968 if let Some(receipt) = receipt.as_mut()
7969 && receipt.wants_capture()
7970 {
7971 receipt.arm_capture(prompt());
7972 }
7973 receipt
7974}
7975
7976fn capture_chat_messages(messages: &[ChatMessage]) -> serde_json::Value {
7980 serde_json::Value::Array(
7981 messages
7982 .iter()
7983 .map(|message| {
7984 let mut row = json!({ "role": message.role, "content": message.content });
7985 if !message.tool_calls.is_empty() {
7986 row["tool_calls"] = serde_json::Value::Array(
7987 message
7988 .tool_calls
7989 .iter()
7990 .map(|call| {
7991 json!({
7992 "id": call.id,
7993 "function": {
7994 "name": call.function.name,
7995 "arguments": call.function.arguments,
7996 },
7997 })
7998 })
7999 .collect(),
8000 );
8001 }
8002 row
8003 })
8004 .collect(),
8005 )
8006}
8007
8008enum BudgetRejection {
8009 Invalid(String),
8010 Insufficient,
8011 Unenrolled,
8012 PrincipalCapped,
8015 Unavailable(String),
8016}
8017
8018impl BudgetRejection {
8019 fn into_response(self) -> (Response, &'static str) {
8020 match self {
8021 Self::Invalid(message) => (bad_request(&message, Some("prompt")), "invalid_request"),
8022 Self::Insufficient => (
8023 error_response_coded(
8024 StatusCode::PAYMENT_REQUIRED,
8025 "tenant prepaid balance is insufficient for this request",
8026 "insufficient_balance",
8027 None,
8028 Some("insufficient_balance"),
8029 ),
8030 "insufficient_balance",
8031 ),
8032 Self::Unenrolled => (
8033 error_response_coded(
8034 StatusCode::PAYMENT_REQUIRED,
8035 "tenant is not enrolled for prepaid billing",
8036 "tenant_not_enrolled",
8037 None,
8038 Some("tenant_not_enrolled"),
8039 ),
8040 "tenant_not_enrolled",
8041 ),
8042 Self::PrincipalCapped => (
8043 error_response_coded(
8044 StatusCode::PAYMENT_REQUIRED,
8045 "this API key's spend cap is reached; raise or clear the key's cap to continue",
8046 "key_spend_cap_reached",
8047 None,
8048 Some("key_spend_cap_reached"),
8049 ),
8050 "key_spend_cap_reached",
8051 ),
8052 Self::Unavailable(err) => {
8053 eprintln!("[budget] ERROR: admission unavailable: {err}");
8054 (
8055 error_response_coded(
8056 StatusCode::SERVICE_UNAVAILABLE,
8057 "tenant budget accounting is unavailable",
8058 "server_error",
8059 None,
8060 Some("tenant_budget_unavailable"),
8061 ),
8062 "tenant_budget_unavailable",
8063 )
8064 }
8065 }
8066 }
8067}
8068
8069fn prepare_budget_prompt(
8070 request: &mut Request,
8071 tokenizer: Option<&Tokenizer>,
8072) -> Result<usize, String> {
8073 if let Some(error) = worker::prompt_source_limit_error(request) {
8074 return Err(error);
8075 }
8076 if request.prepared_prompt.is_none() {
8077 if let Some(trace) = request.ttft.as_ref() {
8078 trace.mark_tokenize_start();
8079 }
8080 let prompt = if !request.prompt_ids.is_empty() {
8081 request.prompt_ids.clone()
8082 } else if !request.chat_turns.is_empty() {
8083 let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
8084 let plain = worker::plain_chat_render_path(
8091 &request.tools_json,
8092 &request.think,
8093 request.reasoning_effort.as_deref(),
8094 &request.chat_turns,
8095 tokenizer.has_qwen_effort_ladder(),
8096 );
8097 let rendered = if plain {
8098 let messages: Vec<_> = request
8099 .chat_turns
8100 .iter()
8101 .map(|turn| (turn.role.as_str(), turn.content.as_str()))
8102 .collect();
8103 tokenizer.apply_chat_template(&messages, true)
8104 } else {
8105 tokenizer
8106 .apply_chat_template_tools_ex(
8107 &request.chat_turns,
8108 true,
8109 &request.tools_json,
8110 &request.tools_struct,
8111 request.think,
8112 request.reasoning_effort.as_deref(),
8113 )
8114 .map_err(|err| format!("chat template: {err}"))?
8115 };
8116 tokenizer.encode(&rendered, true)
8117 } else if request.chat {
8118 let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
8119 let rendered =
8120 tokenizer.apply_chat_template(&[("user", request.prompt_text.as_str())], true);
8121 tokenizer.encode(&rendered, true)
8122 } else {
8123 let tokenizer = tokenizer.ok_or("reservation tokenizer is unavailable")?;
8124 tokenizer.encode(&request.prompt_text, true)
8125 };
8126 if prompt.is_empty() {
8127 return Err("empty prompt after tokenization".into());
8128 }
8129 if let Some(trace) = request.ttft.as_ref() {
8130 trace.mark_tokenize_end(prompt.len());
8131 }
8132 request.prepared_prompt = Some(prompt);
8133 }
8134 let prompt_tokens = request
8135 .prepared_prompt
8136 .as_ref()
8137 .expect("budget prompt was prepared")
8138 .len();
8139 if let Some(limit) = request.max_prompt_tokens
8140 && prompt_tokens > limit
8141 {
8142 return Err(format!(
8143 "prompt ({prompt_tokens} tok) exceeds configured model maximum ({limit})"
8144 ));
8145 }
8146 Ok(prompt_tokens)
8147}
8148
8149fn budget_completion_bound(
8150 request: &Request,
8151 prompt_tokens: usize,
8152 caps: Option<&ModelCaps>,
8153) -> Result<usize, String> {
8154 let max_new = request.params.max_new;
8155 let requested_ctx = match (request.params.max_ctx, max_new) {
8156 (Some(cap), _) => cap,
8157 (None, worker::MAX_NEW_CTX_BOUNDED) => {
8158 let server_ctx = std::env::var("MEMRA_CTX")
8159 .ok()
8160 .and_then(|value| value.parse().ok())
8161 .unwrap_or(8192usize);
8162 let mut cap = server_ctx;
8163 if prompt_tokens.saturating_add(16) > cap {
8164 cap = prompt_tokens.saturating_add(server_ctx);
8165 }
8166 cap
8167 }
8168 (None, max_new) => prompt_tokens
8169 .checked_add(max_new)
8170 .and_then(|value| value.checked_add(8))
8171 .ok_or_else(|| "request context bound overflowed".to_string())?,
8172 };
8173 let ctx_cap = caps
8174 .map(|caps| caps.context_length)
8175 .filter(|&context| context > 0)
8176 .map_or(requested_ctx, |context| requested_ctx.min(context));
8177 if prompt_tokens >= ctx_cap {
8178 return Err(format!(
8179 "prompt ({prompt_tokens} tok) >= context cap ({ctx_cap})"
8180 ));
8181 }
8182 Ok(max_new.min(ctx_cap - prompt_tokens))
8183}
8184
8185struct BudgetAdmission {
8190 permit: Option<metering::Permit>,
8191 reserved_ctx: Option<u64>,
8192}
8193
8194impl std::fmt::Debug for BudgetAdmission {
8196 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
8197 f.debug_struct("BudgetAdmission")
8198 .field("permit", &self.permit.is_some())
8199 .field("reserved_ctx", &self.reserved_ctx)
8200 .finish()
8201 }
8202}
8203
8204fn admit_tenant_budget(
8205 st: &AppState,
8206 tenant: &auth::TenantCtx,
8207 request: &mut Request,
8208) -> Result<BudgetAdmission, BudgetRejection> {
8209 let Some(accounting) = st.metering.as_ref().filter(|m| m.enforces_limits()) else {
8210 return Ok(BudgetAdmission {
8211 permit: None,
8212 reserved_ctx: None,
8213 });
8214 };
8215 match accounting.is_limited(&tenant.tenant) {
8216 Ok(false) => return Err(BudgetRejection::Unenrolled),
8217 Ok(true) => {}
8218 Err(metering::AdmitError::Unavailable(err)) => {
8219 return Err(BudgetRejection::Unavailable(err));
8220 }
8221 Err(other) => {
8222 return Err(BudgetRejection::Unavailable(format!(
8223 "unexpected budget enrollment result: {other:?}"
8224 )));
8225 }
8226 }
8227 let tokenizer = st
8228 .budget_tokenizers
8229 .as_ref()
8230 .and_then(|tokenizers| tokenizers.get(&request.model))
8231 .map(Arc::as_ref);
8232 if request.prompt_ids.is_empty() && tokenizer.is_none() {
8233 return Err(BudgetRejection::Unavailable(format!(
8234 "no reservation tokenizer for model {:?}",
8235 request.model
8236 )));
8237 }
8238 let prompt_tokens =
8239 prepare_budget_prompt(request, tokenizer).map_err(BudgetRejection::Invalid)?;
8240 let completion_tokens =
8241 budget_completion_bound(request, prompt_tokens, st.caps.get(&request.model))
8242 .map_err(BudgetRejection::Invalid)?;
8243 let prompt_tokens = u64::try_from(prompt_tokens)
8244 .map_err(|_| BudgetRejection::Unavailable("prompt token count exceeds u64".into()))?;
8245 let completion_tokens = u64::try_from(completion_tokens)
8246 .map_err(|_| BudgetRejection::Unavailable("completion token bound exceeds u64".into()))?;
8247 match accounting.reserve(
8248 &tenant.tenant,
8249 tenant.key_prefix.as_deref(),
8250 &request.model,
8251 prompt_tokens,
8252 completion_tokens,
8253 ) {
8254 Ok(permit) => Ok(BudgetAdmission {
8255 permit,
8256 reserved_ctx: Some(prompt_tokens.saturating_add(completion_tokens)),
8257 }),
8258 Err(metering::AdmitError::Insufficient) => Err(BudgetRejection::Insufficient),
8259 Err(metering::AdmitError::PrincipalCapped) => Err(BudgetRejection::PrincipalCapped),
8260 Err(metering::AdmitError::Blocked) => Err(BudgetRejection::Insufficient),
8264 Err(metering::AdmitError::Unenrolled) => Err(BudgetRejection::Unenrolled),
8265 Err(metering::AdmitError::Unavailable(err)) => Err(BudgetRejection::Unavailable(err)),
8266 }
8267}
8268
8269fn request_ledger_error_response() -> Response {
8270 error_response_coded(
8271 StatusCode::INTERNAL_SERVER_ERROR,
8272 "request completion could not be committed to the billing ledger",
8273 "server_error",
8274 None,
8275 Some("request_ledger_unavailable"),
8276 )
8277}
8278
8279fn request_ledger_error_body() -> serde_json::Value {
8280 error_body(
8281 "request completion could not be committed to the billing ledger",
8282 "server_error",
8283 None,
8284 Some("request_ledger_unavailable"),
8285 )
8286}
8287
8288fn ledger_rejected(
8289 mut receipt: Option<Box<dyn metering::Receipt>>,
8290 response: Response,
8291 error_code: &str,
8292 request_id: &str,
8293) -> Response {
8294 let status = response.status().as_u16();
8295 if let Some(receipt) = receipt.as_mut()
8296 && let Err(err) = receipt.reject(status, error_code)
8297 {
8298 eprintln!("[ledger] ERROR: request {request_id} rejection receipt failed: {err}");
8299 return with_request_id(request_id, request_ledger_error_response());
8300 }
8301 with_request_id(request_id, response)
8302}
8303
8304fn ledger_unbilled(
8309 mut receipt: Option<Box<dyn metering::Receipt>>,
8310 response: Response,
8311 outcome: &'static str,
8312 error_code: &str,
8313 request_id: &str,
8314) -> Response {
8315 let status = response.status().as_u16();
8316 if let Some(receipt) = receipt.as_mut()
8317 && let Err(err) = receipt.settle_unbilled(outcome, status, error_code)
8318 {
8319 eprintln!("[ledger] ERROR: request {request_id} {outcome} receipt failed: {err}");
8320 return with_request_id(request_id, request_ledger_error_response());
8321 }
8322 with_request_id(request_id, response)
8323}
8324
8325fn engine_error_code(class: worker::ErrClass) -> &'static str {
8326 use worker::ErrClass as C;
8327 match class {
8328 C::InvalidRequest => "invalid_request",
8329 C::ContextLength => "context_length_exceeded",
8330 C::ModelNotFound => "model_not_found",
8331 C::RateLimit => "rate_limit_exceeded",
8332 C::Overloaded => "overloaded",
8333 C::Engine => "engine_error",
8334 }
8335}
8336
8337fn model_not_found_response(models: &[String], requested: &str) -> Response {
8357 error_response_coded(
8358 StatusCode::BAD_REQUEST,
8359 &format!("unknown model {requested:?}; loaded: {models:?}"),
8360 "invalid_request_error",
8361 Some("model"),
8362 Some("model_not_found"),
8363 )
8364}
8365
8366fn validate_prompt_ids(ids: &[u32], caps: Option<&ModelCaps>) -> Result<(), String> {
8375 let Some(n_vocab) = caps.map(|c| c.n_vocab).filter(|&n| n > 0) else {
8376 return Ok(());
8377 };
8378 if let Some((pos, &id)) = ids
8379 .iter()
8380 .enumerate()
8381 .find(|&(_, &id)| id as usize >= n_vocab)
8382 {
8383 return Err(format!(
8384 "prompt_ids[{pos}] = {id} is out of vocabulary (model vocab size {n_vocab})"
8385 ));
8386 }
8387 Ok(())
8388}
8389
8390#[cfg(test)]
8391mod prompt_ids_tests {
8392 use super::*;
8393
8394 #[test]
8395 fn prompt_ids_are_bounded_by_the_model_vocab_at_intake() {
8396 let caps = ModelCaps {
8397 n_vocab: 8,
8398 ..Default::default()
8399 };
8400 assert!(validate_prompt_ids(&[0, 3, 7], Some(&caps)).is_ok());
8402 assert!(validate_prompt_ids(&[], Some(&caps)).is_ok());
8403 let err = validate_prompt_ids(&[1, 8, 2], Some(&caps)).unwrap_err();
8405 assert!(err.contains("prompt_ids[1] = 8"), "{err}");
8406 assert!(err.contains("vocab size 8"), "{err}");
8407 let err = validate_prompt_ids(&[u32::MAX], Some(&caps)).unwrap_err();
8408 assert!(err.contains("4294967295"), "{err}");
8409 let unknown = ModelCaps::default();
8411 assert!(validate_prompt_ids(&[u32::MAX], Some(&unknown)).is_ok());
8412 assert!(validate_prompt_ids(&[u32::MAX], None).is_ok());
8413 }
8414}
8415
8416fn canonical_model_id(models: &[String], requested: &str) -> Option<String> {
8417 if models.iter().any(|m| m == requested) {
8418 return Some(requested.to_string());
8419 }
8420 if requested.is_empty() || requested.contains('/') {
8421 return None;
8422 }
8423 let mut matches = models.iter().filter(|m| {
8424 m.rsplit('/')
8425 .next()
8426 .is_some_and(|suffix| suffix == requested)
8427 });
8428 match (matches.next(), matches.next()) {
8429 (Some(only), None) => Some(only.clone()),
8430 _ => None,
8431 }
8432}
8433
8434async fn completions_admitted(
8435 state: State<AppState>,
8436 headers: axum::http::HeaderMap,
8437 trace: Option<Extension<TtftRequestTrace>>,
8438 AdmittedJson(req, admission): AdmittedJson<CompletionReq>,
8439) -> Response {
8440 completions_with_admission(state, headers, trace, Json(req), Some(admission)).await
8441}
8442
8443#[cfg(test)]
8444async fn completions(
8445 State(st): State<AppState>,
8446 headers: axum::http::HeaderMap,
8447 trace: Option<Extension<TtftRequestTrace>>,
8448 request: Json<CompletionReq>,
8449) -> Response {
8450 completions_with_admission(State(st), headers, trace, request, None).await
8451}
8452
8453async fn completions_with_admission(
8454 State(st): State<AppState>,
8455 headers: axum::http::HeaderMap,
8456 trace: Option<Extension<TtftRequestTrace>>,
8457 Json(mut req): Json<CompletionReq>,
8458 mut body_admission: Option<BodyAdmissionLease>,
8459) -> Response {
8460 let env = Envelope::new(false);
8461 if let Err(msg) = req.stop.validate() {
8462 return with_request_id(&env.id, bad_request(&msg, Some("stop")));
8463 }
8464 if let Err(msg) = validate_client_identifier(req.trace_id.as_deref(), "trace_id") {
8465 return with_request_id(&env.id, bad_request(&msg, Some("trace_id")));
8466 }
8467 match canonical_model_id(&st.models, &req.model) {
8468 Some(canonical) => req.model = canonical,
8469 None => {
8470 return with_request_id(&env.id, model_not_found_response(&st.models, &req.model));
8471 }
8472 }
8473 let ttft = trace.and_then(|Extension(trace)| trace.0);
8476 if let Some(trace) = ttft.as_ref() {
8477 trace.mark_parsed();
8478 trace.bind_request(&env.id, &req.model);
8479 }
8480 let tenant = match authenticate(&st.api_auth, &headers) {
8481 Ok(t) => t,
8482 Err(resp) => return with_request_id(&env.id, resp),
8483 };
8484 let cache_ns = match tenant_namespace(&tenant, &req.cache_salt) {
8485 Ok(ns) => ns,
8486 Err(msg) => return with_request_id(&env.id, bad_request(msg, Some("cache_salt"))),
8487 };
8488 if let Err((msg, param)) = reject_unsupported(&[
8490 (
8491 "logit_bias",
8492 req.logit_bias.is_some(),
8493 " (device-side sampling has no bias hook yet)",
8494 ),
8495 ("logprobs", req.logprobs.is_some(), ""),
8496 (
8497 "n",
8498 req.n.is_some_and(|n| n != 1),
8499 " for n != 1 (single choice only)",
8500 ),
8501 (
8502 "best_of",
8503 req.best_of.is_some_and(|n| n != 1),
8504 " (single choice only)",
8505 ),
8506 ]) {
8507 return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
8508 }
8509 if let Err(msg) = validate_prompt_ids(&req.prompt_ids, st.caps.get(&req.model)) {
8512 return with_request_id(&env.id, bad_request(&msg, Some("prompt_ids")));
8513 }
8514 let deadline = match parse_timeout_ms(req.timeout_ms.as_ref()) {
8517 Ok(ms) => RequestDeadline::starting_now(ms),
8518 Err(msg) => return with_request_id(&env.id, bad_request(&msg, Some("timeout_ms"))),
8519 };
8520 let lane = match lane_for_tenant(&headers, &tenant) {
8521 Ok(l) => l,
8522 Err(resp) => return resp,
8523 };
8524 let (tx, rx) = worker::event_channel();
8525 let model = req.model.clone();
8526 let stream = req.stream;
8527 let affinity = match affinity_key(&req.session_id, &req.user, &headers) {
8528 Ok(affinity) => affinity,
8529 Err(msg) => return with_request_id(&env.id, bad_request(&msg, Some("session_id"))),
8530 };
8531 let mut request = build_request_with_trace(
8532 &req,
8533 tx,
8534 lane,
8535 affinity,
8536 ttft.clone(),
8537 st.sampling_defaults(&model).for_mode(ThinkMode::Default),
8541 );
8542 request.cache_ns = cache_ns;
8543 request.request_id = env.id.clone();
8544 request.wire_deadline = Some(deadline.at.into_std());
8547 if let Err((message, param)) = apply_model_request_limits(
8548 &mut request,
8549 st.openrouter_metadata.get(&model),
8550 st.caps.get(&model),
8551 ) {
8552 return with_request_id(&env.id, bad_request(&message, Some(param)));
8553 }
8554 if let Err(msg) = nonstream_deadline_gate(
8559 &request,
8560 req.stream,
8561 deadline,
8562 req.max_tokens.is_some(),
8563 st.budget_tokenizers
8564 .as_ref()
8565 .and_then(|t| t.get(&req.model))
8566 .map(Arc::as_ref),
8567 ) {
8568 return with_request_id(
8569 &env.id,
8570 error_response_coded(
8571 StatusCode::BAD_REQUEST,
8572 &msg,
8573 "invalid_request_error",
8574 Some("max_tokens"),
8575 Some("nonstream_deadline_infeasible"),
8576 ),
8577 );
8578 }
8579 if draining() {
8582 let receipt = start_request_receipt(
8583 &st,
8584 &env,
8585 &tenant,
8586 &req.model,
8587 "/v1/completions",
8588 lane,
8589 req.stream,
8590 effective_max_tokens(&request),
8591 None,
8592 None,
8593 );
8594 return ledger_rejected(receipt, drain_response(), "draining", &env.id);
8595 }
8596 let budget = match admit_tenant_budget(&st, &tenant, &mut request) {
8597 Ok(budget) => budget,
8598 Err(rejection) => {
8599 let (response, error_code) = rejection.into_response();
8600 let receipt = start_request_receipt(
8601 &st,
8602 &env,
8603 &tenant,
8604 &req.model,
8605 "/v1/completions",
8606 lane,
8607 req.stream,
8608 effective_max_tokens(&request),
8609 None,
8610 None,
8611 );
8612 return ledger_rejected(receipt, response, error_code, &env.id);
8613 }
8614 };
8615 let receipt = start_request_receipt(
8616 &st,
8617 &env,
8618 &tenant,
8619 &req.model,
8620 "/v1/completions",
8621 lane,
8622 req.stream,
8623 effective_max_tokens(&request),
8624 budget.reserved_ctx,
8625 budget.permit,
8626 );
8627 let receipt = arm_capture(receipt, || json!({ "prompt": req.prompt }));
8628 let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
8631 Ok(slot) => slot,
8632 Err(resp) => {
8633 return ledger_rejected(receipt, resp, "rate_limit_exceeded", &env.id);
8634 }
8635 };
8636 if let Some(admission) = body_admission.as_mut() {
8637 admission.release();
8638 }
8639 let pending_admit = match reserve_pending_admit(&st, lane, &rl, deadline) {
8642 Ok(guard) => guard,
8643 Err((resp, outcome)) => {
8644 return ledger_unbilled(receipt, rl.attach(resp), outcome, outcome, &env.id);
8645 }
8646 };
8647 meter_admit(&env, &tenant, &model, lane);
8648 let stop_strings = request.stop_strings.clone();
8649
8650 if let Some(trace) = ttft.as_ref() {
8655 trace.mark_submitted();
8656 }
8657 if st.cmd_tx.send(Cmd::Generate(Box::new(request))).is_err() {
8658 drop(pending_admit);
8659 return ledger_rejected(
8660 receipt,
8661 rl.attach(worker_unavailable_response()),
8662 "worker_unavailable",
8663 &env.id,
8664 );
8665 }
8666 pending_admit.commit();
8667 let rx = match tokio::time::timeout_at(deadline.at, peek_admission(rx)).await {
8671 Ok(Ok(rx)) => rx,
8672 Ok(Err((resp, error_code))) => {
8673 return ledger_rejected(receipt, rl.attach(resp), error_code, &env.id);
8674 }
8675 Err(_) => {
8676 return ledger_unbilled(
8677 receipt,
8678 rl.attach(deadline_exceeded_response(deadline.ms, stream)),
8679 "deadline_exceeded",
8680 "deadline_exceeded",
8681 &env.id,
8682 );
8683 }
8684 };
8685
8686 let resp = if stream {
8687 let rx = match peek_first_token(rx, deadline).await {
8691 Ok(rx) => rx,
8692 Err(()) => {
8693 return ledger_unbilled(
8694 receipt,
8695 rl.attach(deadline_exceeded_response(deadline.ms, true)),
8696 "deadline_exceeded",
8697 "deadline_exceeded",
8698 &env.id,
8699 );
8700 }
8701 };
8702 sse_response_with_receipt(
8703 rx,
8704 model,
8705 false,
8706 None,
8707 env.clone(),
8708 stop_strings,
8709 Some(guard),
8710 receipt,
8711 )
8712 .into_response()
8713 } else {
8714 let mut receipt = receipt;
8720 let resp = blocking_response_with_receipt(
8721 rx,
8722 model,
8723 false,
8724 stop_strings,
8725 None,
8726 env.clone(),
8727 &mut receipt,
8728 Some(deadline),
8729 )
8730 .await;
8731 drop(guard); resp.into_response()
8733 };
8734 rl.attach(with_request_id(&env.id, resp))
8735}
8736
8737async fn chat_completions_admitted(
8738 state: State<AppState>,
8739 headers: axum::http::HeaderMap,
8740 trace: Option<Extension<TtftRequestTrace>>,
8741 AdmittedJson(req, admission): AdmittedJson<ChatCompletionReq>,
8742) -> Response {
8743 chat_completions_with_admission(state, headers, trace, Json(req), Some(admission)).await
8744}
8745
8746#[cfg(test)]
8747async fn chat_completions(
8748 State(st): State<AppState>,
8749 headers: axum::http::HeaderMap,
8750 trace: Option<Extension<TtftRequestTrace>>,
8751 request: Json<ChatCompletionReq>,
8752) -> Response {
8753 chat_completions_with_admission(State(st), headers, trace, request, None).await
8754}
8755
8756async fn chat_completions_with_admission(
8757 State(st): State<AppState>,
8758 headers: axum::http::HeaderMap,
8759 trace: Option<Extension<TtftRequestTrace>>,
8760 Json(mut req): Json<ChatCompletionReq>,
8761 mut body_admission: Option<BodyAdmissionLease>,
8762) -> Response {
8763 let env = Envelope::new(true);
8764 match canonical_model_id(&st.models, &req.model) {
8769 Some(canonical) => req.model = canonical,
8770 None => {
8771 return with_request_id(&env.id, model_not_found_response(&st.models, &req.model));
8772 }
8773 }
8774 let ttft = trace.and_then(|Extension(trace)| trace.0);
8775 if let Some(trace) = ttft.as_ref() {
8776 trace.mark_parsed();
8777 trace.bind_request(&env.id, &req.model);
8778 }
8779 let tenant = match authenticate(&st.api_auth, &headers) {
8780 Ok(t) => t,
8781 Err(resp) => return with_request_id(&env.id, resp),
8782 };
8783 let cache_ns = match tenant_namespace(&tenant, &req.cache_salt) {
8784 Ok(ns) => ns,
8785 Err(msg) => return with_request_id(&env.id, bad_request(msg, Some("cache_salt"))),
8786 };
8787 if req.messages.is_empty()
8788 || req.messages.iter().any(|message| {
8789 !matches!(
8790 message.role.as_str(),
8791 "system" | "developer" | "user" | "assistant" | "tool"
8792 )
8793 })
8794 {
8795 return with_request_id(
8796 &env.id,
8797 bad_request(
8798 "messages must use system/developer/user/assistant/tool roles",
8799 Some("messages"),
8800 ),
8801 );
8802 }
8803 if let Err((msg, param)) = reject_unsupported(&[
8808 (
8809 "logit_bias",
8810 req.logit_bias.is_some(),
8811 " (device-side sampling has no bias hook yet)",
8812 ),
8813 (
8814 "logprobs",
8815 req.logprobs
8816 .as_ref()
8817 .is_some_and(|v| v.as_bool() != Some(false)),
8818 "",
8819 ),
8820 ("top_logprobs", req.top_logprobs.is_some(), ""),
8821 (
8822 "n",
8823 req.n.is_some_and(|n| n != 1),
8824 " for n != 1 (single choice only)",
8825 ),
8826 ]) {
8827 return with_request_id(&env.id, bad_request(&msg, Some(¶m)));
8828 }
8829 let deadline = match parse_timeout_ms(req.timeout_ms.as_ref()) {
8832 Ok(ms) => RequestDeadline::starting_now(ms),
8833 Err(msg) => return with_request_id(&env.id, bad_request(&msg, Some("timeout_ms"))),
8834 };
8835 let lane = match lane_for_tenant(&headers, &tenant) {
8836 Ok(l) => l,
8837 Err(resp) => return resp,
8838 };
8839 let model = req.model.clone();
8840 let stream = req.stream;
8841 let capture_prompt = st
8844 .metering
8845 .as_ref()
8846 .filter(|m| m.captures(&tenant.tenant))
8847 .map(|_| capture_chat_messages(&req.messages));
8848 let declared_max_tokens = req.max_tokens.is_some();
8852 let vision_preprocess_permit = match try_vision_preprocess(request_has_vision(&req)) {
8856 Ok(permit) => permit,
8857 Err(response) => return with_request_id(&env.id, response),
8858 };
8859 let (tx, rx) = worker::event_channel();
8860 let affinity = match affinity_key(&req.session_id, &req.user, &headers) {
8861 Ok(affinity) => affinity,
8862 Err(msg) => return with_request_id(&env.id, bad_request(&msg, Some("session_id"))),
8863 };
8864 let mut plan = match build_chat_request_with_trace(
8865 req,
8866 st.caps.get(&model),
8867 tx,
8868 lane,
8869 affinity,
8870 ttft.clone(),
8871 st.openrouter_metadata
8872 .get(&model)
8873 .and_then(|m| m.default_reasoning_effort.as_deref()),
8874 &st.sampling_defaults(&model),
8875 ) {
8876 Ok(plan) => plan,
8877 Err(err) => {
8878 return with_request_id(&env.id, bad_request(&err, None));
8879 }
8880 };
8881 plan.request.cache_ns = cache_ns;
8882 plan.request.request_id = env.id.clone();
8883 plan.request.wire_deadline = Some(deadline.at.into_std());
8884 if let Err((message, param)) = apply_model_request_limits(
8885 &mut plan.request,
8886 st.openrouter_metadata.get(&model),
8887 st.caps.get(&model),
8888 ) {
8889 return with_request_id(&env.id, bad_request(&message, Some(param)));
8890 }
8891 if let Err(msg) = nonstream_deadline_gate(
8894 &plan.request,
8895 stream,
8896 deadline,
8897 declared_max_tokens,
8898 st.budget_tokenizers
8899 .as_ref()
8900 .and_then(|t| t.get(&model))
8901 .map(Arc::as_ref),
8902 ) {
8903 return with_request_id(
8904 &env.id,
8905 error_response_coded(
8906 StatusCode::BAD_REQUEST,
8907 &msg,
8908 "invalid_request_error",
8909 Some("max_tokens"),
8910 Some("nonstream_deadline_infeasible"),
8911 ),
8912 );
8913 }
8914 plan.vision_memory = match reserve_vision_memory(&plan) {
8915 Ok(permit) => permit,
8916 Err(err) => {
8917 return with_request_id(&env.id, vision_memory_error_response(err, Some("messages")));
8918 }
8919 };
8920 if draining() {
8923 let receipt = start_request_receipt(
8924 &st,
8925 &env,
8926 &tenant,
8927 &model,
8928 "/v1/chat/completions",
8929 lane,
8930 stream,
8931 effective_max_tokens(&plan.request),
8932 None,
8933 None,
8934 );
8935 return ledger_rejected(receipt, drain_response(), "draining", &env.id);
8936 }
8937 let budget = match admit_tenant_budget(&st, &tenant, &mut plan.request) {
8938 Ok(budget) => budget,
8939 Err(rejection) => {
8940 let (response, error_code) = rejection.into_response();
8941 let receipt = start_request_receipt(
8942 &st,
8943 &env,
8944 &tenant,
8945 &model,
8946 "/v1/chat/completions",
8947 lane,
8948 stream,
8949 effective_max_tokens(&plan.request),
8950 None,
8951 None,
8952 );
8953 return ledger_rejected(receipt, response, error_code, &env.id);
8954 }
8955 };
8956 let receipt = start_request_receipt(
8957 &st,
8958 &env,
8959 &tenant,
8960 &model,
8961 "/v1/chat/completions",
8962 lane,
8963 stream,
8964 effective_max_tokens(&plan.request),
8965 budget.reserved_ctx,
8966 budget.permit,
8967 );
8968 let receipt = if let Some(prompt) = capture_prompt {
8969 arm_capture(receipt, move || prompt)
8970 } else {
8971 receipt
8972 };
8973 let (guard, rl) = match acquire_request_slot(&st, lane, &tenant, &env) {
8977 Ok(slot) => slot,
8978 Err(resp) => {
8979 return ledger_rejected(receipt, resp, "rate_limit_exceeded", &env.id);
8980 }
8981 };
8982 if let Some(admission) = body_admission.as_mut() {
8983 admission.release();
8984 }
8985 let pending_admit = match reserve_pending_admit(&st, lane, &rl, deadline) {
8988 Ok(guard) => guard,
8989 Err((resp, outcome)) => {
8990 return ledger_unbilled(receipt, rl.attach(resp), outcome, outcome, &env.id);
8991 }
8992 };
8993 if let Err(err) = decode_pending_vision(&mut plan) {
8998 return ledger_rejected(
8999 receipt,
9000 rl.attach(bad_request(&err, Some("messages"))),
9001 "invalid_request_error",
9002 &env.id,
9003 );
9004 }
9005 plan.request.vision_memory = plan.vision_memory.take();
9006 drop(vision_preprocess_permit);
9007 let constraint_ready = if plan.request.grammar.is_some() {
9008 let (ready_tx, ready_rx) = tokio::sync::oneshot::channel();
9009 plan.request.constraint_ready = Some(ready_tx);
9010 Some(ready_rx)
9011 } else {
9012 None
9013 };
9014 meter_admit(&env, &tenant, &model, lane);
9015 let stop_strings = plan.request.stop_strings.clone();
9016 if let Some(trace) = ttft.as_ref() {
9018 trace.mark_submitted();
9019 }
9020 if st
9021 .cmd_tx
9022 .send(Cmd::Generate(Box::new(plan.request)))
9023 .is_err()
9024 {
9025 drop(pending_admit);
9026 return ledger_rejected(
9027 receipt,
9028 rl.attach(worker_unavailable_response()),
9029 "worker_unavailable",
9030 &env.id,
9031 );
9032 }
9033 pending_admit.commit();
9034 if let Some(ready) = constraint_ready {
9040 let bound = constrained::CONSTRAINT_COMPILE_TIMEOUT.min(deadline.remaining());
9041 match tokio::time::timeout(bound, ready).await {
9042 Ok(Ok(Ok(()))) => {}
9043 Ok(Ok(Err(err))) => {
9044 return ledger_rejected(
9045 receipt,
9046 rl.attach(engine_error_response(&err)),
9047 engine_error_code(err.class),
9048 &env.id,
9049 );
9050 }
9051 Ok(Err(_)) => {
9052 return ledger_rejected(
9053 receipt,
9054 rl.attach(worker_unavailable_response()),
9055 "worker_unavailable",
9056 &env.id,
9057 );
9058 }
9059 Err(_) if deadline.remaining().is_zero() => {
9060 return ledger_unbilled(
9061 receipt,
9062 rl.attach(deadline_exceeded_response(deadline.ms, stream)),
9063 "deadline_exceeded",
9064 "deadline_exceeded",
9065 &env.id,
9066 );
9067 }
9068 Err(_) => {
9069 return ledger_rejected(
9070 receipt,
9071 rl.attach(engine_error_response(&worker::constraint_timeout_error())),
9072 "constraint_compile_timeout",
9073 &env.id,
9074 );
9075 }
9076 }
9077 }
9078 let rx = match tokio::time::timeout_at(deadline.at, peek_admission(rx)).await {
9080 Ok(Ok(rx)) => rx,
9081 Ok(Err((resp, error_code))) => {
9082 return ledger_rejected(receipt, rl.attach(resp), error_code, &env.id);
9083 }
9084 Err(_) => {
9085 return ledger_unbilled(
9086 receipt,
9087 rl.attach(deadline_exceeded_response(deadline.ms, stream)),
9088 "deadline_exceeded",
9089 "deadline_exceeded",
9090 &env.id,
9091 );
9092 }
9093 };
9094 let resp = if stream {
9095 let rx = match peek_first_token(rx, deadline).await {
9097 Ok(rx) => rx,
9098 Err(()) => {
9099 return ledger_unbilled(
9100 receipt,
9101 rl.attach(deadline_exceeded_response(deadline.ms, true)),
9102 "deadline_exceeded",
9103 "deadline_exceeded",
9104 &env.id,
9105 );
9106 }
9107 };
9108 sse_response_with_receipt(
9109 rx,
9110 model,
9111 true,
9112 plan.parser,
9113 env.clone(),
9114 stop_strings,
9115 Some(guard),
9116 receipt,
9117 )
9118 .into_response()
9119 } else {
9120 let mut receipt = receipt;
9123 let resp = blocking_response_with_receipt(
9124 rx,
9125 model,
9126 true,
9127 stop_strings,
9128 plan.parser,
9129 env.clone(),
9130 &mut receipt,
9131 Some(deadline),
9132 )
9133 .await;
9134 drop(guard); resp.into_response()
9136 };
9137 rl.attach(with_request_id(&env.id, resp))
9138}
9139
9140#[cfg(test)]
9149fn sse_response(
9150 rx: worker::EventReceiver,
9151 model: String,
9152 chat: bool,
9153 parser: Option<ToolStreamParser>,
9154 env: Envelope,
9155 stop_strings: Vec<String>,
9156 guard: Option<InflightGuard>,
9157) -> Sse<impl futures_core::Stream<Item = Result<SseEvent, std::convert::Infallible>>> {
9158 sse_response_with_receipt(rx, model, chat, parser, env, stop_strings, guard, None)
9159}
9160
9161#[allow(clippy::too_many_arguments)] fn sse_response_with_receipt(
9163 mut rx: worker::EventReceiver,
9164 model: String,
9165 chat: bool,
9166 mut parser: Option<ToolStreamParser>,
9167 env: Envelope,
9168 stop_strings: Vec<String>,
9169 guard: Option<InflightGuard>,
9170 mut receipt: Option<Box<dyn metering::Receipt>>,
9171) -> Sse<impl futures_core::Stream<Item = Result<SseEvent, std::convert::Infallible>>> {
9172 let mut scrub = (!stop_strings.is_empty() && (chat || openai_compat()))
9176 .then(|| StopScrubber::new(stop_strings));
9177 let stream = async_stream::stream! {
9178 let _guard = guard;
9181 let mut call_index: usize = 0;
9182 let mut role_sent = false;
9185 macro_rules! chat_chunk {
9186 ($delta:expr, $finish:expr) => {{
9187 let mut delta = $delta;
9188 if chat && !role_sent {
9189 role_sent = true;
9190 delta["role"] = json!("assistant");
9191 }
9192 env.stamp(json!({ "object": "chat.completion.chunk", "model": model,
9193 "choices": [{ "index": 0, "delta": delta,
9194 "finish_reason": $finish }] }))
9195 .to_string()
9196 }};
9197 }
9198 macro_rules! piece_chunks {
9200 ($piece:expr) => {{
9201 let mut payloads: Vec<String> = Vec::new();
9202 match $piece {
9203 Piece::Content(text) => {
9204 let text = match scrub.as_mut() {
9205 Some(sc) => sc.push(&text),
9206 None => text,
9207 };
9208 if !text.is_empty() {
9209 payloads.push(chat_chunk!(json!({ "content": text }),
9210 serde_json::Value::Null));
9211 }
9212 }
9213 Piece::Reasoning(text) => payloads.push(
9217 chat_chunk!(json!({ "reasoning": text }), serde_json::Value::Null)),
9218 Piece::Call(call) => {
9219 payloads.push(chat_chunk!(json!({ "tool_calls": [{
9220 "index": call_index, "id": call.id, "type": "function",
9221 "function": { "name": call.name, "arguments": "" } }] }),
9222 serde_json::Value::Null));
9223 payloads.push(chat_chunk!(json!({ "tool_calls": [{
9224 "index": call_index,
9225 "function": { "arguments": call.arguments } }] }),
9226 serde_json::Value::Null));
9227 call_index += 1;
9228 }
9229 }
9230 payloads
9231 }};
9232 }
9233 let mut terminal = false;
9237 while let Some(ev) = rx.recv().await {
9238 match ev {
9239 Event::PromptCapture { .. } => {} Event::PromptUsage { n_prompt, n_cached } => {
9241 if let Some(receipt) = receipt.as_mut()
9242 && let Err(err) = receipt.record_prompt_usage(
9243 n_prompt as u64,
9244 n_cached as u64,
9245 )
9246 {
9247 eprintln!(
9248 "[ledger] ERROR: request {} partial prompt receipt failed: {err}",
9249 env.id
9250 );
9251 let _ = receipt.reject(500, "request_ledger_unavailable");
9254 let payload = request_ledger_error_body().to_string();
9255 if chat || openai_compat() {
9256 yield Ok(SseEvent::default().data(payload));
9257 yield Ok(SseEvent::default().data("[DONE]"));
9258 } else {
9259 yield Ok(SseEvent::default().event("error").data(payload));
9260 }
9261 terminal = true;
9262 break;
9263 }
9264 }
9265 Event::Token { id, text } => {
9266 if let Some(receipt) = receipt.as_mut()
9267 && let Err(err) = receipt.record_completion_token()
9268 {
9269 eprintln!(
9270 "[ledger] ERROR: request {} partial completion receipt failed: {err}",
9271 env.id
9272 );
9273 let _ = receipt.reject(500, "request_ledger_unavailable");
9274 let payload = request_ledger_error_body().to_string();
9275 if chat || openai_compat() {
9276 yield Ok(SseEvent::default().data(payload));
9277 yield Ok(SseEvent::default().data("[DONE]"));
9278 } else {
9279 yield Ok(SseEvent::default().event("error").data(payload));
9280 }
9281 terminal = true;
9282 break;
9283 }
9284 if let Some(receipt) = receipt.as_mut() {
9287 receipt.capture_completion_delta(&text);
9288 }
9289 if let Some(p) = parser.as_mut() {
9290 for piece in p.push(&text) {
9291 for payload in piece_chunks!(piece) {
9292 yield Ok(SseEvent::default().data(payload));
9293 }
9294 }
9295 continue;
9296 }
9297 let text = match scrub.as_mut() {
9298 Some(sc) => sc.push(&text),
9299 None => text,
9300 };
9301 if text.is_empty() && scrub.is_some() {
9302 continue; }
9304 let payload = if chat {
9305 chat_chunk!(json!({ "content": text }), serde_json::Value::Null)
9306 } else if openai_compat() {
9307 env.stamp(json!({ "object": "text_completion", "model": model,
9308 "choices": [{ "index": 0, "text": text, "finish_reason": null }] }))
9309 .to_string()
9310 } else {
9311 json!({ "model": model, "id": id, "text": text }).to_string()
9312 };
9313 yield Ok(SseEvent::default().data(payload));
9314 }
9315 Event::TokenSnapshot(_) => {}
9319 Event::Done { stop_reason, n_tokens, n_prompt, n_cached, elapsed_s, spec } => {
9320 let mut finish = stop_reason_to_finish(&stop_reason);
9321 if let Some(p) = parser.as_mut() {
9322 for piece in p.finish() {
9323 for payload in piece_chunks!(piece) {
9324 yield Ok(SseEvent::default().data(payload));
9325 }
9326 }
9327 if p.n_calls() > 0 { finish = "tool_calls"; }
9328 }
9329 if let Some(sc) = scrub.as_mut() {
9331 let tail = sc.finish();
9332 if !tail.is_empty() {
9333 let payload = if chat {
9334 chat_chunk!(json!({ "content": tail }),
9335 serde_json::Value::Null)
9336 } else {
9337 env.stamp(json!({ "object": "text_completion",
9338 "model": model,
9339 "choices": [{ "index": 0, "text": tail,
9340 "finish_reason": null }] })).to_string()
9341 };
9342 yield Ok(SseEvent::default().data(payload));
9343 }
9344 }
9345 if let Some(receipt) = receipt.as_mut()
9346 && let Err(err) = receipt.complete(
9347 metering::UsageCounts {
9348 prompt_tokens: n_prompt as u64,
9349 cached_prompt_tokens: n_cached as u64,
9350 completion_tokens: n_tokens as u64,
9351 },
9352 elapsed_s,
9353 )
9354 {
9355 eprintln!(
9356 "[ledger] ERROR: request {} completion receipt failed: {err}",
9357 env.id
9358 );
9359 let _ = receipt.reject(500, "request_ledger_unavailable");
9363 let payload = request_ledger_error_body().to_string();
9364 if chat || openai_compat() {
9365 yield Ok(SseEvent::default().data(payload));
9366 yield Ok(SseEvent::default().data("[DONE]"));
9367 } else {
9368 yield Ok(SseEvent::default().event("error").data(payload));
9369 }
9370 terminal = true;
9371 break;
9372 }
9373 if chat || openai_compat() {
9374 let usage = usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec);
9375 let fin = if chat {
9376 let mut v = env.stamp(json!({
9377 "object": "chat.completion.chunk", "model": model,
9378 "choices": [{ "index": 0, "delta": {},
9379 "finish_reason": finish }],
9380 "usage": usage }));
9381 if !role_sent {
9383 v["choices"][0]["delta"]["role"] = json!("assistant");
9384 }
9385 v
9386 } else {
9387 env.stamp(json!({ "object": "text_completion", "model": model,
9388 "choices": [{ "index": 0, "text": "",
9389 "finish_reason": finish }],
9390 "usage": usage }))
9391 }.to_string();
9392 yield Ok(SseEvent::default().data(fin));
9393 yield Ok(SseEvent::default().data("[DONE]"));
9394 } else {
9395 let payload = json!({
9396 "stop_reason": stop_reason, "n_tokens": n_tokens,
9397 "prompt_tokens": n_prompt, "cached_tokens": n_cached,
9398 "elapsed_s": elapsed_s
9399 }).to_string();
9400 yield Ok(SseEvent::default().event("done").data(payload));
9401 }
9402 terminal = true;
9403 break;
9404 }
9405 Event::Error(err) => {
9406 let ledger_error = if let Some(receipt) = receipt.as_mut() {
9416 receipt
9417 .reject(class_http(err.class).0.as_u16(), engine_error_code(err.class))
9418 .err()
9419 } else {
9420 None
9421 };
9422 if let Some(ref ledger_error) = ledger_error {
9423 eprintln!(
9424 "[ledger] ERROR: request {} failure receipt failed: {ledger_error}",
9425 env.id
9426 );
9427 }
9428 let payload = if ledger_error.is_some() {
9429 request_ledger_error_body().to_string()
9430 } else {
9431 engine_error_body(&err).to_string()
9432 };
9433 if chat || openai_compat() {
9434 yield Ok(SseEvent::default().data(payload));
9437 yield Ok(SseEvent::default().data("[DONE]"));
9438 } else {
9439 yield Ok(SseEvent::default().event("error").data(payload));
9442 }
9443 terminal = true;
9444 break;
9445 }
9446 }
9447 }
9448 if !terminal {
9449 let e = worker::EngineError::overloaded(
9455 "worker closed the stream without completing (worker restart in progress)",
9456 );
9457 if let Some(receipt) = receipt.as_mut()
9458 && let Err(ledger_err) = receipt.reject(
9459 class_http(e.class).0.as_u16(),
9460 engine_error_code(e.class),
9461 )
9462 {
9463 eprintln!(
9464 "[ledger] ERROR: request {} closed-stream receipt failed: {ledger_err}",
9465 env.id
9466 );
9467 }
9468 let payload = engine_error_body(&e).to_string();
9469 if chat || openai_compat() {
9470 yield Ok(SseEvent::default().data(payload));
9471 yield Ok(SseEvent::default().data("[DONE]"));
9472 } else {
9473 yield Ok(SseEvent::default().event("error").data(payload));
9474 }
9475 }
9476 };
9477 Sse::new(stream).keep_alive(
9478 axum::response::sse::KeepAlive::new().interval(std::time::Duration::from_secs(5)),
9481 )
9482}
9483
9484fn truncate_at_stop(text: &mut String, stop_strings: &[String]) {
9486 if let Some(offset) = stop_strings.iter().filter_map(|stop| text.find(stop)).min() {
9487 text.truncate(offset);
9488 }
9489}
9490
9491fn partial_stop_suffix(s: &str, tag: &str) -> usize {
9494 let mut best = 0;
9495 for (k, _) in tag.char_indices().skip(1) {
9496 if k <= s.len() && s.ends_with(&tag[..k]) {
9497 best = k;
9498 }
9499 }
9500 best
9501}
9502
9503struct StopScrubber {
9509 stops: Vec<String>,
9510 buf: String,
9511 done: bool,
9512}
9513
9514impl StopScrubber {
9515 fn new(stops: Vec<String>) -> Self {
9516 Self {
9517 stops,
9518 buf: String::new(),
9519 done: false,
9520 }
9521 }
9522
9523 fn push(&mut self, text: &str) -> String {
9525 if self.done {
9526 return String::new();
9527 }
9528 self.buf.push_str(text);
9529 if let Some(i) = self
9530 .stops
9531 .iter()
9532 .filter_map(|s| self.buf.find(s.as_str()))
9533 .min()
9534 {
9535 self.done = true;
9536 let out = self.buf[..i].to_string();
9537 self.buf.clear();
9538 return out;
9539 }
9540 let keep = self
9541 .stops
9542 .iter()
9543 .map(|s| partial_stop_suffix(&self.buf, s))
9544 .max()
9545 .unwrap_or(0);
9546 let emit_to = self.buf.len() - keep;
9547 let out = self.buf[..emit_to].to_string();
9548 self.buf.drain(..emit_to);
9549 out
9550 }
9551
9552 fn finish(&mut self) -> String {
9554 if self.done {
9555 self.buf.clear();
9556 return String::new();
9557 }
9558 std::mem::take(&mut self.buf)
9559 }
9560}
9561
9562#[cfg(test)]
9563async fn blocking_response(
9564 rx: worker::EventReceiver,
9565 model: String,
9566 chat: bool,
9567 stop_strings: Vec<String>,
9568 parser: Option<ToolStreamParser>,
9569 env: Envelope,
9570) -> Response {
9571 blocking_response_with_receipt(rx, model, chat, stop_strings, parser, env, &mut None, None)
9572 .await
9573}
9574
9575struct BlockingPayload<'a> {
9579 env: &'a Envelope,
9580 model: String,
9581 chat: bool,
9582 finish: &'static str,
9583 text: String,
9584 reasoning: String,
9585 calls: Vec<ParsedToolCall>,
9586 tokens: Vec<u32>,
9587 stop_reason: String,
9588 n_prompt: usize,
9589 n_tokens: usize,
9590 n_cached: usize,
9591 elapsed_s: f64,
9592 spec: Option<worker::SpecUsage>,
9593 deadline_error: Option<serde_json::Value>,
9599}
9600
9601fn blocking_payload(p: BlockingPayload<'_>) -> Response {
9602 let BlockingPayload {
9603 env,
9604 model,
9605 chat,
9606 finish,
9607 text,
9608 reasoning,
9609 calls,
9610 tokens,
9611 stop_reason,
9612 n_prompt,
9613 n_tokens,
9614 n_cached,
9615 elapsed_s,
9616 spec,
9617 deadline_error,
9618 } = p;
9619 if chat {
9620 let content = if !calls.is_empty() && text.is_empty() {
9622 serde_json::Value::Null
9623 } else {
9624 serde_json::Value::String(text)
9625 };
9626 let mut message = json!({ "role": "assistant", "content": content });
9627 if !reasoning.is_empty() {
9630 message["reasoning"] = json!(reasoning);
9631 message["reasoning_details"] = json!([{
9632 "type": "reasoning.text", "text": reasoning }]);
9633 }
9634 if !calls.is_empty() {
9635 message["tool_calls"] =
9636 serde_json::Value::Array(calls.iter().map(tool_call_json).collect());
9637 }
9638 let mut body = json!({
9639 "object": "chat.completion", "model": model,
9640 "choices": [{ "index": 0,
9641 "message": message,
9642 "finish_reason": finish }],
9643 "usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
9644 });
9645 if let Some(err) = deadline_error {
9646 body["choices"][0]["native_finish_reason"] = json!("deadline_exceeded");
9647 body["error"] = err;
9648 }
9649 return Json(env.stamp(body)).into_response();
9650 }
9651 if openai_compat() {
9652 let mut body = json!({
9653 "object": "text_completion", "model": model,
9654 "choices": [{ "index": 0, "text": text,
9655 "finish_reason": finish }],
9656 "usage": usage_json(n_prompt, n_tokens, n_cached, elapsed_s, spec)
9657 });
9658 if let Some(err) = deadline_error {
9659 body["choices"][0]["native_finish_reason"] = json!("deadline_exceeded");
9660 body["error"] = err;
9661 }
9662 return Json(env.stamp(body)).into_response();
9663 }
9664 Json(CompletionResp {
9665 model,
9666 text,
9667 tokens,
9668 stop_reason,
9669 error: deadline_error,
9670 n_tokens,
9671 prompt_tokens: n_prompt,
9672 cached_tokens: n_cached,
9673 elapsed_s,
9674 })
9675 .into_response()
9676}
9677
9678#[allow(clippy::too_many_arguments)] async fn blocking_response_with_receipt(
9704 mut rx: worker::EventReceiver,
9705 model: String,
9706 chat: bool,
9707 stop_strings: Vec<String>,
9708 mut parser: Option<ToolStreamParser>,
9709 env: Envelope,
9710 receipt: &mut Option<Box<dyn metering::Receipt>>,
9711 deadline: Option<RequestDeadline>,
9712) -> Response {
9713 let mut text = String::new();
9714 let mut reasoning = String::new();
9715 let mut tokens: Vec<u32> = Vec::new();
9716 let mut calls: Vec<ParsedToolCall> = Vec::new();
9717 let consume = |pieces: Vec<Piece>,
9718 text: &mut String,
9719 reasoning: &mut String,
9720 calls: &mut Vec<ParsedToolCall>| {
9721 for piece in pieces {
9722 match piece {
9723 Piece::Content(t) => text.push_str(&t),
9724 Piece::Reasoning(t) => reasoning.push_str(&t),
9725 Piece::Call(c) => calls.push(c),
9726 }
9727 }
9728 };
9729 let started = std::time::Instant::now();
9731 let mut seen_prompt: usize = 0;
9732 let mut seen_cached: usize = 0;
9733 let mut seen_tokens: usize = 0;
9734 loop {
9735 let ev = match deadline {
9736 Some(d) => tokio::select! {
9737 biased;
9738 ev = rx.recv() => ev,
9739 () = tokio::time::sleep_until(d.at) => {
9740 drop(rx);
9743 if seen_tokens == 0 {
9744 if let Some(receipt) = receipt.as_mut()
9748 && let Err(err) = receipt.settle_unbilled(
9749 "deadline_exceeded",
9750 StatusCode::REQUEST_TIMEOUT.as_u16(),
9751 "deadline_exceeded",
9752 )
9753 {
9754 eprintln!(
9755 "[ledger] ERROR: request {} deadline receipt failed: {err}",
9756 env.id
9757 );
9758 return request_ledger_error_response();
9759 }
9760 return deadline_exceeded_response(d.ms, false);
9761 }
9762 if let Some(p) = parser.as_mut() {
9763 consume(p.finish(), &mut text, &mut reasoning, &mut calls);
9764 }
9765 truncate_at_stop(&mut text, &stop_strings);
9766 let elapsed_s = started.elapsed().as_secs_f64();
9767 if let Some(receipt) = receipt.as_mut()
9770 && let Err(err) = receipt.complete_deadline_partial(
9771 metering::UsageCounts {
9772 prompt_tokens: seen_prompt as u64,
9773 cached_prompt_tokens: seen_cached as u64,
9774 completion_tokens: seen_tokens as u64,
9775 },
9776 elapsed_s,
9777 )
9778 {
9779 eprintln!(
9780 "[ledger] ERROR: request {} partial-deadline receipt failed: {err}",
9781 env.id
9782 );
9783 let _ = receipt.reject(500, "request_ledger_unavailable");
9784 return request_ledger_error_response();
9785 }
9786 eprintln!(
9787 "[deadline] request {} delivered PARTIAL: {} tokens in {:.1}s of a \
9788 {} ms deadline (prompt {}); non-streaming caller advised to stream",
9789 env.id, seen_tokens, elapsed_s, d.ms, seen_prompt
9790 );
9791 let err_obj = json!({
9792 "message": format!(
9793 "deadline of {} ms (timeout_ms; default {}) elapsed mid-generation; \
9794 the {} tokens produced before the cut are delivered above and are \
9795 billed. Set \"stream\": true for work this long — a stream's \
9796 deadline bounds only the time to first token — or lower max_tokens.",
9797 d.ms, TIMEOUT_MS_DEFAULT, seen_tokens
9798 ),
9799 "code": "deadline_exceeded",
9800 "metadata": { "error_type": "timeout", "provider_name": "memra" }
9801 });
9802 return blocking_payload(BlockingPayload {
9803 env: &env,
9804 model,
9805 chat,
9806 finish: "error",
9807 text,
9808 reasoning,
9809 calls,
9810 tokens,
9811 stop_reason: "Deadline".to_string(),
9812 n_prompt: seen_prompt,
9813 n_tokens: seen_tokens,
9814 n_cached: seen_cached,
9815 elapsed_s,
9816 spec: None,
9817 deadline_error: Some(err_obj),
9818 });
9819 }
9820 },
9821 None => rx.recv().await,
9822 };
9823 let Some(ev) = ev else { break };
9824 match ev {
9825 Event::PromptCapture { .. } => {} Event::PromptUsage { n_prompt, n_cached } => {
9827 if let Some(receipt) = receipt.as_mut()
9828 && let Err(err) = receipt.record_prompt_usage(n_prompt as u64, n_cached as u64)
9829 {
9830 eprintln!(
9831 "[ledger] ERROR: request {} partial prompt receipt failed: {err}",
9832 env.id
9833 );
9834 let _ = receipt.reject(500, "request_ledger_unavailable");
9837 return request_ledger_error_response();
9838 }
9839 seen_prompt = n_prompt;
9840 seen_cached = n_cached;
9841 }
9842 Event::Token { id, text: delta } => {
9843 if let Some(receipt) = receipt.as_mut()
9844 && let Err(err) = receipt.record_completion_token()
9845 {
9846 eprintln!(
9847 "[ledger] ERROR: request {} partial completion receipt failed: {err}",
9848 env.id
9849 );
9850 let _ = receipt.reject(500, "request_ledger_unavailable");
9851 return request_ledger_error_response();
9852 }
9853 if let Some(receipt) = receipt.as_mut() {
9855 receipt.capture_completion_delta(&delta);
9856 }
9857 tokens.push(id);
9858 seen_tokens += 1;
9859 match parser.as_mut() {
9860 Some(p) => consume(p.push(&delta), &mut text, &mut reasoning, &mut calls),
9861 None => text.push_str(&delta),
9862 }
9863 }
9864 Event::TokenSnapshot(ids) => tokens = ids,
9865 Event::Done {
9866 stop_reason,
9867 n_tokens,
9868 n_prompt,
9869 n_cached,
9870 elapsed_s,
9871 spec,
9872 } => {
9873 if let Some(p) = parser.as_mut() {
9874 consume(p.finish(), &mut text, &mut reasoning, &mut calls);
9875 }
9876 truncate_at_stop(&mut text, &stop_strings);
9877 let finish = if calls.is_empty() {
9878 stop_reason_to_finish(&stop_reason)
9879 } else {
9880 "tool_calls"
9881 };
9882 if let Some(receipt) = receipt.as_mut()
9883 && let Err(err) = receipt.complete(
9884 metering::UsageCounts {
9885 prompt_tokens: n_prompt as u64,
9886 cached_prompt_tokens: n_cached as u64,
9887 completion_tokens: n_tokens as u64,
9888 },
9889 elapsed_s,
9890 )
9891 {
9892 eprintln!(
9893 "[ledger] ERROR: request {} completion receipt failed: {err}",
9894 env.id
9895 );
9896 let _ = receipt.reject(500, "request_ledger_unavailable");
9899 return request_ledger_error_response();
9900 }
9901 return blocking_payload(BlockingPayload {
9902 env: &env,
9903 model,
9904 chat,
9905 finish,
9906 text,
9907 reasoning,
9908 calls,
9909 tokens,
9910 stop_reason,
9911 n_prompt,
9912 n_tokens,
9913 n_cached,
9914 elapsed_s,
9915 spec,
9916 deadline_error: None,
9917 });
9918 }
9919 Event::Error(err) => {
9920 if let Some(receipt) = receipt.as_mut()
9924 && let Err(ledger_err) = receipt.reject(
9925 class_http(err.class).0.as_u16(),
9926 engine_error_code(err.class),
9927 )
9928 {
9929 eprintln!(
9930 "[ledger] ERROR: request {} failure receipt failed: {ledger_err}",
9931 env.id
9932 );
9933 return request_ledger_error_response();
9934 }
9935 return engine_error_response(&err);
9936 }
9937 }
9938 }
9939 let e = worker::EngineError::overloaded(
9944 "worker closed the stream without completing (worker restart in progress)",
9945 );
9946 if let Some(receipt) = receipt.as_mut()
9947 && let Err(ledger_err) =
9948 receipt.reject(class_http(e.class).0.as_u16(), engine_error_code(e.class))
9949 {
9950 eprintln!(
9951 "[ledger] ERROR: request {} closed-stream receipt failed: {ledger_err}",
9952 env.id
9953 );
9954 return request_ledger_error_response();
9955 }
9956 engine_error_response(&e)
9957}
9958
9959#[cfg(test)]
9960mod tests {
9961 use super::*;
9962
9963 #[test]
9969 fn capture_children_are_distinct_ledger_identities_under_the_parent_id() {
9970 let parent = Envelope::new(false);
9971 assert!(parent.id.starts_with("cmpl-"));
9972 let a = parent.capture_child(0);
9973 let b = parent.capture_child(1);
9974 let c = parent.capture_child(2);
9975 assert_eq!(a.id, format!("{}.0", parent.id));
9976 assert_eq!(b.id, format!("{}.1", parent.id));
9977 assert_eq!(c.id, format!("{}.2", parent.id));
9978 assert_ne!(a.id, b.id);
9979 assert_ne!(b.id, c.id);
9980 for child in [&a, &b, &c] {
9981 assert!(
9982 child.id.starts_with(&parent.id),
9983 "child nests under the parent by prefix"
9984 );
9985 assert_ne!(
9986 child.id, parent.id,
9987 "a child never reuses the parent's ledger id"
9988 );
9989 assert_eq!(child.created, parent.created);
9990 }
9991 assert_eq!(parent.capture_child(1).id, b.id);
9994 }
9995
9996 #[derive(Debug, Clone, PartialEq)]
10004 enum MeterEvent {
10005 Reserve {
10006 tenant: String,
10007 principal: Option<String>,
10008 model: String,
10009 },
10010 Open {
10011 request_id: String,
10012 tenant: String,
10013 model: String,
10014 route: &'static str,
10015 stream: bool,
10016 with_permit: bool,
10017 },
10018 PromptUsage {
10019 prompt: u64,
10020 cached: u64,
10021 },
10022 Token,
10023 CapturePrompt(serde_json::Value),
10024 CaptureDelta(String),
10025 Complete {
10026 prompt: u64,
10027 cached: u64,
10028 completion: u64,
10029 },
10030 DeadlinePartial {
10031 prompt: u64,
10032 cached: u64,
10033 completion: u64,
10034 },
10035 Reject {
10036 status: u16,
10037 code: String,
10038 },
10039 Unbilled {
10040 outcome: &'static str,
10041 status: u16,
10042 code: String,
10043 },
10044 Dropped {
10047 prompt: u64,
10048 cached: u64,
10049 completion: u64,
10050 },
10051 }
10052
10053 enum ReserveScript {
10056 Admit { with_permit: bool },
10057 Insufficient,
10058 Blocked,
10059 PrincipalCapped,
10060 }
10061
10062 struct MockMetering {
10063 events: Arc<std::sync::Mutex<Vec<MeterEvent>>>,
10064 limits: bool,
10065 limited: bool,
10066 reserve_script: std::sync::Mutex<std::collections::VecDeque<ReserveScript>>,
10067 captures: bool,
10068 }
10069
10070 impl MockMetering {
10071 fn admit_all() -> Arc<Self> {
10072 Arc::new(MockMetering {
10073 events: Arc::new(std::sync::Mutex::new(Vec::new())),
10074 limits: false,
10075 limited: true,
10076 reserve_script: std::sync::Mutex::new(std::collections::VecDeque::new()),
10077 captures: false,
10078 })
10079 }
10080
10081 fn with_limits(script: Vec<ReserveScript>) -> Arc<Self> {
10082 Arc::new(MockMetering {
10083 events: Arc::new(std::sync::Mutex::new(Vec::new())),
10084 limits: true,
10085 limited: true,
10086 reserve_script: std::sync::Mutex::new(script.into()),
10087 captures: false,
10088 })
10089 }
10090
10091 fn capturing() -> Arc<Self> {
10092 Arc::new(MockMetering {
10093 events: Arc::new(std::sync::Mutex::new(Vec::new())),
10094 limits: false,
10095 limited: true,
10096 reserve_script: std::sync::Mutex::new(std::collections::VecDeque::new()),
10097 captures: true,
10098 })
10099 }
10100
10101 fn events(&self) -> Vec<MeterEvent> {
10102 self.events.lock().unwrap().clone()
10103 }
10104 }
10105
10106 impl metering::Metering for MockMetering {
10107 fn enforces_limits(&self) -> bool {
10108 self.limits
10109 }
10110
10111 fn is_limited(&self, _tenant: &str) -> Result<bool, metering::AdmitError> {
10112 Ok(self.limited)
10113 }
10114
10115 fn reserve(
10116 &self,
10117 tenant: &str,
10118 principal: Option<&str>,
10119 model: &str,
10120 _prompt_tokens: u64,
10121 _completion_bound: u64,
10122 ) -> Result<Option<metering::Permit>, metering::AdmitError> {
10123 self.events.lock().unwrap().push(MeterEvent::Reserve {
10124 tenant: tenant.into(),
10125 principal: principal.map(str::to_owned),
10126 model: model.into(),
10127 });
10128 match self.reserve_script.lock().unwrap().pop_front() {
10129 None | Some(ReserveScript::Admit { with_permit: false }) => Ok(None),
10130 Some(ReserveScript::Admit { with_permit: true }) => {
10131 Ok(Some(Box::new(()) as metering::Permit))
10132 }
10133 Some(ReserveScript::Insufficient) => Err(metering::AdmitError::Insufficient),
10134 Some(ReserveScript::Blocked) => Err(metering::AdmitError::Blocked),
10135 Some(ReserveScript::PrincipalCapped) => Err(metering::AdmitError::PrincipalCapped),
10136 }
10137 }
10138
10139 fn open(
10140 &self,
10141 meta: &metering::RequestMeta<'_>,
10142 permit: Option<metering::Permit>,
10143 ) -> Box<dyn metering::Receipt> {
10144 self.events.lock().unwrap().push(MeterEvent::Open {
10145 request_id: meta.request_id.into(),
10146 tenant: meta.tenant.into(),
10147 model: meta.model.into(),
10148 route: meta.route,
10149 stream: meta.stream,
10150 with_permit: permit.is_some(),
10151 });
10152 Box::new(MockReceipt {
10153 events: self.events.clone(),
10154 wants_capture: self.captures,
10155 prompt: 0,
10156 cached: 0,
10157 completion: 0,
10158 finalized: false,
10159 })
10160 }
10161
10162 fn captures(&self, _tenant: &str) -> bool {
10163 self.captures
10164 }
10165
10166 fn limits_health(&self) -> Option<metering::LimitsHealth> {
10167 self.limits.then_some(metering::LimitsHealth {
10168 source_reload_failed: 0,
10169 source_reload_consecutive: 0,
10170 source_available: true,
10171 })
10172 }
10173 }
10174
10175 struct MockReceipt {
10176 events: Arc<std::sync::Mutex<Vec<MeterEvent>>>,
10177 wants_capture: bool,
10178 prompt: u64,
10179 cached: u64,
10180 completion: u64,
10181 finalized: bool,
10182 }
10183
10184 impl metering::Receipt for MockReceipt {
10185 fn wants_capture(&self) -> bool {
10186 self.wants_capture
10187 }
10188
10189 fn arm_capture(&mut self, prompt: serde_json::Value) {
10190 self.events
10191 .lock()
10192 .unwrap()
10193 .push(MeterEvent::CapturePrompt(prompt));
10194 }
10195
10196 fn capture_completion_delta(&mut self, text: &str) {
10197 if self.wants_capture {
10198 self.events
10199 .lock()
10200 .unwrap()
10201 .push(MeterEvent::CaptureDelta(text.into()));
10202 }
10203 }
10204
10205 fn record_prompt_usage(&mut self, prompt: u64, cached: u64) -> Result<(), String> {
10206 self.prompt = prompt;
10207 self.cached = cached;
10208 self.events
10209 .lock()
10210 .unwrap()
10211 .push(MeterEvent::PromptUsage { prompt, cached });
10212 Ok(())
10213 }
10214
10215 fn record_completion_token(&mut self) -> Result<(), String> {
10216 self.completion += 1;
10217 self.events.lock().unwrap().push(MeterEvent::Token);
10218 Ok(())
10219 }
10220
10221 fn complete(
10222 &mut self,
10223 usage: metering::UsageCounts,
10224 _worker_elapsed_s: f64,
10225 ) -> Result<(), String> {
10226 self.finalized = true;
10227 self.events.lock().unwrap().push(MeterEvent::Complete {
10228 prompt: usage.prompt_tokens,
10229 cached: usage.cached_prompt_tokens,
10230 completion: usage.completion_tokens,
10231 });
10232 Ok(())
10233 }
10234
10235 fn complete_deadline_partial(
10236 &mut self,
10237 usage: metering::UsageCounts,
10238 _worker_elapsed_s: f64,
10239 ) -> Result<(), String> {
10240 self.finalized = true;
10241 self.events
10242 .lock()
10243 .unwrap()
10244 .push(MeterEvent::DeadlinePartial {
10245 prompt: usage.prompt_tokens,
10246 cached: usage.cached_prompt_tokens,
10247 completion: usage.completion_tokens,
10248 });
10249 Ok(())
10250 }
10251
10252 fn reject(&mut self, status: u16, error_code: &str) -> Result<(), String> {
10253 self.finalized = true;
10254 self.events.lock().unwrap().push(MeterEvent::Reject {
10255 status,
10256 code: error_code.into(),
10257 });
10258 Ok(())
10259 }
10260
10261 fn settle_unbilled(
10262 &mut self,
10263 outcome: &'static str,
10264 status: u16,
10265 error_code: &str,
10266 ) -> Result<(), String> {
10267 self.finalized = true;
10268 self.events.lock().unwrap().push(MeterEvent::Unbilled {
10269 outcome,
10270 status,
10271 code: error_code.into(),
10272 });
10273 Ok(())
10274 }
10275 }
10276
10277 impl Drop for MockReceipt {
10278 fn drop(&mut self) {
10279 if !self.finalized {
10280 self.events.lock().unwrap().push(MeterEvent::Dropped {
10281 prompt: self.prompt,
10282 cached: self.cached,
10283 completion: self.completion,
10284 });
10285 }
10286 }
10287 }
10288
10289 static GATE_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
10294
10295 fn gate_env_lock() -> std::sync::MutexGuard<'static, ()> {
10302 let guard = GATE_ENV_LOCK.lock().unwrap_or_else(|poisoned| {
10303 GATE_ENV_LOCK.clear_poison();
10306 poisoned.into_inner()
10307 });
10308 unsafe { std::env::remove_var("MEMRA_NONSTREAM_DEADLINE_GATE") };
10309 guard
10310 }
10311
10312 fn gate_request(max_new: usize, prompt_ids: usize) -> worker::Request {
10315 let req: CompletionReq = serde_json::from_value(json!({
10316 "model": "qwen/qwen3.8-27b",
10317 "prompt_ids": vec![7u32; prompt_ids],
10318 }))
10319 .unwrap();
10320 let (tx, _rx) = worker::event_channel();
10321 let mut request = build_request(&req, tx, lanes::Lane::Interactive, None);
10322 request.params.max_new = max_new;
10323 request
10324 }
10325
10326 #[test]
10333 fn the_feasibility_gate_boundary_matches_the_measured_ladder() {
10334 let prompt = 30_278u64;
10335 let deadline_ms = TIMEOUT_MS_DEFAULT;
10336 let margin = |max_new: u64| {
10337 let prefill_ms = prompt * 1_000 / PREFILL_FLOOR_TOK_S;
10338 let decode_ms = max_new * 1_000 / DECODE_FLOOR_TOK_S;
10339 (prefill_ms + decode_ms) <= deadline_ms * DEADLINE_INFEASIBLE_MARGIN_PCT / 100
10340 };
10341 for allowed in [64u64, 2048, 4096, 5120, 6144] {
10342 assert!(margin(allowed), "{allowed} measured OK and must be allowed");
10343 }
10344 for refused in [8192u64, 16384, 262_144] {
10345 assert!(
10346 !margin(refused),
10347 "{refused} measured as a 408 and must be refused"
10348 );
10349 }
10350 }
10351
10352 #[test]
10353 fn the_gate_names_a_max_tokens_that_actually_fits() {
10354 let fits = deadline_fitting_max_tokens(30_278, TIMEOUT_MS_DEFAULT).unwrap();
10357 assert!(
10358 fits > 0 && fits < 7_800,
10359 "advice {fits} must fit the measured ceiling"
10360 );
10361 assert_eq!(
10363 deadline_fitting_max_tokens(400_000, TIMEOUT_MS_DEFAULT),
10364 None
10365 );
10366 }
10367
10368 #[test]
10369 fn streaming_is_never_gated_and_the_gate_can_be_switched_off() {
10370 let req = gate_request(262_144, 30_000);
10371 let deadline = RequestDeadline::starting_now(TIMEOUT_MS_DEFAULT);
10372 let err = nonstream_deadline_gate(&req, false, deadline, true, None).unwrap_err();
10374 assert!(
10375 err.contains("stream"),
10376 "message must name the streaming alternative: {err}"
10377 );
10378 assert!(
10379 err.contains("max_tokens"),
10380 "message must name the knob: {err}"
10381 );
10382 assert!(nonstream_deadline_gate(&req, true, deadline, true, None).is_ok());
10384 let _l = gate_env_lock(); for off in ["0", "off", "false"] {
10391 unsafe { std::env::set_var("MEMRA_NONSTREAM_DEADLINE_GATE", off) };
10392 assert!(
10393 nonstream_deadline_gate(&req, false, deadline, true, None).is_ok(),
10394 "MEMRA_NONSTREAM_DEADLINE_GATE={off} must disable the gate"
10395 );
10396 }
10397 unsafe { std::env::set_var("MEMRA_NONSTREAM_DEADLINE_GATE", "1") };
10398 assert!(nonstream_deadline_gate(&req, false, deadline, true, None).is_err());
10399 unsafe { std::env::remove_var("MEMRA_NONSTREAM_DEADLINE_GATE") };
10400 assert!(
10401 nonstream_deadline_gate(&req, false, deadline, true, None).is_err(),
10402 "unset means ON (the documented default)"
10403 );
10404 }
10405
10406 #[test]
10413 fn the_feasibility_gate_is_wired_on_every_surface_not_just_the_two_i_remembered() {
10414 let strip = |src: &str| -> String {
10421 src.lines()
10422 .map(|line| match line.find("//") {
10423 Some(i) => line[..i].to_string(),
10424 None => line.to_string(),
10425 })
10426 .collect::<Vec<_>>()
10427 .join("\n")
10428 };
10429 fn body<'a>(src: &'a str, signature: &str) -> &'a str {
10431 let start = src
10432 .find(signature)
10433 .unwrap_or_else(|| panic!("{signature} not found — did the handler get renamed?"));
10434 let rest = &src[start + signature.len()..];
10435 let end = rest.find("\nasync fn ").unwrap_or(rest.len());
10436 let end = rest[..end].find("\npub(crate) async fn ").unwrap_or(end);
10437 &rest[..end]
10438 }
10439 let main_src = strip(include_str!("lib.rs"));
10440 let surfaces_src = strip(include_str!("surfaces.rs"));
10441 for (surface, src, signature) in [
10442 (
10443 "/v1/completions",
10444 &main_src,
10445 "async fn completions_with_admission(",
10446 ),
10447 (
10448 "/v1/chat/completions",
10449 &main_src,
10450 "async fn chat_completions_with_admission(",
10451 ),
10452 (
10453 "/v1/messages + /v1/responses (shared admission)",
10454 &surfaces_src,
10455 "pub(crate) async fn admit_translated(",
10456 ),
10457 ] {
10458 let handler = body(src, signature);
10459 assert!(
10460 handler.contains("nonstream_deadline_gate("),
10461 "{surface} must CALL the feasibility gate inside {signature}"
10462 );
10463 let limits = handler
10466 .find("apply_model_request_limits(")
10467 .unwrap_or_else(|| panic!("{surface}: no apply_model_request_limits call"));
10468 let gate = handler.find("nonstream_deadline_gate(").unwrap();
10469 assert!(
10470 limits < gate,
10471 "{surface}: the gate must run after apply_model_request_limits"
10472 );
10473 }
10474 }
10475
10476 #[test]
10480 fn the_native_shape_carries_the_deadline_error_and_omits_it_otherwise() {
10481 let err = json!({"code": "deadline_exceeded",
10482 "metadata": {"error_type": "timeout"}});
10483 let cut = CompletionResp {
10484 model: "m".into(),
10485 text: "partial".into(),
10486 tokens: vec![1, 2],
10487 stop_reason: "Deadline".into(),
10488 error: Some(err.clone()),
10489 n_tokens: 2,
10490 prompt_tokens: 9,
10491 cached_tokens: 0,
10492 elapsed_s: 1.0,
10493 };
10494 let v = serde_json::to_value(&cut).unwrap();
10495 assert_eq!(v["stop_reason"], "Deadline");
10496 assert_eq!(v["error"]["code"], "deadline_exceeded");
10497 assert_eq!(v["error"]["metadata"]["error_type"], "timeout");
10498 let whole = CompletionResp {
10500 error: None,
10501 stop_reason: "Eos".into(),
10502 ..cut
10503 };
10504 let v = serde_json::to_value(&whole).unwrap();
10505 assert!(
10506 v.get("error").is_none(),
10507 "a complete response must not grow an error key: {v}"
10508 );
10509 }
10510
10511 #[test]
10512 fn a_ctx_bounded_request_is_not_gated_because_context_is_its_only_limit() {
10513 let _l = gate_env_lock();
10514 let req = gate_request(worker::MAX_NEW_CTX_BOUNDED, 30_000);
10518 assert!(
10519 nonstream_deadline_gate(
10520 &req,
10521 false,
10522 RequestDeadline::starting_now(TIMEOUT_MS_DEFAULT),
10523 false,
10524 None,
10525 )
10526 .is_ok(),
10527 "an omitted max_tokens is never gated — context is its only limit"
10528 );
10529 let resolved = gate_request(32_768, 30_000);
10534 assert!(
10535 nonstream_deadline_gate(
10536 &resolved,
10537 false,
10538 RequestDeadline::starting_now(TIMEOUT_MS_DEFAULT),
10539 false,
10540 None,
10541 )
10542 .is_ok(),
10543 "a resolved-but-undeclared cap is not the caller's number to be refused over"
10544 );
10545 assert!(
10547 nonstream_deadline_gate(
10548 &resolved,
10549 false,
10550 RequestDeadline::starting_now(TIMEOUT_MS_DEFAULT),
10551 true,
10552 None,
10553 )
10554 .is_err()
10555 );
10556 }
10557
10558 #[test]
10559 fn the_prompt_estimate_is_exact_for_ids_and_a_proxy_otherwise() {
10560 let req = gate_request(64, 1234);
10561 assert_eq!(prompt_tokens_estimate(&req, None), 1234, "ids are exact");
10562 let mut text = gate_request(64, 0);
10563 text.prompt_ids.clear();
10564 text.prompt_text = "x".repeat(6_000);
10565 assert_eq!(
10566 prompt_tokens_estimate(&text, None),
10567 1_000,
10568 "the fallback under-counts on purpose (bytes/6): an over-count refuses work \
10569 that would have succeeded"
10570 );
10571 }
10572
10573 #[test]
10574 fn vision_memory_reservation_is_bounded_and_released() {
10575 let permit = try_reserve_vision_memory(MAX_VISION_PATCH_BYTES).unwrap();
10576 let Err(capacity) = try_reserve_vision_memory(1) else {
10577 panic!("a full process vision budget admitted another request");
10578 };
10579 assert!(matches!(capacity, VisionMemoryError::Capacity(_)));
10580 let response = vision_memory_error_response(capacity, Some("messages"));
10581 assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
10582 assert_eq!(response.headers()["retry-after"], "5");
10583 assert_eq!(response.headers()["retry-after-ms"], "5000");
10584 drop(permit);
10585 assert!(try_reserve_vision_memory(1).is_ok());
10586 let Err(request) = try_reserve_vision_memory(MAX_VISION_PATCH_BYTES + 1) else {
10587 panic!("an over-limit vision request was admitted");
10588 };
10589 assert!(matches!(request, VisionMemoryError::Request(_)));
10590 let response = vision_memory_error_response(request, Some("messages"));
10591 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
10592 assert_eq!(response.headers()["x-should-retry"], "false");
10593 let _ = try_reserve_vision_memory(1);
10594 }
10595
10596 #[test]
10597 fn header_auth_gate_covers_only_inference_dialects() {
10598 for path in [
10599 "/v1/auth/check",
10600 "/v1/completions",
10601 "/v1/chat/completions",
10602 "/v1/messages",
10603 "/v1/responses",
10604 "/v1/embeddings",
10605 "/v1/rerank",
10606 ] {
10607 assert!(protected_inference_path(path), "{path}");
10608 }
10609 for path in ["/health", "/readyz", "/models", "/v1/models", "/metrics"] {
10610 assert!(!protected_inference_path(path), "{path}");
10611 }
10612 }
10613 #[tokio::test]
10620 async fn served_completion_capture_is_byte_exact_and_armed_receipts_only() {
10621 use crate::metering::Metering as _;
10622 let prompt = json!([{ "role": "user", "content": "capture me — exactly" }]);
10623
10624 let drive = |receipt: Option<Box<dyn metering::Receipt>>| async {
10625 let (tx, rx) = worker::event_channel();
10626 tx.send(Event::PromptUsage {
10627 n_prompt: 7,
10628 n_cached: 0,
10629 })
10630 .unwrap();
10631 tx.send(Event::Token {
10632 id: 1,
10633 text: "Hel".into(),
10634 })
10635 .unwrap();
10636 tx.send(Event::Token {
10637 id: 2,
10638 text: "lo".into(),
10639 })
10640 .unwrap();
10641 tx.send(Event::Done {
10642 stop_reason: "eos".into(),
10643 n_tokens: 2,
10644 n_prompt: 7,
10645 n_cached: 0,
10646 elapsed_s: 0.05,
10647 spec: None,
10648 })
10649 .unwrap();
10650 drop(tx);
10651 let mut receipt = receipt;
10652 blocking_response_with_receipt(
10653 rx,
10654 "m".into(),
10655 true,
10656 Vec::new(),
10657 None,
10658 Envelope::new(true),
10659 &mut receipt,
10660 None,
10661 )
10662 .await
10663 };
10664
10665 let plain = MockMetering::admit_all();
10667 let receipt = plain.open(
10668 &metering::RequestMeta {
10669 request_id: "cap-unmarked",
10670 tenant: "unmarked",
10671 principal: None,
10672 model: "m",
10673 route: "/v1/chat/completions",
10674 lane: "interactive",
10675 stream: false,
10676 max_tokens: None,
10677 reserved_ctx: None,
10678 },
10679 None,
10680 );
10681 let response = drive(Some(receipt)).await;
10682 assert_eq!(response.status(), StatusCode::OK);
10683 assert!(
10684 !plain.events().iter().any(|e| matches!(
10685 e,
10686 MeterEvent::CaptureDelta(_) | MeterEvent::CapturePrompt(_)
10687 )),
10688 "an unarmed receipt must see no capture traffic: {:?}",
10689 plain.events()
10690 );
10691
10692 let capturing = MockMetering::capturing();
10695 let mut receipt = capturing.open(
10696 &metering::RequestMeta {
10697 request_id: "cap-marked",
10698 tenant: "marked",
10699 principal: None,
10700 model: "m",
10701 route: "/v1/chat/completions",
10702 lane: "interactive",
10703 stream: false,
10704 max_tokens: None,
10705 reserved_ctx: None,
10706 },
10707 None,
10708 );
10709 assert!(receipt.wants_capture());
10710 receipt.arm_capture(prompt.clone());
10711 let response = drive(Some(receipt)).await;
10712 assert_eq!(response.status(), StatusCode::OK);
10713 let body = axum::body::to_bytes(response.into_body(), usize::MAX)
10714 .await
10715 .unwrap();
10716 let body: serde_json::Value = serde_json::from_slice(&body).unwrap();
10717 assert_eq!(body["choices"][0]["message"]["content"], "Hello");
10718
10719 let events = capturing.events();
10720 assert!(
10721 events.contains(&MeterEvent::CapturePrompt(prompt.clone())),
10722 "prompt must arm byte-exact: {events:?}"
10723 );
10724 let completion: String = events
10725 .iter()
10726 .filter_map(|e| match e {
10727 MeterEvent::CaptureDelta(text) => Some(text.as_str()),
10728 _ => None,
10729 })
10730 .collect();
10731 assert_eq!(
10732 completion, "Hello",
10733 "the deltas must reassemble the served completion byte-exact: {events:?}"
10734 );
10735 assert!(
10736 events.contains(&MeterEvent::Complete {
10737 prompt: 7,
10738 cached: 0,
10739 completion: 2,
10740 }),
10741 "worker-truth usage settles alongside the capture: {events:?}"
10742 );
10743 }
10744
10745 fn tool_caps() -> ModelCaps {
10746 ModelCaps {
10747 tools_branch: true,
10748 qwen_think: true,
10749 think_switch: true,
10750 chat_ok: true,
10751 ..Default::default()
10752 }
10753 }
10754
10755 fn ladder_caps() -> ModelCaps {
10760 ModelCaps {
10761 qwen_effort: true,
10762 ..tool_caps()
10763 }
10764 }
10765
10766 fn gemma_tool_caps() -> ModelCaps {
10767 ModelCaps {
10768 tools_branch: true,
10769 gemma_think: true,
10770 chat_ok: true,
10771 instruct_type: Some("gemma".into()),
10772 ..Default::default()
10773 }
10774 }
10775
10776 fn hy3_tool_caps() -> ModelCaps {
10777 ModelCaps {
10778 tools_branch: true,
10779 hy3: true,
10780 chat_ok: true,
10781 effort_levels: true,
10782 instruct_type: Some("hy3".into()),
10783 ..Default::default()
10784 }
10785 }
10786
10787 fn gemma_template(kind: &str) -> String {
10788 let file = match kind {
10789 "qat" => "qat-trunk-template.jinja",
10790 _ => "official-tooluse-template.jinja",
10791 };
10792 let path = format!(
10793 "{}/../../research/gemma4-tools-20260817/{file}",
10794 env!("CARGO_MANIFEST_DIR")
10795 );
10796 std::fs::read_to_string(&path).unwrap_or_else(|e| panic!("read {path}: {e}"))
10797 }
10798
10799 fn render_fixture(request: &serde_json::Value, template: &str) -> String {
10804 let tools_arr = request
10805 .get("tools")
10806 .and_then(|t| t.as_array())
10807 .cloned()
10808 .unwrap_or_default();
10809 let (tools_json, tools_struct, _schemas) = if tools_arr.is_empty() {
10810 (Vec::new(), Vec::new(), HashMap::new())
10811 } else {
10812 prepare_tools(&tools_arr).unwrap()
10813 };
10814 let effort = request
10815 .get("reasoning_effort")
10816 .and_then(|v| v.as_str())
10817 .map(String::from);
10818 let (think, _lvl, _explicit) =
10819 parse_think(&effort, &None, None, None, None, false).unwrap();
10820
10821 let mut turns: Vec<TmplTurn> = Vec::new();
10822 for msg in request["messages"].as_array().unwrap() {
10823 let role = msg["role"].as_str().unwrap();
10824 let role = if role == "developer" { "system" } else { role };
10825 let content =
10826 content_to_text(msg.get("content").unwrap_or(&serde_json::Value::Null)).unwrap();
10827 let tool_calls = msg
10828 .get("tool_calls")
10829 .and_then(|a| a.as_array())
10830 .map(|a| {
10831 a.iter()
10832 .map(|tc| {
10833 let rtc: ReqToolCall = serde_json::from_value(tc.clone()).unwrap();
10834 render_req_tool_call(&rtc).unwrap()
10835 })
10836 .collect()
10837 })
10838 .unwrap_or_default();
10839 let tool_responses = msg
10840 .get("tool_responses")
10841 .and_then(|a| a.as_array())
10842 .map(|a| {
10843 a.iter()
10844 .map(|tr| {
10845 (
10846 tr.get("name").and_then(|n| n.as_str()).unwrap().to_string(),
10847 json_to_val(&tr["response"]),
10848 )
10849 })
10850 .collect()
10851 })
10852 .unwrap_or_default();
10853 turns.push(TmplTurn {
10854 role: role.to_string(),
10855 content,
10856 tool_calls,
10857 reasoning: msg
10858 .get("reasoning")
10859 .and_then(|r| r.as_str())
10860 .map(String::from)
10861 .filter(|s| !s.is_empty()),
10862 tool_call_id: msg
10863 .get("tool_call_id")
10864 .and_then(|s| s.as_str())
10865 .map(String::from),
10866 tool_name: msg.get("name").and_then(|s| s.as_str()).map(String::from),
10867 tool_responses,
10868 task: None,
10869 tools: Vec::new(),
10870 });
10871 }
10872 chat::apply_chat_template_tools_ex(
10873 Some(template),
10874 &turns,
10875 true,
10876 &tools_json,
10877 &tools_struct,
10878 think,
10879 None,
10880 None,
10881 )
10882 .unwrap()
10883 }
10884
10885 #[test]
10889 fn gemma4_tools_fixtures_match_the_official_jinja() {
10890 let dir = format!(
10891 "{}/../../research/gemma4-tools-20260817/fixtures",
10892 env!("CARGO_MANIFEST_DIR")
10893 );
10894 let mut entries: Vec<_> = std::fs::read_dir(&dir)
10895 .unwrap_or_else(|e| panic!("read fixtures dir {dir}: {e}"))
10896 .map(|e| e.unwrap().path())
10897 .filter(|p| p.is_dir())
10898 .collect();
10899 entries.sort();
10900 assert!(
10901 entries.len() >= 14,
10902 "expected >=14 fixtures, found {}",
10903 entries.len()
10904 );
10905 let (mut official, mut qat) = (0u32, 0u32);
10906 for d in entries {
10907 let input: serde_json::Value =
10908 serde_json::from_str(&std::fs::read_to_string(d.join("input.json")).unwrap())
10909 .unwrap();
10910 let expected = std::fs::read_to_string(d.join("expected.txt")).unwrap();
10911 let kind = input
10912 .get("template")
10913 .and_then(|t| t.as_str())
10914 .unwrap_or("official");
10915 match kind {
10916 "qat" => qat += 1,
10917 _ => official += 1,
10918 }
10919 let tmpl = gemma_template(kind);
10920 let got = render_fixture(&input["request"], &tmpl);
10921 assert_eq!(
10922 got, expected,
10923 "fixture {:?} diverged from the jinja oracle",
10924 d
10925 );
10926 }
10927 assert!(
10928 official >= 12 && qat >= 2,
10929 "coverage: {official} official, {qat} qat"
10930 );
10931 }
10932
10933 #[test]
10939 fn gemma4_tools_flow_through_build_chat_request() {
10940 let tmpl = gemma_template("official");
10941 for name in [
10942 "01-system-tools-basic",
10943 "04-single-call-cycle",
10944 "07-multi-cycle-agentic",
10945 ] {
10946 let path = format!(
10947 "{}/../../research/gemma4-tools-20260817/fixtures/{name}/input.json",
10948 env!("CARGO_MANIFEST_DIR")
10949 );
10950 let input: serde_json::Value =
10951 serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
10952 let expected_path = format!(
10953 "{}/../../research/gemma4-tools-20260817/fixtures/{name}/expected.txt",
10954 env!("CARGO_MANIFEST_DIR")
10955 );
10956 let expected = std::fs::read_to_string(&expected_path).unwrap();
10957 let req: ChatCompletionReq = serde_json::from_value(input["request"].clone()).unwrap();
10958 let (tx, _rx) = worker::event_channel();
10959 let plan = build_chat_request(
10960 req,
10961 Some(&gemma_tool_caps()),
10962 tx,
10963 lanes::Lane::Interactive,
10964 None,
10965 )
10966 .unwrap();
10967 let got = chat::apply_chat_template_tools_ex(
10968 Some(&tmpl),
10969 &plan.request.chat_turns,
10970 true,
10971 &plan.request.tools_json,
10972 &plan.request.tools_struct,
10973 plan.request.think,
10974 plan.request.reasoning_effort.as_deref(),
10975 None,
10976 )
10977 .unwrap();
10978 assert_eq!(got, expected, "pipeline render diverged for {name}");
10979 }
10980 }
10981
10982 fn glm5_template() -> String {
10991 let path = format!(
10992 "{}/../../research/glm53-flash-bringup-20260827/chat_template.jinja",
10993 env!("CARGO_MANIFEST_DIR")
10994 );
10995 std::fs::read_to_string(&path).unwrap_or_else(|e| panic!("read {path}: {e}"))
10996 }
10997
10998 fn glm5_caps() -> ModelCaps {
11003 ModelCaps {
11004 tools_branch: true,
11005 qwen_think: true,
11006 think_switch: false,
11007 chat_ok: true,
11008 context_length: 1_048_576,
11009 tokenizer: "glm4".into(),
11010 instruct_type: Some("glm".into()),
11011 effort_levels: true,
11012 glm5: true,
11013 ..Default::default()
11014 }
11015 }
11016
11017 fn glm5_render(body: serde_json::Value) -> Result<String, String> {
11019 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
11020 let (tx, _rx) = worker::event_channel();
11021 let plan = build_chat_request(req, Some(&glm5_caps()), tx, lanes::Lane::Interactive, None)?;
11022 chat::apply_chat_template_tools_ex(
11023 Some(&glm5_template()),
11024 &plan.request.chat_turns,
11025 true,
11026 &plan.request.tools_json,
11027 &plan.request.tools_struct,
11028 plan.request.think,
11029 plan.request.reasoning_effort.as_deref(),
11030 None,
11031 )
11032 }
11033
11034 #[test]
11039 fn glm5_fixtures_match_the_vendor_jinja() {
11040 let dir = format!(
11041 "{}/../../research/glm53-flash-bringup-20260827/surface-fixtures",
11042 env!("CARGO_MANIFEST_DIR")
11043 );
11044 let mut entries: Vec<_> = std::fs::read_dir(&dir)
11045 .unwrap_or_else(|e| panic!("read fixtures dir {dir}: {e}"))
11046 .map(|e| e.unwrap().path())
11047 .filter(|p| p.is_dir())
11048 .collect();
11049 entries.sort();
11050 assert!(
11051 entries.len() >= 22,
11052 "expected >=22 fixtures, found {}",
11053 entries.len()
11054 );
11055 for d in entries {
11056 let input: serde_json::Value =
11057 serde_json::from_str(&std::fs::read_to_string(d.join("input.json")).unwrap())
11058 .unwrap();
11059 let expected = std::fs::read_to_string(d.join("expected.txt")).unwrap();
11060 let got = glm5_render(input["request"].clone())
11061 .unwrap_or_else(|e| panic!("fixture {d:?} refused: {e}"));
11062 assert_eq!(
11063 got, expected,
11064 "fixture {d:?} diverged from the jinja oracle"
11065 );
11066 }
11067 }
11068
11069 #[test]
11075 fn glm5_never_renders_chatml() {
11076 let tmpl = glm5_template();
11077 assert!(tmpl.contains("<think>") && tmpl.contains("add_generation_prompt"));
11079 assert!(tmpl.contains("<tools>"));
11080 assert!(chat::template_is_glm5(&tmpl));
11081 for body in [
11082 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}),
11083 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
11084 "tools": [{"type": "function", "function": {"name": "f",
11085 "parameters": {"type": "object", "properties": {}}}}]}),
11086 ] {
11087 let got = glm5_render(body).unwrap();
11088 assert!(
11089 !got.contains("<|im_start|>") && !got.contains("<|im_end|>"),
11090 "glm5 rendered ChatML frames: {got:?}"
11091 );
11092 assert!(
11093 got.starts_with("[gMASK]<sop><|system|>Reasoning Effort: "),
11094 "{got:?}"
11095 );
11096 assert!(got.ends_with("<|assistant|><think>"), "{got:?}");
11097 }
11098 }
11099
11100 #[test]
11104 fn glm5_reasoning_effort_renders_and_keeps_its_max_tier() {
11105 for (sent, line) in [
11106 (None, "Max"),
11107 (Some("low"), "Low"),
11108 (Some("medium"), "Low"),
11111 (Some("high"), "High"),
11112 (Some("xhigh"), "Max"),
11113 (Some("max"), "Max"),
11114 (Some("ultra"), "Max"),
11115 ] {
11116 let mut body = json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]});
11117 if let Some(v) = sent {
11118 body["reasoning_effort"] = json!(v);
11119 }
11120 let got = glm5_render(body).unwrap();
11121 assert!(
11122 got.starts_with(&format!("[gMASK]<sop><|system|>Reasoning Effort: {line}<|")),
11123 "reasoning_effort {sent:?} should render {line:?}: {got:?}"
11124 );
11125 }
11126 let sampler_of = |v: &str| {
11129 let req: ChatCompletionReq = serde_json::from_value(
11130 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
11133 "reasoning_effort": v, "seed": 7}),
11134 )
11135 .unwrap();
11136 let (tx, _rx) = worker::event_channel();
11137 let plan =
11138 build_chat_request(req, Some(&glm5_caps()), tx, lanes::Lane::Interactive, None)
11139 .unwrap();
11140 format!("{:?}", plan.request.sampler_cfg)
11141 };
11142 assert_eq!(sampler_of("low"), sampler_of("max"));
11143 assert_eq!(canonical_effort_for("max", true), Some("max"));
11145 assert_eq!(canonical_effort_for("xhigh", true), Some("max"));
11146 assert_eq!(canonical_effort_for("max", false), Some("high"));
11147 }
11148
11149 #[test]
11153 fn glm5_refuses_what_its_template_cannot_honour() {
11154 for (value, needle) in [
11155 ("none", "cannot disable reasoning"),
11156 ("minimal", "cannot disable reasoning"),
11157 ("bogus", "bad reasoning_effort"),
11158 ] {
11159 let err = glm5_render(json!({"model": "m",
11160 "messages": [{"role": "user", "content": "hi"}],
11161 "reasoning_effort": value}))
11162 .err()
11163 .unwrap_or_else(|| panic!("reasoning_effort {value:?} must be refused"));
11164 assert!(err.contains(needle), "{value}: {err}");
11165 }
11166 }
11167
11168 #[test]
11172 fn one_glm5_request_renders_identical_bytes_on_all_three_surfaces() {
11173 let chat = json!({
11180 "model": "m",
11181 "reasoning_effort": "high",
11182 "messages": [
11183 {"role": "user", "content": "Weather in Paris and Rome?"},
11184 {"role": "assistant", "content": null,
11185 "tool_calls": [
11186 {"id": "c1", "type": "function",
11187 "function": {"name": "get_weather",
11188 "arguments": "{\"city\": \"Paris\"}"}},
11189 {"id": "c2", "type": "function",
11190 "function": {"name": "get_weather",
11191 "arguments": "{\"city\": \"Rome\"}"}}]},
11192 {"role": "tool", "tool_call_id": "c2", "content": "rome:27"},
11193 {"role": "tool", "tool_call_id": "c1", "content": "paris:21"}
11194 ],
11195 "tools": [{"type": "function", "function": {
11196 "name": "get_weather", "description": "Get the current weather for a city",
11197 "parameters": {"type": "object",
11198 "properties": {"city": {"type": "string"}},
11199 "required": ["city"]}}}]
11200 });
11201 let responses = responses_api::translate(&json!({
11202 "model": "m",
11203 "reasoning": {"effort": "high"},
11204 "input": [
11205 {"type": "message", "role": "user",
11206 "content": [{"type": "input_text", "text": "Weather in Paris and Rome?"}]},
11207 {"type": "function_call", "call_id": "c1", "name": "get_weather",
11208 "arguments": "{\"city\": \"Paris\"}"},
11209 {"type": "function_call", "call_id": "c2", "name": "get_weather",
11210 "arguments": "{\"city\": \"Rome\"}"},
11211 {"type": "function_call_output", "call_id": "c2", "output": "rome:27"},
11212 {"type": "function_call_output", "call_id": "c1", "output": "paris:21"}
11213 ],
11214 "tools": [{"type": "function", "name": "get_weather",
11215 "description": "Get the current weather for a city",
11216 "parameters": {"type": "object",
11217 "properties": {"city": {"type": "string"}},
11218 "required": ["city"]}}]
11219 }))
11220 .expect("/v1/responses translate");
11221 let messages = anthropic::translate(&json!({
11222 "model": "m",
11223 "max_tokens": 256,
11224 "output_config": {"effort": "high"},
11225 "messages": [
11226 {"role": "user", "content": "Weather in Paris and Rome?"},
11227 {"role": "assistant", "content": [
11228 {"type": "tool_use", "id": "c1", "name": "get_weather",
11229 "input": {"city": "Paris"}},
11230 {"type": "tool_use", "id": "c2", "name": "get_weather",
11231 "input": {"city": "Rome"}}]},
11232 {"role": "user", "content": [
11233 {"type": "tool_result", "tool_use_id": "c2", "content": "rome:27"},
11234 {"type": "tool_result", "tool_use_id": "c1", "content": "paris:21"}]}
11235 ],
11236 "tools": [{"name": "get_weather",
11237 "description": "Get the current weather for a city",
11238 "input_schema": {"type": "object",
11239 "properties": {"city": {"type": "string"}},
11240 "required": ["city"]}}]
11241 }))
11242 .expect("/v1/messages translate");
11243 let want = glm5_render(chat).expect("chat");
11244 assert!(
11246 want.contains(
11247 "<tool_call>get_weather<arg_key>city</arg_key><arg_value>Paris</arg_value>\
11248 </tool_call><tool_call>get_weather<arg_key>city</arg_key>\
11249 <arg_value>Rome</arg_value></tool_call>"
11250 ),
11251 "{want:?}"
11252 );
11253 assert!(
11257 want.contains(
11258 "<|observation|><tool_response>paris:21</tool_response>\
11259 <tool_response>rome:27</tool_response>"
11260 ),
11261 "{want:?}"
11262 );
11263 assert!(
11264 want.contains("<|system|>Reasoning Effort: High"),
11265 "{want:?}"
11266 );
11267 for (surface, body) in [
11268 ("/v1/responses", responses),
11269 ("/v1/messages", messages.clone()),
11270 ] {
11271 let got = glm5_render(body).unwrap_or_else(|e| panic!("{surface}: {e}"));
11272 assert_eq!(
11273 got, want,
11274 "{surface} rendered DIFFERENT glm5 prompt bytes than /v1/chat/completions"
11275 );
11276 }
11277 let mut idless = messages;
11283 for m in idless["messages"].as_array_mut().unwrap() {
11284 if m["role"] == "tool" {
11285 m.as_object_mut().unwrap().remove("tool_call_id");
11286 }
11287 }
11288 let got = glm5_render(idless).expect("id-less render");
11289 assert_ne!(
11290 got, want,
11291 "dropping tool_call_id must change the rendered order — this test cannot detect \
11292 a surface that loses ids otherwise"
11293 );
11294 assert!(
11295 got.contains(
11296 "<|observation|><tool_response>rome:27</tool_response>\
11297 <tool_response>paris:21</tool_response>"
11298 ),
11299 "{got:?}"
11300 );
11301 }
11302
11303 #[test]
11306 fn glm5_chat_arms_the_native_tool_parser() {
11307 let req: ChatCompletionReq = serde_json::from_value(json!({
11308 "model": "m", "messages": [{"role": "user", "content": "weather?"}],
11309 "tools": [{"type": "function", "function": {"name": "get_weather",
11310 "parameters": {"type": "object",
11311 "properties": {"city": {"type": "string"}}}}}]}))
11312 .unwrap();
11313 let (tx, _rx) = worker::event_channel();
11314 let plan = build_chat_request(req, Some(&glm5_caps()), tx, lanes::Lane::Interactive, None)
11315 .unwrap();
11316 let mut parser = plan.parser.expect("glm5 tools request must carry a parser");
11317 let pieces = parser.push(
11318 "reasoning here</think><tool_call>get_weather<arg_key>city</arg_key>\
11319 <arg_value>Paris</arg_value></tool_call>",
11320 );
11321 let calls: Vec<_> = pieces
11322 .iter()
11323 .filter_map(|p| match p {
11324 toolcall::Piece::Call(c) => Some((c.name.as_str(), c.arguments.as_str())),
11325 _ => None,
11326 })
11327 .collect();
11328 assert_eq!(
11329 calls,
11330 vec![("get_weather", r#"{"city":"Paris"}"#)],
11331 "{pieces:?}"
11332 );
11333 assert!(
11334 pieces
11335 .iter()
11336 .any(|p| matches!(p, toolcall::Piece::Reasoning(r) if r == "reasoning here")),
11337 "{pieces:?}"
11338 );
11339 assert!(
11341 !pieces
11342 .iter()
11343 .any(|p| matches!(p, toolcall::Piece::Content(_))),
11344 "{pieces:?}"
11345 );
11346 let req: ChatCompletionReq = serde_json::from_value(
11350 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}),
11351 )
11352 .unwrap();
11353 let (tx, _rx) = worker::event_channel();
11354 let plan = build_chat_request(req, Some(&glm5_caps()), tx, lanes::Lane::Interactive, None)
11355 .unwrap();
11356 let mut parser = plan
11357 .parser
11358 .expect("glm5 non-tools request must still split reasoning");
11359 let pieces = parser.push("weighing it</think>The answer.");
11360 assert!(
11361 pieces
11362 .iter()
11363 .any(|p| matches!(p, toolcall::Piece::Reasoning(r) if r == "weighing it")),
11364 "{pieces:?}"
11365 );
11366 assert!(
11367 pieces
11368 .iter()
11369 .any(|p| matches!(p, toolcall::Piece::Content(c) if c == "The answer.")),
11370 "{pieces:?}"
11371 );
11372 }
11373
11374 #[test]
11381 fn glm5_plain_fast_path_never_drops_replayed_reasoning() {
11382 let with_reasoning = vec![
11383 chat::Turn {
11384 role: "user".into(),
11385 content: "a".into(),
11386 ..Default::default()
11387 },
11388 chat::Turn {
11389 role: "assistant".into(),
11390 content: "A".into(),
11391 reasoning: Some("I considered a.".into()),
11392 ..Default::default()
11393 },
11394 chat::Turn {
11395 role: "user".into(),
11396 content: "b".into(),
11397 ..Default::default()
11398 },
11399 ];
11400 assert!(!worker::plain_chat_render_path(
11402 &[],
11403 &chat::ThinkMode::Default,
11404 None,
11405 &with_reasoning,
11406 false,
11407 ));
11408 let plain_turns: Vec<chat::Turn> = with_reasoning
11411 .iter()
11412 .cloned()
11413 .map(|mut t| {
11414 t.reasoning = None;
11415 t
11416 })
11417 .collect();
11418 assert!(worker::plain_chat_render_path(
11419 &[],
11420 &chat::ThinkMode::Default,
11421 None,
11422 &plain_turns,
11423 false,
11424 ));
11425 let tmpl = glm5_template();
11428 let via_tools = chat::apply_chat_template_tools_ex(
11429 Some(&tmpl),
11430 &with_reasoning,
11431 true,
11432 &[],
11433 &[],
11434 chat::ThinkMode::Default,
11435 None,
11436 None,
11437 )
11438 .unwrap();
11439 let msgs: Vec<(&str, &str)> = with_reasoning
11440 .iter()
11441 .map(|t| (t.role.as_str(), t.content.as_str()))
11442 .collect();
11443 let via_plain = chat::apply_chat_template_str(Some(&tmpl), &msgs, true);
11444 assert!(
11445 via_tools.contains("<think>I considered a.</think>"),
11446 "{via_tools:?}"
11447 );
11448 assert_ne!(via_tools, via_plain);
11449 let plain_msgs: Vec<(&str, &str)> = plain_turns
11452 .iter()
11453 .map(|t| (t.role.as_str(), t.content.as_str()))
11454 .collect();
11455 assert_eq!(
11456 chat::apply_chat_template_tools_ex(
11457 Some(&tmpl),
11458 &plain_turns,
11459 true,
11460 &[],
11461 &[],
11462 chat::ThinkMode::Default,
11463 None,
11464 None,
11465 )
11466 .unwrap(),
11467 chat::apply_chat_template_str(Some(&tmpl), &plain_msgs, true)
11468 );
11469 }
11470
11471 #[test]
11475 fn glm5_model_row_does_not_claim_structured_output() {
11476 let caps = glm5_caps();
11477 let row = model_entry_v1("zai/glm-5.3-flash", Some(&caps), None);
11478 assert_eq!(row["capabilities"]["structured_output"], json!(false));
11479 assert_eq!(row["capabilities"]["tools"], json!(true));
11480 assert_eq!(row["capabilities"]["reasoning"], json!(true));
11481 let err = glm5_render(json!({"model": "m",
11483 "messages": [{"role": "user", "content": "hi"}],
11484 "response_format": {"type": "json_object"}}))
11485 .expect_err("response_format must be refused on a switchless think template");
11486 assert!(
11490 err.contains("neither an enable_thinking switch nor a recognizable"),
11491 "{err}"
11492 );
11493 let switchable = model_entry_v1("q", Some(&tool_caps()), None);
11495 assert_eq!(switchable["capabilities"]["structured_output"], json!(true));
11496 let glm_params = openrouter_supported_parameters(Some(&caps), None, true);
11499 assert!(
11500 glm_params.get("structured_outputs").is_none(),
11501 "{glm_params}"
11502 );
11503 let step_like = ModelCaps {
11509 chat_ok: true,
11510 qwen_think: true,
11511 think_switch: false,
11512 think_close: vec![128799],
11513 ..caps.clone()
11514 };
11515 let step_row = model_entry_v1("stepfun/step-3.7-flash", Some(&step_like), None);
11516 assert_eq!(step_row["capabilities"]["structured_output"], json!(true));
11517 let step_params = openrouter_supported_parameters(Some(&step_like), None, true);
11518 assert!(
11519 step_params.get("structured_outputs").is_some(),
11520 "{step_params}"
11521 );
11522 assert!(glm_params.get("json_mode").is_none(), "{glm_params}");
11523 assert!(glm_params.get("tools").is_some(), "{glm_params}");
11524 let qwen_params = openrouter_supported_parameters(Some(&tool_caps()), None, true);
11525 assert!(
11526 qwen_params.get("structured_outputs").is_some(),
11527 "{qwen_params}"
11528 );
11529 assert!(qwen_params.get("json_mode").is_some(), "{qwen_params}");
11530 }
11531
11532 #[test]
11540 fn catalog_context_claim_is_capped_by_the_deployment_envelope() {
11541 let caps = glm5_caps();
11542 assert_eq!(caps.context_length, 1_048_576);
11543 let metadata = OpenRouterModelMetadata {
11544 max_prompt_length: Some(126_976),
11545 max_output_length: Some(4_096),
11546 ..Default::default()
11547 };
11548 let row = model_entry_v1("zai/glm-5.3-flash", Some(&caps), Some(&metadata));
11550 assert_eq!(row["context_length"], json!(131_072));
11551 let or_row = model_entry_openrouter("zai/glm-5.3-flash", Some(&caps), Some(&metadata));
11552 assert_eq!(
11553 or_row["input_modalities"][0]["supported_inputs"]["max_context_length"]["value"],
11554 json!(131_072)
11555 );
11556 assert_eq!(
11557 published_context_length(Some(&caps), Some(&metadata)),
11558 Some(131_072)
11559 );
11560 assert_eq!(published_context_length(Some(&caps), None), Some(1_048_576));
11562 let half = OpenRouterModelMetadata {
11563 max_output_length: Some(4_096),
11564 ..Default::default()
11565 };
11566 assert_eq!(
11567 published_context_length(Some(&caps), Some(&half)),
11568 Some(1_048_576)
11569 );
11570 let wide = OpenRouterModelMetadata {
11572 max_prompt_length: Some(2_000_000),
11573 max_output_length: Some(2_000_000),
11574 ..Default::default()
11575 };
11576 assert_eq!(
11577 published_context_length(Some(&caps), Some(&wide)),
11578 Some(1_048_576)
11579 );
11580 }
11581
11582 fn dsv4_sentinel() -> String {
11590 let path = format!(
11591 "{}/../../research/dsv4-template-20260818/dsv4-chat-template.sentinel.jinja",
11592 env!("CARGO_MANIFEST_DIR")
11593 );
11594 std::fs::read_to_string(&path).unwrap_or_else(|e| panic!("read {path}: {e}"))
11595 }
11596
11597 fn dsv4_turn(msg: &serde_json::Value) -> TmplTurn {
11602 let role = msg["role"].as_str().unwrap().to_string();
11603 let content =
11604 content_to_text(msg.get("content").unwrap_or(&serde_json::Value::Null)).unwrap();
11605 let reasoning = msg
11606 .get("reasoning")
11607 .or_else(|| msg.get("reasoning_content"))
11608 .and_then(|r| r.as_str())
11609 .map(String::from)
11610 .filter(|s| !s.is_empty());
11611 let tool_calls = msg
11612 .get("tool_calls")
11613 .and_then(|a| a.as_array())
11614 .map(|a| {
11615 a.iter()
11616 .map(|tc| {
11617 let rtc: ReqToolCall = serde_json::from_value(tc.clone()).unwrap();
11618 render_req_tool_call(&rtc).unwrap()
11619 })
11620 .collect()
11621 })
11622 .unwrap_or_default();
11623 let tools = msg
11624 .get("tools")
11625 .and_then(|a| a.as_array())
11626 .map(|a| {
11627 a.iter()
11628 .filter_map(|t| t.get("function").map(json_to_val))
11629 .collect()
11630 })
11631 .unwrap_or_default();
11632 TmplTurn {
11633 role,
11634 content,
11635 tool_calls,
11636 reasoning,
11637 tool_call_id: msg
11638 .get("tool_call_id")
11639 .and_then(|s| s.as_str())
11640 .map(String::from),
11641 tool_name: msg.get("name").and_then(|s| s.as_str()).map(String::from),
11642 tool_responses: Vec::new(),
11643 task: msg.get("task").and_then(|s| s.as_str()).map(String::from),
11644 tools,
11645 }
11646 }
11647
11648 fn dsv4_req_tools(v: Option<&serde_json::Value>) -> Vec<chat::Val> {
11649 v.and_then(|t| t.as_array())
11650 .map(|a| {
11651 a.iter()
11652 .filter_map(|t| t.get("function").map(json_to_val))
11653 .collect()
11654 })
11655 .unwrap_or_default()
11656 }
11657
11658 fn dsv4_run_fixture_dir(subdir: &str, encoding: chat::Dsv4Encoding, min_fixtures: usize) {
11662 let dir = format!(
11663 "{}/../../research/dsv4-template-20260818/{subdir}",
11664 env!("CARGO_MANIFEST_DIR")
11665 );
11666 let tmpl = dsv4_sentinel();
11667 let mut entries: Vec<_> = std::fs::read_dir(&dir)
11668 .unwrap_or_else(|e| panic!("read fixtures dir {dir}: {e}"))
11669 .map(|e| e.unwrap().path())
11670 .filter(|p| p.is_dir())
11671 .collect();
11672 entries.sort();
11673 assert!(
11674 entries.len() >= min_fixtures,
11675 "expected >={min_fixtures} fixtures, found {}",
11676 entries.len()
11677 );
11678 for d in &entries {
11679 let input: serde_json::Value =
11680 serde_json::from_str(&std::fs::read_to_string(d.join("input.json")).unwrap())
11681 .unwrap();
11682 let expected = std::fs::read_to_string(d.join("expected.txt")).unwrap();
11683 let turns: Vec<TmplTurn> = input["turns"]
11684 .as_array()
11685 .unwrap()
11686 .iter()
11687 .map(dsv4_turn)
11688 .collect();
11689 let think = match input["think"].as_str().unwrap() {
11690 "chat" => ThinkMode::NoThink,
11691 _ => ThinkMode::Think,
11692 };
11693 let effort = input
11694 .get("reasoning_effort")
11695 .and_then(|v| v.as_str())
11696 .map(String::from);
11697 let req_tools = dsv4_req_tools(input.get("req_tools"));
11698 let agp = input["add_generation_prompt"].as_bool().unwrap_or(true);
11699 let got = chat::apply_chat_template_tools_ex(
11700 Some(&tmpl),
11701 &turns,
11702 agp,
11703 &[],
11704 &req_tools,
11705 think,
11706 effort.as_deref(),
11707 Some(encoding),
11708 )
11709 .unwrap();
11710 assert_eq!(got, expected, "fixture {:?} diverged from the oracle", d);
11711 }
11712 }
11713
11714 #[test]
11715 fn dsv4_template_fixtures_match_the_oracle() {
11716 dsv4_run_fixture_dir("fixtures", chat::Dsv4Encoding::Preview, 20);
11717 }
11718
11719 #[test]
11725 fn dsv4_0731_fixtures_match_the_oracle() {
11726 dsv4_run_fixture_dir("fixtures-0731", chat::Dsv4Encoding::V0731, 40);
11727 }
11728
11729 #[test]
11730 fn dsv4_artifact_fixtures_are_byte_identical() {
11731 let base = format!(
11735 "{}/../../research/dsv4-template-20260818/ref/artifact-encoding/tests",
11736 env!("CARGO_MANIFEST_DIR")
11737 );
11738 let tmpl = dsv4_sentinel();
11739 for (n, think) in [
11740 (1u32, ThinkMode::Think),
11741 (2, ThinkMode::Think),
11742 (3, ThinkMode::Think),
11743 (4, ThinkMode::NoThink),
11744 ] {
11745 let td: serde_json::Value = serde_json::from_str(
11746 &std::fs::read_to_string(format!("{base}/test_input_{n}.json")).unwrap(),
11747 )
11748 .unwrap();
11749 let (messages, tools) = if td.is_object() {
11750 (td["messages"].clone(), td.get("tools").cloned())
11751 } else {
11752 (td.clone(), None)
11753 };
11754 let mut turns: Vec<TmplTurn> = Vec::new();
11755 for (i, msg) in messages.as_array().unwrap().iter().enumerate() {
11756 let mut t = dsv4_turn(msg);
11757 if i == 0
11758 && let Some(tl) = &tools
11759 {
11760 t.tools = tl
11761 .as_array()
11762 .unwrap()
11763 .iter()
11764 .filter_map(|x| x.get("function").map(json_to_val))
11765 .collect();
11766 }
11767 turns.push(t);
11768 }
11769 let expected = std::fs::read_to_string(format!("{base}/test_output_{n}.txt")).unwrap();
11770 for encoding in [chat::Dsv4Encoding::Preview, chat::Dsv4Encoding::V0731] {
11774 let got = chat::apply_chat_template_tools_ex(
11775 Some(&tmpl),
11776 &turns,
11777 true,
11778 &[],
11779 &[],
11780 think,
11781 None,
11782 Some(encoding),
11783 )
11784 .unwrap();
11785 assert_eq!(
11786 got, expected,
11787 "artifact fixture {n} diverged from the oracle under {encoding:?}"
11788 );
11789 }
11790 }
11791 }
11792
11793 #[test]
11794 fn dsv4_default_thinkmode_renders_thinking() {
11795 let tmpl = dsv4_sentinel();
11798 let turns = vec![TmplTurn {
11799 role: "user".into(),
11800 content: "Hi".into(),
11801 ..Default::default()
11802 }];
11803 let dflt = chat::apply_chat_template_tools_ex(
11804 Some(&tmpl),
11805 &turns,
11806 true,
11807 &[],
11808 &[],
11809 ThinkMode::Default,
11810 None,
11811 None,
11812 )
11813 .unwrap();
11814 let think = chat::apply_chat_template_tools_ex(
11815 Some(&tmpl),
11816 &turns,
11817 true,
11818 &[],
11819 &[],
11820 ThinkMode::Think,
11821 None,
11822 None,
11823 )
11824 .unwrap();
11825 assert_eq!(dflt, think);
11826 assert!(
11827 dflt.ends_with("<\u{ff5c}Assistant\u{ff5c}><think>"),
11828 "{dflt:?}"
11829 );
11830 let chat_mode = chat::apply_chat_template_tools_ex(
11831 Some(&tmpl),
11832 &turns,
11833 true,
11834 &[],
11835 &[],
11836 ThinkMode::NoThink,
11837 None,
11838 None,
11839 )
11840 .unwrap();
11841 assert!(
11842 chat_mode.ends_with("<\u{ff5c}Assistant\u{ff5c}></think>"),
11843 "{chat_mode:?}"
11844 );
11845 }
11846
11847 fn dsv4_run_tokenization_crosscheck(subdir: &str) {
11852 let base = format!(
11853 "{}/../../research/dsv4-template-20260818",
11854 env!("CARGO_MANIFEST_DIR")
11855 );
11856 let refdir = std::path::Path::new(&base).join("ref");
11857 let tok = memra_tokenizer::Tokenizer::from_hf_dir(&refdir)
11858 .expect("load dsv4 tokenizer from ref dir");
11859 assert_eq!(tok.pre(), "deepseek-v3", "pre-tokenizer family detection");
11860 let banked: serde_json::Value = serde_json::from_str(
11861 &std::fs::read_to_string(format!("{base}/{subdir}/tokenization-crosscheck.json"))
11862 .unwrap(),
11863 )
11864 .unwrap();
11865 let obj = banked.as_object().unwrap();
11866 assert!(obj.len() >= 3, "expected >=3 cross-check fixtures");
11867 for (name, ids_v) in obj {
11868 let rendered =
11869 std::fs::read_to_string(format!("{base}/{subdir}/{name}/expected.txt")).unwrap();
11870 let want: Vec<u32> = ids_v
11871 .as_array()
11872 .unwrap()
11873 .iter()
11874 .map(|v| v.as_u64().unwrap() as u32)
11875 .collect();
11876 let got = tok.encode(&rendered, true);
11877 assert_eq!(got, want, "tokenization diverged for {name}");
11878 }
11879 }
11880
11881 #[test]
11882 fn dsv4_tokenization_crosscheck_matches_official_ids() {
11883 dsv4_run_tokenization_crosscheck("fixtures");
11884 }
11885
11886 #[test]
11890 fn dsv4_0731_tokenization_crosscheck_matches_official_ids() {
11891 dsv4_run_tokenization_crosscheck("fixtures-0731");
11892 }
11893
11894 #[test]
11895 fn dsv4_tool_result_long_runs_render_tokenize_roundtrip() {
11896 let base = format!(
11906 "{}/../../research/dsv4-template-20260818",
11907 env!("CARGO_MANIFEST_DIR")
11908 );
11909 let refdir = std::path::Path::new(&base).join("ref");
11910 let tok = memra_tokenizer::Tokenizer::from_hf_dir(&refdir)
11911 .expect("load dsv4 tokenizer from ref dir");
11912 assert_eq!(tok.pre(), "deepseek-v3", "pre-tokenizer family detection");
11913 let tmpl = dsv4_sentinel();
11914 let req_tools = dsv4_req_tools(Some(&serde_json::json!([
11915 {"type": "function", "function": {
11916 "name": "get_data",
11917 "description": "Fetch a blob",
11918 "parameters": {"type": "object", "properties": {"key": {"type": "string"}},
11919 "required": ["key"]}
11920 }}
11921 ])));
11922
11923 let cases: Vec<(&str, String)> = vec![
11924 ("ascii-letter-131k", "Z".repeat(131_072)), ("ascii-letter-1m", "Z".repeat(1_048_576)),
11926 ("space-131k", " ".repeat(131_072)),
11927 ("digit-131k", "7".repeat(131_072)),
11928 (
11929 "mixed-runs",
11930 format!(
11931 "{}{}{}{}",
11932 "Z".repeat(65_536),
11933 " ".repeat(65_536),
11934 "7".repeat(65_536),
11935 "\n".repeat(65_536)
11936 ),
11937 ),
11938 ("cjk-64k", "中".repeat(65_536)),
11939 ("accented-letter-64k", "é".repeat(65_536)),
11940 ];
11941 for (name, blob) in &cases {
11942 let msgs = serde_json::json!([
11943 {"role": "system", "content": "You are a tool-using assistant."},
11944 {"role": "user", "content": "Fetch the blob."},
11945 {"role": "assistant", "reasoning": "Use get_data.", "content": "",
11946 "tool_calls": [{"id": "call_001", "type": "function",
11947 "function": {"name": "get_data",
11948 "arguments": "{\"key\": \"blob\"}"}}]},
11949 {"role": "tool", "tool_call_id": "call_001", "content": blob}
11950 ]);
11951 let turns: Vec<TmplTurn> = msgs.as_array().unwrap().iter().map(dsv4_turn).collect();
11952 let rendered = chat::apply_chat_template_tools_ex(
11953 Some(&tmpl),
11954 &turns,
11955 true,
11956 &[],
11957 &req_tools,
11958 ThinkMode::Think,
11959 None,
11960 None,
11961 )
11962 .unwrap_or_else(|e| panic!("{name}: render failed: {e}"));
11963 assert!(
11964 rendered.contains(blob.as_str()),
11965 "{name}: tool result missing from render"
11966 );
11967 let t0 = std::time::Instant::now();
11968 let ids = tok.encode(&rendered, true);
11969 let encode_dt = t0.elapsed();
11970 assert!(!ids.is_empty(), "{name}: empty encode");
11971 let back = tok.decode(&ids);
11972 assert_eq!(back, rendered, "{name}: decode(encode(x)) != x");
11973 assert!(
11977 encode_dt < std::time::Duration::from_secs(60),
11978 "{name}: encode took {encode_dt:?}"
11979 );
11980 if *name == "ascii-letter-131k"
11984 && let Ok(dir) = std::env::var("DSV4_LONGRUN_DUMP_DIR")
11985 {
11986 std::fs::write(format!("{dir}/rendered-131k.txt"), &rendered).unwrap();
11987 let csv: Vec<String> = ids.iter().map(|i| i.to_string()).collect();
11988 std::fs::write(format!("{dir}/memra-ids-131k.csv"), csv.join(",")).unwrap();
11989 }
11990 }
11991 }
11992
11993 #[test]
11994 fn models_v1_entry_advertises_thinking_support() {
11995 let step_caps = ModelCaps {
11998 effort_levels: true,
11999 ..tool_caps()
12000 };
12001 let entry = model_entry_v1("stepfun/step-3.7-flash", Some(&step_caps), None);
12002 assert_eq!(entry["capabilities"]["reasoning"], true);
12003 assert_eq!(entry["capabilities"]["tools"], true);
12004
12005 let plain = ModelCaps {
12007 chat_ok: true,
12008 ..Default::default()
12009 };
12010 let entry = model_entry_v1("plain", Some(&plain), None);
12011 assert_eq!(entry["capabilities"]["reasoning"], false);
12012 assert_eq!(entry["capabilities"]["tools"], false);
12013 let entry = model_entry_v1("unknown", None, None);
12015 assert_eq!(entry["capabilities"]["reasoning"], false);
12016 assert_eq!(entry["capabilities"]["streaming"], true);
12017 }
12018
12019 #[test]
12020 fn chat_request_preserves_turns_and_openai_stop_forms() {
12021 let payload = serde_json::json!({
12022 "model": "plain_quant",
12023 "messages": [
12024 {"role": "system", "content": "rules"},
12025 {"role": "developer", "content": "dev rules"},
12026 {"role": "user", "content": "task"},
12027 {"role": "assistant", "content": "work"}
12028 ],
12029 "max_tokens": 64,
12030 "temperature": 0.0,
12031 "stop": "<stop>"
12032 });
12033 let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
12034 let (tx, _rx) = worker::event_channel();
12035 let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
12036 let request = plan.request;
12037 assert!(
12038 plan.parser.is_none(),
12039 "no tools -> no parser (isolation contract)"
12040 );
12041 assert!(request.tools_json.is_empty());
12042 assert_eq!(request.think, ThinkMode::Default);
12043 assert_eq!(request.model, "plain_quant");
12044 assert_eq!(request.params.max_new, 64);
12045 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12047 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}]
12048 }))
12049 .unwrap();
12050 let (tx, _rx) = worker::event_channel();
12051 let plan = build_chat_request(req, None, tx, lanes::Lane::Interactive, None).unwrap();
12052 assert_eq!(plan.request.params.max_new, worker::MAX_NEW_CTX_BOUNDED);
12053 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12055 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
12056 "max_completion_tokens": 7
12057 }))
12058 .unwrap();
12059 let (tx, _rx) = worker::event_channel();
12060 assert_eq!(
12061 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
12062 .unwrap()
12063 .request
12064 .params
12065 .max_new,
12066 7
12067 );
12068 let req: CompletionReq = serde_json::from_value(serde_json::json!({
12070 "model": "plain_quant", "prompt": "task"
12071 }))
12072 .unwrap();
12073 let (tx, _rx) = worker::event_channel();
12074 assert_eq!(
12075 build_request(&req, tx, lanes::Lane::Interactive, None)
12076 .params
12077 .max_new,
12078 worker::MAX_NEW_CTX_BOUNDED
12079 );
12080 let turns: Vec<(String, String)> = request
12081 .chat_turns
12082 .iter()
12083 .map(|t| (t.role.clone(), t.content.clone()))
12084 .collect();
12085 assert_eq!(
12086 turns,
12087 vec![
12088 ("system".into(), "rules".into()),
12089 ("system".into(), "dev rules".into()), ("user".into(), "task".into()),
12091 ("assistant".into(), "work".into()),
12092 ]
12093 );
12094 assert!(request.chat_turns.iter().all(|t| t.tool_calls.is_empty()));
12095 assert_eq!(request.stop_strings, vec!["<stop>"]);
12096
12097 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12098 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
12099 "stop": ["a", "b"]
12100 }))
12101 .unwrap();
12102 assert_eq!(req.stop.into_vec(), vec!["a", "b"]);
12103
12104 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12108 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
12109 "stop": ["", "real", ""]
12110 }))
12111 .unwrap();
12112 assert_eq!(req.stop.into_vec(), vec!["real"]);
12113 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12114 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
12115 "stop": ""
12116 }))
12117 .unwrap();
12118 assert!(req.stop.into_vec().is_empty());
12119
12120 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
12121 "model": "plain_quant", "messages": [{"role": "user", "content": "task"}],
12122 "stop": null
12123 }))
12124 .unwrap();
12125 assert!(req.stop.into_vec().is_empty());
12126 }
12127
12128 #[test]
12129 fn stop_sequence_limits_bound_count_individual_and_aggregate_work() {
12130 let at_limit = StopSequences::Many(vec!["x".repeat(256); MAX_STOP_SEQUENCES]);
12131 assert!(at_limit.validate().is_ok());
12132 assert!(
12133 StopSequences::Many(vec![String::new(); MAX_STOP_SEQUENCES + 1])
12134 .validate()
12135 .unwrap_err()
12136 .contains("at most")
12137 );
12138 assert!(
12139 StopSequences::One("x".repeat(MAX_STOP_SEQUENCE_BYTES + 1))
12140 .validate()
12141 .unwrap_err()
12142 .contains("each stop")
12143 );
12144 assert!(
12145 StopSequences::Many(vec!["x".repeat(300); MAX_STOP_SEQUENCES])
12146 .validate()
12147 .unwrap_err()
12148 .contains("total at most")
12149 );
12150 }
12151
12152 #[tokio::test]
12153 async fn chat_response_has_openai_message_shape() {
12154 let (tx, rx) = worker::event_channel();
12155 tx.send(Event::Token {
12156 id: 1,
12157 text: "hello".into(),
12158 })
12159 .unwrap();
12160 tx.send(Event::Done {
12161 stop_reason: "Eos".into(),
12162 n_tokens: 1,
12163 n_prompt: 42,
12164 n_cached: 30,
12165 elapsed_s: 0.5,
12166 spec: None,
12167 })
12168 .unwrap();
12169 drop(tx);
12170 let response = blocking_response(
12171 rx,
12172 "plain_quant".into(),
12173 true,
12174 Vec::new(),
12175 None,
12176 Envelope::new(true),
12177 )
12178 .await;
12179 assert_eq!(response.status(), StatusCode::OK);
12180 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
12181 .await
12182 .unwrap();
12183 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
12184 assert_eq!(payload["object"], "chat.completion");
12185 assert!(payload["id"].as_str().unwrap().starts_with("chatcmpl-"));
12187 assert!(payload["created"].as_u64().unwrap() > 1_700_000_000);
12188 let fingerprint = payload["system_fingerprint"].as_str().unwrap();
12192 assert!(
12193 build_id::fingerprint_is_well_formed(fingerprint),
12194 "system_fingerprint {fingerprint:?} is not memra-<version>-<12 hex>"
12195 );
12196 assert_eq!(payload["choices"][0]["message"]["role"], "assistant");
12197 assert_eq!(payload["choices"][0]["message"]["content"], "hello");
12198 assert_eq!(payload["choices"][0]["finish_reason"], "stop");
12199 assert_eq!(payload["usage"]["prompt_tokens"], 42);
12201 assert_eq!(payload["usage"]["completion_tokens"], 1);
12202 assert_eq!(payload["usage"]["total_tokens"], 43);
12203 assert_eq!(
12204 payload["usage"]["prompt_tokens_details"]["cached_tokens"],
12205 30
12206 );
12207 assert!(payload["usage"].get("spec").is_none());
12210 }
12211
12212 #[tokio::test]
12213 async fn native_response_uses_terminal_token_snapshot_for_coalesced_events() {
12214 let (tx, rx) = worker::event_channel();
12215 tx.send(Event::Token {
12217 id: 4,
12218 text: "hello".into(),
12219 })
12220 .unwrap();
12221 tx.send(Event::TokenSnapshot(vec![1, 2, 3, 4])).unwrap();
12222 tx.send(Event::Done {
12223 stop_reason: "MaxNew".into(),
12224 n_tokens: 4,
12225 n_prompt: 2,
12226 n_cached: 0,
12227 elapsed_s: 0.5,
12228 spec: None,
12229 })
12230 .unwrap();
12231 drop(tx);
12232
12233 let response = blocking_response(
12234 rx,
12235 "plain_quant".into(),
12236 false,
12237 Vec::new(),
12238 None,
12239 Envelope::new(false),
12240 )
12241 .await;
12242 assert_eq!(response.status(), StatusCode::OK);
12243 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
12244 .await
12245 .unwrap();
12246 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
12247 assert_eq!(payload["text"], "hello");
12248 assert_eq!(payload["tokens"], serde_json::json!([1, 2, 3, 4]));
12249 assert_eq!(payload["n_tokens"], 4);
12250 }
12251
12252 #[tokio::test]
12255 async fn chat_usage_carries_spec_acceptance_summary() {
12256 let (tx, rx) = worker::event_channel();
12257 tx.send(Event::Token {
12258 id: 1,
12259 text: "hello".into(),
12260 })
12261 .unwrap();
12262 tx.send(Event::Done {
12263 stop_reason: "Eos".into(),
12264 n_tokens: 1,
12265 n_prompt: 42,
12266 n_cached: 0,
12267 elapsed_s: 0.5,
12268 spec: Some(worker::SpecUsage {
12269 rounds: 10,
12270 drafted: 30,
12271 accepted: 21,
12272 }),
12273 })
12274 .unwrap();
12275 drop(tx);
12276 let response = blocking_response(
12277 rx,
12278 "plain_quant".into(),
12279 true,
12280 Vec::new(),
12281 None,
12282 Envelope::new(true),
12283 )
12284 .await;
12285 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
12286 .await
12287 .unwrap();
12288 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
12289 let sp = &payload["usage"]["spec"];
12290 assert_eq!(sp["rounds"], 10);
12291 assert_eq!(sp["drafted"], 30);
12292 assert_eq!(sp["accepted"], 21);
12293 assert!((sp["acceptance_rate"].as_f64().unwrap() - 0.7).abs() < 1e-9);
12294 assert_eq!(payload["usage"]["total_tokens"], 43);
12296 }
12297
12298 fn weather_request(extra: serde_json::Value) -> ChatCompletionReq {
12299 let mut payload = serde_json::json!({
12300 "model": "m",
12301 "messages": [{"role": "user", "content": "Weather in Paris?"}],
12302 "tools": [{"type": "function", "function": {
12303 "name": "get_weather",
12304 "description": "Get current weather",
12305 "parameters": {"type": "object",
12306 "properties": {"city": {"type": "string"},
12307 "days": {"type": "integer"}},
12308 "required": ["city"]}}}],
12309 });
12310 if let Some(obj) = extra.as_object() {
12311 for (k, v) in obj {
12312 payload[k] = v.clone();
12313 }
12314 }
12315 serde_json::from_value(payload).unwrap()
12316 }
12317
12318 #[test]
12322 fn glm5_vision_decode_is_deferred_and_grid_pinned() {
12323 let (tx, _rx) = worker::event_channel();
12324 let req: ChatCompletionReq = serde_json::from_value(json!({
12325 "model": "m", "messages": [{"role": "user", "content": "hi"}],
12326 }))
12327 .unwrap();
12328 let mut plan = build_chat_request(
12329 req,
12330 Some(&ModelCaps {
12331 chat_ok: true,
12332 ..Default::default()
12333 }),
12334 tx,
12335 lanes::Lane::Interactive,
12336 None,
12337 )
12338 .unwrap();
12339 let bmp = |w: u32, h: u32| -> Vec<u8> {
12342 let row = (w * 3).div_ceil(4) * 4;
12343 let size = 54 + row * h;
12344 let mut b = vec![0x42u8, 0x4d];
12345 b.extend_from_slice(&size.to_le_bytes());
12346 b.extend_from_slice(&[0; 4]);
12347 b.extend_from_slice(&54u32.to_le_bytes());
12348 b.extend_from_slice(&40u32.to_le_bytes());
12349 b.extend_from_slice(&w.to_le_bytes());
12350 b.extend_from_slice(&h.to_le_bytes());
12351 b.extend_from_slice(&1u16.to_le_bytes());
12352 b.extend_from_slice(&24u16.to_le_bytes());
12353 b.extend_from_slice(&[0u8; 24]);
12354 b.extend(std::iter::repeat_n(0x7fu8, (row * h) as usize));
12355 b
12356 };
12357 let bytes = bmp(112, 112);
12358 let (gh, gw) = memra_engine::vision_glm5::glm5_plan_image(&bytes).unwrap();
12359 assert_eq!((gh, gw), (8, 8), "identity resize grid");
12360 assert_eq!(memra_engine::vision_glm5::n_merged_for_grid(gh, gw), 16);
12361 plan.pending_glm5.push(PendingGlm5Image {
12362 bytes: bytes.clone(),
12363 gh,
12364 gw,
12365 });
12366 decode_pending_vision(&mut plan).unwrap();
12367 assert_eq!(plan.request.glm5_images.len(), 1);
12368 let unit = &plan.request.glm5_images[0];
12369 assert_eq!((unit.gh, unit.gw), (gh, gw));
12370 assert_eq!(
12371 unit.patches.len(),
12372 gh * gw * memra_engine::vision_glm5::G5V_PATCH_IN
12373 );
12374 plan.request.glm5_images.clear();
12376 plan.pending_glm5.push(PendingGlm5Image {
12377 bytes,
12378 gh: gh + 2,
12379 gw,
12380 });
12381 let err = decode_pending_vision(&mut plan).unwrap_err();
12382 assert!(err.contains("header-planned"), "got: {err}");
12383 }
12384
12385 #[test]
12386 fn vision_decode_is_deferred_and_grid_pinned() {
12387 let (tx, _rx) = worker::event_channel();
12392 let req: ChatCompletionReq = serde_json::from_value(json!({
12393 "model": "m", "messages": [{"role": "user", "content": "hi"}],
12394 }))
12395 .unwrap();
12396 let mut plan = build_chat_request(
12397 req,
12398 Some(&ModelCaps {
12399 chat_ok: true,
12400 ..Default::default()
12401 }),
12402 tx,
12403 lanes::Lane::Interactive,
12404 None,
12405 )
12406 .unwrap();
12407 let bmp = |w: i32, h: i32, with_pixels: bool| -> Vec<u8> {
12411 let mut b = Vec::new();
12412 b.extend_from_slice(b"BM");
12413 b.extend_from_slice(&54u32.to_le_bytes());
12414 b.extend_from_slice(&0u32.to_le_bytes());
12415 b.extend_from_slice(&54u32.to_le_bytes());
12416 b.extend_from_slice(&40u32.to_le_bytes());
12417 b.extend_from_slice(&w.to_le_bytes());
12418 b.extend_from_slice(&h.to_le_bytes());
12419 b.extend_from_slice(&1u16.to_le_bytes());
12420 b.extend_from_slice(&24u16.to_le_bytes());
12421 b.extend_from_slice(&[0u8; 24]);
12422 if with_pixels {
12423 b.extend(std::iter::repeat_n(0x7fu8, (w * h * 3) as usize));
12424 }
12425 b
12426 };
12427 let bytes = bmp(64, 64, true);
12428 let (gh, gw) = memra_engine::vision_pre::plan_image_bytes(&bytes).unwrap();
12429 plan.pending_images.push(PendingVisionUnit::Still {
12430 bytes: bytes.clone(),
12431 gh,
12432 gw,
12433 });
12434 decode_pending_vision(&mut plan).unwrap();
12435 assert_eq!(plan.request.images.len(), 1);
12436 assert_eq!(
12437 (
12438 plan.request.images[0].prep.gh,
12439 plan.request.images[0].prep.gw
12440 ),
12441 (gh, gw),
12442 "decoded grid must equal the header-planned grid the pad run was rendered from"
12443 );
12444 plan.request.images.clear();
12446 plan.pending_images.push(PendingVisionUnit::Still {
12447 bytes,
12448 gh: gh + 2,
12449 gw,
12450 });
12451 let err = decode_pending_vision(&mut plan).unwrap_err();
12452 assert!(err.contains("header-planned"), "got: {err}");
12453 let bomb = bmp(16_000, 16_000, false);
12456 plan.pending_images.clear();
12457 plan.pending_images.push(PendingVisionUnit::Still {
12458 bytes: bomb,
12459 gh: 2,
12460 gw: 2,
12461 });
12462 let err = decode_pending_vision(&mut plan).unwrap_err();
12463 assert!(err.contains("exceeds the decode budget"), "got: {err}");
12464 }
12465
12466 #[test]
12467 fn tools_request_renders_client_key_order_and_arms_parser() {
12468 let (tx, _rx) = worker::event_channel();
12469 let plan = build_chat_request(
12470 weather_request(json!({})),
12471 Some(&tool_caps()),
12472 tx,
12473 lanes::Lane::Interactive,
12474 None,
12475 )
12476 .unwrap();
12477 assert!(plan.parser.is_some());
12478 assert_eq!(plan.request.tools_json.len(), 1);
12479 assert_eq!(
12481 plan.request.tools_json[0],
12482 "{\"type\": \"function\", \"function\": {\"name\": \"get_weather\", \
12483 \"description\": \"Get current weather\", \"parameters\": {\"type\": \"object\", \
12484 \"properties\": {\"city\": {\"type\": \"string\"}, \"days\": {\"type\": \
12485 \"integer\"}}, \"required\": [\"city\"]}}}"
12486 );
12487 }
12488
12489 #[test]
12490 fn hy3_tools_and_reasoning_flow_through_the_real_chat_plan() {
12491 let (tx, _rx) = worker::event_channel();
12492 let plan = build_chat_request(
12493 weather_request(json!({"reasoning_effort": "high"})),
12494 Some(&hy3_tool_caps()),
12495 tx,
12496 lanes::Lane::Interactive,
12497 None,
12498 )
12499 .unwrap();
12500 assert_eq!(plan.request.think, ThinkMode::Think);
12501 assert_eq!(plan.request.reasoning_effort.as_deref(), Some("high"));
12502 assert!(
12503 plan.request
12504 .stop_strings
12505 .iter()
12506 .any(|stop| stop == "</tool_calls:opensource>")
12507 );
12508 let rendered = chat::apply_chat_template_tools_ex(
12509 Some("... hy_User ... <tools> ..."),
12510 &plan.request.chat_turns,
12511 true,
12512 &plan.request.tools_json,
12513 &plan.request.tools_struct,
12514 plan.request.think,
12515 plan.request.reasoning_effort.as_deref(),
12516 None,
12517 )
12518 .unwrap();
12519 assert!(rendered.contains("<tool_calls:opensource>"));
12520 assert!(rendered.ends_with("<think:opensource>"));
12521
12522 let mut parser = plan.parser.expect("HY3 tools arm its native parser");
12523 let pieces = parser.push(concat!(
12524 "Need weather.</think:opensource>",
12525 "<tool_calls:opensource><tool_call:opensource>get_weather",
12526 "<tool_sep:opensource>\n<arg_key:opensource>city</arg_key:opensource>\n",
12527 "<arg_value:opensource>Paris</arg_value:opensource>\n",
12528 "</tool_call:opensource></tool_calls:opensource>",
12529 ));
12530 assert!(pieces.contains(&Piece::Reasoning("Need weather.".into())));
12531 assert!(pieces.iter().any(|piece| matches!(piece, Piece::Call(call)
12532 if call.name == "get_weather" && call.arguments == r#"{"city":"Paris"}"#)));
12533 }
12534
12535 #[test]
12536 fn tool_choice_none_strips_tools_and_parser() {
12537 let (tx, _rx) = worker::event_channel();
12538 let plan = build_chat_request(
12539 weather_request(json!({"tool_choice": "none"})),
12540 Some(&tool_caps()),
12541 tx,
12542 lanes::Lane::Interactive,
12543 None,
12544 )
12545 .unwrap();
12546 let mut p = plan
12549 .parser
12550 .expect("think-open chat arms the reasoning splitter");
12551 let pieces = p.push("x</think>\n\n<tool_call> stays prose");
12552 assert_eq!(
12553 pieces,
12554 vec![
12555 Piece::Reasoning("x".into()),
12556 Piece::Content("<tool_call> stays prose".into()),
12557 ]
12558 );
12559 assert!(plan.request.tools_json.is_empty());
12560 let (tx, _rx) = worker::event_channel();
12562 assert!(
12563 build_chat_request(
12564 weather_request(json!({"tool_choice": "required"})),
12565 Some(&tool_caps()),
12566 tx,
12567 lanes::Lane::Interactive,
12568 None
12569 )
12570 .is_err()
12571 );
12572 let (tx, _rx) = worker::event_channel();
12573 assert!(
12574 build_chat_request(
12575 weather_request(json!({"tool_choice":
12576 {"type": "function", "function": {"name": "get_weather"}}})),
12577 Some(&tool_caps()),
12578 tx,
12579 lanes::Lane::Interactive,
12580 None
12581 )
12582 .is_err()
12583 );
12584 }
12585
12586 #[test]
12587 fn model_plan_accepts_st_dir_and_rejects_bogus_dir() {
12588 let root = std::env::temp_dir().join(format!("memra_plan_test_{}", std::process::id()));
12589 let _ = std::fs::remove_dir_all(&root);
12590
12591 let st = root.join("st_single");
12593 std::fs::create_dir_all(&st).unwrap();
12594 std::fs::write(st.join("config.json"), "{}").unwrap();
12595 std::fs::write(st.join("model.safetensors"), b"x").unwrap();
12596 assert!(validate_model_path(st.to_str().unwrap()).is_ok());
12597
12598 let sh = root.join("st_sharded");
12600 std::fs::create_dir_all(&sh).unwrap();
12601 std::fs::write(sh.join("config.json"), "{}").unwrap();
12602 std::fs::write(sh.join("model.safetensors.index.json"), "{}").unwrap();
12603 assert!(validate_model_path(sh.to_str().unwrap()).is_ok());
12604
12605 let rp = root.join("repack");
12607 std::fs::create_dir_all(&rp).unwrap();
12608 std::fs::write(rp.join("manifest.json"), "{}").unwrap();
12609 assert!(validate_model_path(rp.to_str().unwrap()).is_ok());
12610
12611 let bogus = root.join("bogus");
12613 std::fs::create_dir_all(&bogus).unwrap();
12614 let err = validate_model_path(bogus.to_str().unwrap()).unwrap_err();
12615 assert!(
12616 err.contains("model.safetensors"),
12617 "error should say what is missing: {err}"
12618 );
12619 assert!(
12620 err.contains("manifest.json"),
12621 "error should mention the repack form: {err}"
12622 );
12623
12624 let nc = root.join("no_config");
12626 std::fs::create_dir_all(&nc).unwrap();
12627 std::fs::write(nc.join("model.safetensors"), b"x").unwrap();
12628 let err = validate_model_path(nc.to_str().unwrap()).unwrap_err();
12629 assert!(
12630 err.contains("config.json"),
12631 "error should name config.json: {err}"
12632 );
12633
12634 let err = validate_model_path(root.join("nowhere").to_str().unwrap()).unwrap_err();
12636 assert!(err.contains("does not exist"), "{err}");
12637
12638 let f = root.join("model.gguf");
12640 std::fs::write(&f, b"g").unwrap();
12641 assert!(validate_model_path(f.to_str().unwrap()).is_ok());
12642
12643 let _ = std::fs::remove_dir_all(&root);
12644 }
12645
12646 #[test]
12647 fn chat_on_templateless_dir_checkpoint_is_rejected_with_clear_message() {
12648 let caps = ModelCaps {
12651 tools_branch: false,
12652 qwen_think: false,
12653 think_switch: false,
12654 chat_ok: false,
12655 ..Default::default()
12656 };
12657 let payload = serde_json::json!({
12658 "model": "st_model",
12659 "messages": [{"role": "user", "content": "hello"}],
12660 });
12661 let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
12662 let (tx, _rx) = worker::event_channel();
12663 let err = match build_chat_request(req, Some(&caps), tx, lanes::Lane::Interactive, None) {
12664 Err(e) => e,
12665 Ok(_) => panic!("templateless dir checkpoint must reject chat"),
12666 };
12667 assert!(
12668 err.contains("no chat template"),
12669 "message should name the cause: {err}"
12670 );
12671 assert!(
12672 err.contains("/v1/completions"),
12673 "message should point at the raw-prompt escape hatch: {err}"
12674 );
12675 }
12676
12677 #[test]
12678 fn tools_on_model_without_tools_branch_is_rejected() {
12679 let (tx, _rx) = worker::event_channel();
12680 let caps = ModelCaps {
12681 chat_ok: true,
12682 ..Default::default()
12683 };
12684 assert!(
12685 build_chat_request(
12686 weather_request(json!({})),
12687 Some(&caps),
12688 tx,
12689 lanes::Lane::Interactive,
12690 None
12691 )
12692 .is_err()
12693 );
12694 let (tx, _rx) = worker::event_channel();
12695 assert!(
12696 build_chat_request(
12697 weather_request(json!({})),
12698 None,
12699 tx,
12700 lanes::Lane::Interactive,
12701 None
12702 )
12703 .is_err()
12704 );
12705 }
12706
12707 #[test]
12708 fn reasoning_effort_maps_to_think_switch() {
12709 for (extra, want) in [
12715 (json!({}), ThinkMode::Default),
12716 (json!({"reasoning_effort": "low"}), ThinkMode::Think),
12717 (json!({"reasoning_effort": "none"}), ThinkMode::NoThink),
12718 (json!({"reasoning_effort": "minimal"}), ThinkMode::NoThink),
12719 (json!({"reasoning_effort": "high"}), ThinkMode::Think),
12720 (json!({"reasoning_effort": "medium"}), ThinkMode::Think),
12721 (json!({"reasoning": {"enabled": false}}), ThinkMode::NoThink),
12722 (json!({"reasoning": {"effort": "low"}}), ThinkMode::Think),
12723 (json!({"reasoning": {"enabled": true}}), ThinkMode::Think),
12724 (json!({"reasoning_effort": "xhigh"}), ThinkMode::Think),
12728 (json!({"reasoning_effort": "max"}), ThinkMode::Think),
12729 (json!({"reasoning_effort": "ultra"}), ThinkMode::Think),
12730 (
12734 json!({"reasoning": {"enabled": true, "effort": "none"}}),
12735 ThinkMode::Think,
12736 ),
12737 (
12738 json!({"reasoning": {"enabled": false, "effort": "high"}}),
12739 ThinkMode::NoThink,
12740 ),
12741 ] {
12742 let (tx, _rx) = worker::event_channel();
12743 let plan = build_chat_request(
12744 weather_request(extra.clone()),
12745 Some(&ladder_caps()),
12750 tx,
12751 lanes::Lane::Interactive,
12752 None,
12753 )
12754 .unwrap();
12755 assert_eq!(plan.request.think, want, "extra={extra}");
12756 }
12757 for extra in [
12761 json!({"reasoning_effort": "extreme"}),
12762 json!({"reasoning": {"effort": "banana"}}),
12763 json!({"reasoning": {"enabled": false, "effort": "banana"}}),
12764 json!({"reasoning": {"enabled": true, "effort": ""}}),
12765 ] {
12766 let (tx, _rx) = worker::event_channel();
12767 assert!(
12768 build_chat_request(
12769 weather_request(extra.clone()),
12770 Some(&tool_caps()),
12771 tx,
12772 lanes::Lane::Interactive,
12773 None
12774 )
12775 .is_err(),
12776 "extra={extra} must be rejected by the one allowlist"
12777 );
12778 }
12779 for (raw, want) in [
12782 ("none", Some("none")),
12783 ("minimal", Some("minimal")),
12784 ("low", Some("low")),
12785 ("medium", Some("medium")),
12786 ("high", Some("high")),
12787 ("xhigh", Some("high")),
12788 ("max", Some("high")),
12789 ("ultra", Some("high")),
12790 ("banana", None),
12791 ("", None),
12792 ("HIGH", None),
12793 ] {
12794 assert_eq!(canonical_effort(raw), want, "canonical_effort({raw:?})");
12795 }
12796 for (raw, want) in [
12799 ("none", Some("none")),
12800 ("minimal", Some("minimal")),
12801 ("low", Some("low")),
12802 ("medium", Some("medium")),
12803 ("high", Some("high")),
12804 ("xhigh", Some("max")),
12805 ("max", Some("max")),
12806 ("ultra", Some("max")),
12807 ("banana", None),
12808 ("", None),
12809 ("MAX", None),
12810 ] {
12811 assert_eq!(
12812 canonical_effort_for(raw, true),
12813 want,
12814 "canonical_effort_for({raw:?}, dsv4)"
12815 );
12816 }
12817 }
12818
12819 #[test]
12820 fn dsv4_reasoning_effort_max_survives_canonicalization() {
12821 let dsv4_caps = ModelCaps {
12827 chat_ok: true,
12828 dsv4: true,
12829 ..Default::default()
12830 };
12831 let build = |caps: &ModelCaps, effort: &str| {
12832 let (tx, _rx) = worker::event_channel();
12833 let req: ChatCompletionReq = serde_json::from_value(json!({
12834 "model": "m",
12835 "messages": [{"role": "user", "content": "hi"}],
12836 "reasoning_effort": effort,
12837 }))
12838 .unwrap();
12839 build_chat_request(req, Some(caps), tx, lanes::Lane::Interactive, None)
12840 };
12841 for raw in ["max", "xhigh", "ultra"] {
12842 let plan = build(&dsv4_caps, raw).unwrap();
12843 assert_eq!(
12844 plan.request.reasoning_effort.as_deref(),
12845 Some("max"),
12846 "dsv4 {raw:?} must reach the renderer as the max rung"
12847 );
12848 assert_eq!(plan.request.think, chat::ThinkMode::Think);
12849 }
12850 let plan = build(&dsv4_caps, "high").unwrap();
12852 assert_eq!(plan.request.reasoning_effort.as_deref(), Some("high"));
12853 let step_caps = ModelCaps {
12855 chat_ok: true,
12856 effort_levels: true,
12857 ..Default::default()
12858 };
12859 let plan = build(&step_caps, "max").unwrap();
12860 assert_eq!(plan.request.reasoning_effort.as_deref(), Some("high"));
12861 }
12862
12863 #[test]
12864 fn default_reasoning_effort_flips_only_the_unset_request() {
12865 let build = |extra: serde_json::Value, default_effort: Option<&str>| {
12872 let (tx, _rx) = worker::event_channel();
12873 build_chat_request_with_trace(
12874 weather_request(extra),
12875 Some(&ladder_caps()),
12876 tx,
12877 lanes::Lane::Interactive,
12878 None,
12879 None,
12880 default_effort,
12881 &ModelSamplingDefaults::default(),
12882 )
12883 .unwrap()
12884 };
12885 for (extra, want) in [
12886 (json!({}), ThinkMode::Think),
12888 (json!({"reasoning": {"exclude": true}}), ThinkMode::NoThink),
12892 (json!({"include_reasoning": false}), ThinkMode::NoThink),
12893 (json!({"reasoning": {"exclude": false}}), ThinkMode::Think),
12895 (json!({"include_reasoning": true}), ThinkMode::Think),
12896 (json!({"reasoning_effort": "none"}), ThinkMode::NoThink),
12898 (json!({"reasoning_effort": "minimal"}), ThinkMode::NoThink),
12899 (json!({"reasoning": {"enabled": false}}), ThinkMode::NoThink),
12900 (json!({"reasoning_effort": "low"}), ThinkMode::Think),
12902 (json!({"reasoning_effort": "high"}), ThinkMode::Think),
12903 (json!({"reasoning": {"enabled": true}}), ThinkMode::Think),
12904 ] {
12905 let plan = build(extra.clone(), Some("high"));
12906 assert_eq!(plan.request.think, want, "extra={extra}");
12907 }
12908 assert_eq!(
12910 build(json!({}), Some("none")).request.think,
12911 ThinkMode::NoThink
12912 );
12913 assert_eq!(
12914 build(json!({"reasoning_effort": "high"}), Some("none"))
12915 .request
12916 .think,
12917 ThinkMode::Think
12918 );
12919 assert_eq!(build(json!({}), None).request.think, ThinkMode::Default);
12923 }
12924
12925 const SWITCHED_QWEN_TMPL: &str = "<tools> ... add_generation_prompt ... \
12931 {%- if enable_thinking is defined and enable_thinking is false %}'<think>\\n\\n</think>\\n\\n'\
12932 {%- else %}'<think>\\n'{%- endif %}";
12933
12934 #[test]
12935 fn vllm_enable_thinking_switch_is_wired_not_ignored() {
12936 let build = |extra: serde_json::Value| {
12942 let (tx, _rx) = worker::event_channel();
12943 build_chat_request(
12944 weather_request(extra),
12945 Some(&tool_caps()),
12946 tx,
12947 lanes::Lane::Interactive,
12948 None,
12949 )
12950 };
12951 for (extra, want) in [
12952 (json!({"enable_thinking": false}), ThinkMode::NoThink),
12953 (json!({"enable_thinking": true}), ThinkMode::Think),
12954 (
12955 json!({"chat_template_kwargs": {"enable_thinking": false}}),
12956 ThinkMode::NoThink,
12957 ),
12958 (
12959 json!({"chat_template_kwargs": {"enable_thinking": true}}),
12960 ThinkMode::Think,
12961 ),
12962 (
12965 json!({"enable_thinking": false, "reasoning_effort": "high"}),
12966 ThinkMode::NoThink,
12967 ),
12968 (
12970 json!({"enable_thinking": false,
12971 "chat_template_kwargs": {"enable_thinking": false}}),
12972 ThinkMode::NoThink,
12973 ),
12974 ] {
12975 let plan = build(extra.clone()).unwrap_or_else(|e| {
12976 panic!("{extra} must be accepted and honored, got 400: {e}");
12977 });
12978 assert_eq!(
12979 plan.request.think, want,
12980 "{extra} was ACCEPTED AND IGNORED — the banned silent-accept class"
12981 );
12982 }
12983 let render = |extra: serde_json::Value| -> String {
12986 let plan = build(extra).unwrap();
12987 chat::apply_chat_template_tools_ex(
12988 Some(SWITCHED_QWEN_TMPL),
12989 &plan.request.chat_turns,
12990 true,
12991 &plan.request.tools_json,
12992 &plan.request.tools_struct,
12993 plan.request.think,
12994 plan.request.reasoning_effort.as_deref(),
12995 None,
12996 )
12997 .unwrap()
12998 };
12999 let off = render(json!({"enable_thinking": false}));
13000 assert!(
13001 off.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"),
13002 "enable_thinking:false must render the CLOSED think pair: {off:?}"
13003 );
13004 let on = render(json!({}));
13005 assert!(
13006 on.ends_with("<|im_start|>assistant\n<think>\n"),
13007 "an unset request must still render the template's OPEN think tail: {on:?}"
13008 );
13009 assert_eq!(
13010 off,
13011 render(json!({"chat_template_kwargs": {"enable_thinking": false}})),
13012 "both vLLM spellings must render byte-identically"
13013 );
13014 assert_eq!(
13015 off,
13016 render(json!({"reasoning_effort": "none"})),
13017 "the vLLM spelling must render byte-identically to the OpenAI spelling"
13018 );
13019 }
13020
13021 #[test]
13022 fn unknown_chat_template_kwarg_refuses_by_name() {
13023 let build = |extra: serde_json::Value| {
13026 let (tx, _rx) = worker::event_channel();
13027 build_chat_request(
13028 weather_request(extra),
13029 Some(&tool_caps()),
13030 tx,
13031 lanes::Lane::Interactive,
13032 None,
13033 )
13034 };
13035 let refusal = |extra: serde_json::Value, why: &str| -> String {
13036 build(extra).err().unwrap_or_else(|| panic!("{why}"))
13037 };
13038 let err = refusal(
13039 json!({"chat_template_kwargs": {"add_generation_prompt": false}}),
13040 "an unimplementable template kwarg must not be accepted",
13041 );
13042 assert!(
13043 err.contains("add_generation_prompt") && err.contains("enable_thinking"),
13044 "the refusal must name the offending key AND the supported one: {err}"
13045 );
13046 let err = refusal(
13047 json!({"chat_template_kwargs": "enable_thinking=false"}),
13048 "a non-object chat_template_kwargs must not be accepted",
13049 );
13050 assert!(
13051 err.contains("must be an object"),
13052 "refusal must say what shape is expected: {err}"
13053 );
13054 let err = refusal(
13055 json!({"chat_template_kwargs": {"enable_thinking": "false"}}),
13056 "a stringly-typed switch must not be accepted",
13057 );
13058 assert!(
13059 err.contains("true or false"),
13060 "refusal must name the expected type: {err}"
13061 );
13062 let plan = build(json!({"chat_template_kwargs": null}))
13064 .expect("null chat_template_kwargs is the unset case");
13065 assert_eq!(plan.request.think, ThinkMode::Default);
13066 }
13067
13068 const Q38_TMPL: &str =
13083 include_str!("../../../research/reasoning-schema-20260823/qwen38-27b.chat_template.jinja");
13084
13085 fn render_with(
13088 tmpl: &str,
13089 caps: &ModelCaps,
13090 extra: serde_json::Value,
13091 default_effort: Option<&str>,
13092 ) -> Result<String, String> {
13093 let mut payload = serde_json::json!({
13094 "model": "m",
13095 "messages": [{"role": "user", "content": "hi"}],
13096 });
13097 if let Some(obj) = extra.as_object() {
13098 for (k, v) in obj {
13099 payload[k] = v.clone();
13100 }
13101 }
13102 let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
13103 let (tx, _rx) = worker::event_channel();
13104 let plan = build_chat_request_with_trace(
13105 req,
13106 Some(caps),
13107 tx,
13108 lanes::Lane::Interactive,
13109 None,
13110 None,
13111 default_effort,
13112 &ModelSamplingDefaults::default(),
13113 )?;
13114 Ok(chat::apply_chat_template_tools_ex(
13115 Some(tmpl),
13116 &plan.request.chat_turns,
13117 true,
13118 &plan.request.tools_json,
13119 &plan.request.tools_struct,
13120 plan.request.think,
13121 plan.request.reasoning_effort.as_deref(),
13122 None,
13123 )
13124 .unwrap())
13125 }
13126
13127 #[test]
13128 fn qwen38_effort_ladder_reaches_prompt_bytes_through_the_whole_api() {
13129 let r = |extra: serde_json::Value| render_with(Q38_TMPL, &ladder_caps(), extra, None);
13136 let xhigh = "Reasoning effort is set to xhigh.";
13137 let low = "Reasoning effort is set to low.";
13138 assert!(r(json!({"reasoning_effort": "low"})).unwrap().contains(low));
13140 assert!(
13141 r(json!({"reasoning_effort": "high"}))
13142 .unwrap()
13143 .contains(xhigh)
13144 );
13145 let medium = r(json!({"reasoning_effort": "medium"})).unwrap();
13148 assert!(!medium.contains("Reasoning effort is set to"), "{medium:?}");
13149 let low_p = r(json!({"reasoning_effort": "low"})).unwrap();
13151 let high_p = r(json!({"reasoning_effort": "high"})).unwrap();
13152 assert_ne!(low_p, high_p);
13153 assert_ne!(low_p, medium);
13154 assert_ne!(high_p, medium);
13155 for alias in ["xhigh", "max", "ultra"] {
13158 assert_eq!(r(json!({"reasoning_effort": alias})).unwrap(), high_p);
13159 }
13160 assert_eq!(r(json!({})).unwrap(), high_p);
13163 assert_eq!(
13166 render_with(Q38_TMPL, &ladder_caps(), json!({}), Some("medium")).unwrap(),
13167 medium
13168 );
13169 let off = r(json!({"reasoning_effort": "none"})).unwrap();
13172 assert!(off.ends_with("<think>\n\n</think>\n\n"), "{off:?}");
13173 assert!(!off.contains("Reasoning effort is set to"), "{off:?}");
13174 }
13175
13176 #[test]
13177 fn the_effort_sentence_is_measurable_on_the_deployed_binary_without_a_deploy() {
13178 const LOW_SENTENCE: &str = "Reasoning effort is set to low. Keep your thinking brief and \
13191focused, moving directly to the conclusion without unnecessary elaboration.";
13192 let expected = format!(
13193 "<|im_start|>system\n{LOW_SENTENCE}<|im_end|>\n\
13194 <|im_start|>user\nhi<|im_end|>\n<|im_start|>assistant\n<think>\n"
13195 );
13196 let after_fix = render_with(
13198 Q38_TMPL,
13199 &ladder_caps(),
13200 json!({"reasoning_effort": "low"}),
13201 None,
13202 )
13203 .unwrap();
13204 assert_eq!(
13205 after_fix, expected,
13206 "the shipped prompt for reasoning_effort:\"low\""
13207 );
13208 const ORNITH_TMPL: &str = include_str!(
13211 "../../../research/reasoning-schema-20260823/ornith15.chat_template.jinja"
13212 );
13213 let on_deployed_binary = render_with(
13214 ORNITH_TMPL,
13215 &tool_caps(),
13216 json!({"messages": [{"role": "system", "content": LOW_SENTENCE},
13217 {"role": "user", "content": "hi"}]}),
13218 None,
13219 )
13220 .unwrap();
13221 assert_eq!(
13222 on_deployed_binary, expected,
13223 "the live cell's system-message stand-in must render the SAME bytes as the post-fix \
13224 level, or its reasoning-volume numbers do not describe the shipped prompt"
13225 );
13226 let ladderless_unset = render_with(ORNITH_TMPL, &tool_caps(), json!({}), None).unwrap();
13229 assert!(
13230 !ladderless_unset.contains("Reasoning effort is set to"),
13231 "pre-lane q38 injected no effort instruction at any level: {ladderless_unset:?}"
13232 );
13233 assert_eq!(
13234 ladderless_unset,
13235 render_with(
13236 Q38_TMPL,
13237 &ladder_caps(),
13238 json!({"reasoning_effort": "medium"}),
13239 None
13240 )
13241 .unwrap(),
13242 "medium is the vendor's zero-steering rung and therefore the pre-lane byte baseline"
13243 );
13244 }
13245
13246 #[test]
13247 fn include_reasoning_false_stops_reasoning_it_does_not_hide_it() {
13248 let off = render_with(
13255 Q38_TMPL,
13256 &ladder_caps(),
13257 json!({"reasoning_effort": "none"}),
13258 None,
13259 )
13260 .unwrap();
13261 for extra in [
13262 json!({"include_reasoning": false}),
13263 json!({"reasoning": {"exclude": true}}),
13264 ] {
13265 let got = render_with(Q38_TMPL, &ladder_caps(), extra.clone(), None).unwrap();
13266 assert!(
13267 got.ends_with("<think>\n\n</think>\n\n"),
13268 "{extra} must render the CLOSED think pair, not a hidden reasoning block: {got:?}"
13269 );
13270 assert_eq!(got, off, "{extra} must be byte-identical to reasoning-off");
13271 }
13272 for extra in [
13277 json!({"enable_thinking": true, "include_reasoning": false}),
13278 json!({"reasoning": {"enabled": true}, "include_reasoning": false}),
13279 json!({"reasoning": {"enabled": true, "exclude": true}}),
13280 ] {
13281 let e = render_with(Q38_TMPL, &ladder_caps(), extra.clone(), None)
13282 .err()
13283 .unwrap_or_else(|| panic!("{extra} must be refused as contradictory"));
13284 assert!(e.contains("contradictory"), "{extra}: {e}");
13285 assert!(
13286 e.contains("include_reasoning") || e.contains("exclude"),
13287 "{extra}: the refusal must name the suppression field the caller sent: {e}"
13288 );
13289 }
13290 let dflt = render_with(Q38_TMPL, &ladder_caps(), json!({}), None).unwrap();
13293 for extra in [
13294 json!({"include_reasoning": true}),
13295 json!({"reasoning": {"exclude": false}}),
13296 ] {
13297 assert_eq!(
13298 render_with(Q38_TMPL, &ladder_caps(), extra.clone(), None).unwrap(),
13299 dflt,
13300 "{extra} must not perturb the model's default"
13301 );
13302 }
13303 let switchless = ModelCaps {
13307 think_switch: false,
13308 ..tool_caps()
13309 };
13310 let err = render_with(
13311 Q38_TMPL,
13312 &switchless,
13313 json!({"include_reasoning": false}),
13314 None,
13315 )
13316 .expect_err("include_reasoning:false must not silently bill for hidden reasoning");
13317 assert!(err.contains("cannot disable reasoning"), "{err}");
13318 }
13319
13320 #[test]
13321 fn the_reasoning_object_refuses_every_key_it_cannot_honour() {
13322 let build = |extra: serde_json::Value| {
13323 let (tx, _rx) = worker::event_channel();
13324 build_chat_request(
13325 weather_request(extra),
13326 Some(&ladder_caps()),
13327 tx,
13328 lanes::Lane::Interactive,
13329 None,
13330 )
13331 };
13332 let err = |extra: serde_json::Value, why: &str| -> String {
13333 build(extra).err().unwrap_or_else(|| panic!("{why}"))
13334 };
13335 let e = err(
13339 json!({"reasoning": {"max_tokens": 1024}}),
13340 "reasoning.max_tokens must not be accepted-and-ignored",
13341 );
13342 assert!(e.contains("reasoning.max_tokens"), "{e}");
13343 assert!(e.contains("ONE output budget"), "{e}");
13344 for extra in [
13349 json!({"reasoning": {"max_tokens": null}}),
13350 json!({"reasoning": {"banana": null}}),
13351 ] {
13352 let e = err(
13353 extra.clone(),
13354 "a null-valued unhonourable key must still refuse",
13355 );
13356 assert!(
13357 e.contains("max_tokens") || e.contains("banana"),
13358 "{extra}: {e}"
13359 );
13360 }
13361 let e = err(
13363 json!({"reasoning": {"budget": 5}}),
13364 "an unknown reasoning key must not be accepted",
13365 );
13366 assert!(
13367 e.contains("reasoning.budget") && e.contains("enabled"),
13368 "{e}"
13369 );
13370 for (extra, want) in [
13374 (json!({"reasoning": {"enabled": "false"}}), "true or false"),
13375 (json!({"reasoning": {"exclude": 1}}), "true or false"),
13376 (json!({"reasoning": {"effort": 3}}), "must be a string"),
13377 ] {
13378 let e = err(
13379 extra.clone(),
13380 "a wrong-typed reasoning key must not be ignored",
13381 );
13382 assert!(e.contains(want), "{extra}: {e}");
13383 }
13384 for extra in [
13389 json!({"reasoning": {"enabled": true}}),
13390 json!({"reasoning": {"effort": "low"}}),
13391 json!({"reasoning": {"exclude": false}}),
13392 json!({"reasoning": null}),
13393 json!({"reasoning": {"effort": null}}),
13394 json!({"reasoning": {"enabled": null, "exclude": null}}),
13395 ] {
13396 build(extra.clone()).unwrap_or_else(|e| panic!("{extra} must be served: {e}"));
13397 }
13398 }
13399
13400 #[test]
13401 fn a_graded_level_on_a_binary_model_translates_to_reasoning_on() {
13402 const ORNITH_TMPL: &str = include_str!(
13410 "../../../research/reasoning-schema-20260823/ornith15.chat_template.jinja"
13411 );
13412 let explicit_on = render_with(
13416 ORNITH_TMPL,
13417 &tool_caps(),
13418 json!({"reasoning": {"enabled": true}}),
13419 None,
13420 )
13421 .unwrap();
13422 assert!(explicit_on.ends_with("<think>\n"), "{explicit_on:?}");
13423 for extra in [
13424 json!({"reasoning_effort": "low"}),
13425 json!({"reasoning_effort": "medium"}),
13426 json!({"reasoning_effort": "high"}),
13427 json!({"reasoning_effort": "xhigh"}),
13429 json!({"reasoning": {"effort": "xhigh"}}),
13430 ] {
13431 let got = render_with(ORNITH_TMPL, &tool_caps(), extra.clone(), None)
13432 .unwrap_or_else(|e| panic!("{extra} must TRANSLATE to reasoning-on, got 400: {e}"));
13433 assert_eq!(
13434 got, explicit_on,
13435 "{extra} must render byte-identical to reasoning:{{enabled:true}} — the \
13436 documented translation, not a decorative accept"
13437 );
13438 }
13439 for extra in [
13441 json!({}),
13442 json!({"reasoning_effort": "none"}),
13443 json!({"reasoning_effort": "minimal"}),
13444 json!({"enable_thinking": false}),
13445 ] {
13446 render_with(ORNITH_TMPL, &tool_caps(), extra.clone(), None)
13447 .unwrap_or_else(|e| panic!("{extra} must still be served: {e}"));
13448 }
13449 let minimal = render_with(
13452 ORNITH_TMPL,
13453 &tool_caps(),
13454 json!({"reasoning_effort": "minimal"}),
13455 None,
13456 )
13457 .unwrap();
13458 assert!(
13459 minimal.ends_with("<think>\n\n</think>\n\n"),
13460 "minimal must close the think pair (OFF), not clamp to a reasoning level: {minimal:?}"
13461 );
13462 let ladder_low = render_with(
13465 Q38_TMPL,
13466 &ladder_caps(),
13467 json!({"reasoning_effort": "low"}),
13468 None,
13469 )
13470 .unwrap();
13471 assert!(
13472 ladder_low.contains("Reasoning effort is set to low."),
13473 "{ladder_low:?}"
13474 );
13475 assert_ne!(
13476 ladder_low,
13477 render_with(
13478 Q38_TMPL,
13479 &ladder_caps(),
13480 json!({"reasoning_effort": "high"}),
13481 None
13482 )
13483 .unwrap(),
13484 "the ladder model's rungs stay distinct prompts"
13485 );
13486 }
13487
13488 #[test]
13489 fn one_semantic_reasoning_request_renders_identical_bytes_on_all_three_surfaces() {
13490 let render_chat = |body: serde_json::Value| -> Result<String, String> {
13501 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
13502 let (tx, _rx) = worker::event_channel();
13503 let plan = build_chat_request(
13504 req,
13505 Some(&ladder_caps()),
13506 tx,
13507 lanes::Lane::Interactive,
13508 None,
13509 )?;
13510 Ok(chat::apply_chat_template_tools_ex(
13511 Some(Q38_TMPL),
13512 &plan.request.chat_turns,
13513 true,
13514 &plan.request.tools_json,
13515 &plan.request.tools_struct,
13516 plan.request.think,
13517 plan.request.reasoning_effort.as_deref(),
13518 None,
13519 )
13520 .unwrap())
13521 };
13522 for (intent, chat_body, responses_body, messages_body) in [
13527 (
13528 "reasoning OFF",
13529 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
13530 "reasoning_effort": "none"}),
13531 json!({"model": "m", "input": "hi", "reasoning": {"effort": "none"}}),
13532 json!({"model": "m", "max_tokens": 16,
13533 "messages": [{"role": "user", "content": "hi"}],
13534 "thinking": {"type": "disabled"}}),
13535 ),
13536 (
13537 "reasoning ON at the top rung",
13538 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
13539 "reasoning_effort": "xhigh"}),
13540 json!({"model": "m", "input": "hi", "reasoning": {"effort": "xhigh"}}),
13541 json!({"model": "m", "max_tokens": 16,
13542 "messages": [{"role": "user", "content": "hi"}],
13543 "output_config": {"effort": "xhigh"}}),
13544 ),
13545 (
13546 "reasoning ON at the bottom rung",
13547 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
13548 "reasoning_effort": "low"}),
13549 json!({"model": "m", "input": "hi", "reasoning": {"effort": "low"}}),
13550 json!({"model": "m", "max_tokens": 16,
13551 "messages": [{"role": "user", "content": "hi"}],
13552 "output_config": {"effort": "low"}}),
13553 ),
13554 (
13555 "the model's own default",
13556 json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}),
13557 json!({"model": "m", "input": "hi"}),
13558 json!({"model": "m", "max_tokens": 16,
13559 "messages": [{"role": "user", "content": "hi"}]}),
13560 ),
13561 ] {
13562 let chat = render_chat(chat_body).unwrap_or_else(|e| panic!("{intent} on chat: {e}"));
13563 let via_responses = responses_api::translate(&responses_body)
13564 .unwrap_or_else(|e| panic!("{intent} on /v1/responses: {e:?}"));
13565 let via_messages = anthropic::translate(&messages_body)
13566 .unwrap_or_else(|e| panic!("{intent} on /v1/messages: {e}"));
13567 for (surface, translated) in [
13568 ("/v1/responses", via_responses),
13569 ("/v1/messages", via_messages),
13570 ] {
13571 let got = render_chat(translated)
13572 .unwrap_or_else(|e| panic!("{intent} via {surface}: {e}"));
13573 assert_eq!(
13574 got, chat,
13575 "{intent}: {surface} rendered DIFFERENT prompt bytes than \
13576 /v1/chat/completions — the parameter is honoured on one format and not \
13577 the other"
13578 );
13579 }
13580 }
13581 let switchless = ModelCaps {
13584 think_switch: false,
13585 ..ladder_caps()
13586 };
13587 let render_switchless = |body: serde_json::Value| -> Result<String, String> {
13588 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
13589 let (tx, _rx) = worker::event_channel();
13590 let plan =
13591 build_chat_request(req, Some(&switchless), tx, lanes::Lane::Interactive, None)?;
13592 Ok(format!("{:?}", plan.request.think))
13593 };
13594 for (surface, body) in [
13595 (
13596 "/v1/responses",
13597 responses_api::translate(&json!({
13598 "model": "m", "input": "hi", "reasoning": {"effort": "none"}}))
13599 .unwrap(),
13600 ),
13601 (
13602 "/v1/messages",
13603 anthropic::translate(&json!({
13604 "model": "m", "max_tokens": 16,
13605 "messages": [{"role": "user", "content": "hi"}],
13606 "thinking": {"type": "disabled"}}))
13607 .unwrap(),
13608 ),
13609 ] {
13610 let err = render_switchless(body)
13611 .err()
13612 .unwrap_or_else(|| panic!("{surface} must refuse an unhonourable off-request"));
13613 assert!(err.contains("cannot disable reasoning"), "{surface}: {err}");
13614 }
13615 }
13616
13617 #[test]
13618 fn preserve_thinking_true_is_the_implemented_default_and_false_refuses() {
13619 let build = |extra: serde_json::Value| {
13627 let (tx, _rx) = worker::event_channel();
13628 build_chat_request(
13629 weather_request(extra),
13630 Some(&ladder_caps()),
13631 tx,
13632 lanes::Lane::Interactive,
13633 None,
13634 )
13635 };
13636 build(json!({"chat_template_kwargs": {"preserve_thinking": true}}))
13637 .expect("preserve_thinking:true is the vendor default the renderer implements");
13638 let e = build(json!({"chat_template_kwargs": {"preserve_thinking": false}}))
13639 .err()
13640 .expect("preserve_thinking:false (the strip arm) must refuse");
13641 assert!(e.contains("preserve_thinking"), "{e}");
13642 assert!(e.contains("strip"), "{e}");
13643 assert_eq!(
13646 build(json!({"chat_template_kwargs": {"enable_thinking": false}}))
13647 .unwrap()
13648 .request
13649 .think,
13650 ThinkMode::NoThink
13651 );
13652 let e = build(json!({"chat_template_kwargs": {"preserve_thinking": "false"}}))
13654 .err()
13655 .expect("a stringly-typed preserve_thinking must not be accepted");
13656 assert!(e.contains("true or false"), "{e}");
13657 }
13658
13659 #[test]
13660 fn dsv4_is_exempt_from_the_switchless_off_refusal() {
13661 let dsv4_caps = ModelCaps {
13667 qwen_think: true,
13668 think_switch: false,
13669 dsv4: true,
13670 ..tool_caps()
13671 };
13672 for extra in [
13673 json!({"reasoning_effort": "none"}),
13674 json!({"reasoning": {"enabled": false}}),
13675 json!({"enable_thinking": false}),
13676 json!({"include_reasoning": false}),
13677 ] {
13678 let (tx, _rx) = worker::event_channel();
13679 let plan = build_chat_request(
13680 weather_request(extra.clone()),
13681 Some(&dsv4_caps),
13682 tx,
13683 lanes::Lane::Interactive,
13684 None,
13685 )
13686 .unwrap_or_else(|e| panic!("{extra} must be served on dsv4: {e}"));
13687 assert_eq!(plan.request.think, ThinkMode::NoThink, "extra={extra}");
13688 }
13689 }
13690
13691 #[test]
13692 fn contradictory_think_switches_refuse_instead_of_picking_one() {
13693 let build = |extra: serde_json::Value| {
13696 let (tx, _rx) = worker::event_channel();
13697 build_chat_request(
13698 weather_request(extra),
13699 Some(&tool_caps()),
13700 tx,
13701 lanes::Lane::Interactive,
13702 None,
13703 )
13704 };
13705 for extra in [
13706 json!({"enable_thinking": true, "reasoning": {"enabled": false}}),
13707 json!({"enable_thinking": false, "reasoning": {"enabled": true}}),
13708 json!({"enable_thinking": false, "chat_template_kwargs": {"enable_thinking": true}}),
13709 ] {
13710 match build(extra.clone()) {
13711 Err(err) => assert!(
13712 err.contains("contradictory"),
13713 "the refusal must say the switches contradict: {err}"
13714 ),
13715 Ok(plan) => panic!(
13716 "{extra} must be rejected as contradictory; it silently resolved to {:?}",
13717 plan.request.think
13718 ),
13719 }
13720 }
13721 for extra in [
13723 json!({"enable_thinking": false, "reasoning": {"enabled": false}}),
13724 json!({"enable_thinking": true, "reasoning": {"enabled": true}}),
13725 json!({"enable_thinking": false, "reasoning": {"effort": "high"}}),
13726 ] {
13727 build(extra.clone())
13728 .unwrap_or_else(|e| panic!("{extra} is not a contradiction, but got 400: {e}"));
13729 }
13730 }
13731
13732 #[test]
13733 fn explicit_reasoning_off_on_a_switchless_template_refuses_loudly() {
13734 let switchless = ModelCaps {
13739 tools_branch: true,
13740 qwen_think: true,
13741 think_switch: false,
13742 chat_ok: true,
13743 ..Default::default()
13744 };
13745 let build = |extra: serde_json::Value, caps: &ModelCaps, default_effort: Option<&str>| {
13746 let (tx, _rx) = worker::event_channel();
13747 build_chat_request_with_trace(
13748 weather_request(extra),
13749 Some(caps),
13750 tx,
13751 lanes::Lane::Interactive,
13752 None,
13753 None,
13754 default_effort,
13755 &ModelSamplingDefaults::default(),
13756 )
13757 };
13758 for extra in [
13759 json!({"reasoning_effort": "none"}),
13760 json!({"reasoning_effort": "minimal"}),
13761 json!({"reasoning": {"enabled": false}}),
13762 json!({"enable_thinking": false}),
13763 json!({"chat_template_kwargs": {"enable_thinking": false}}),
13764 ] {
13765 let err = build(extra.clone(), &switchless, None)
13766 .err()
13767 .unwrap_or_else(|| {
13768 panic!(
13769 "{extra} on a switchless think template must not be accepted-and-ignored"
13770 )
13771 });
13772 assert!(
13773 err.contains("cannot disable reasoning"),
13774 "the refusal must say the model cannot disable reasoning: {err}"
13775 );
13776 }
13777 for (extra, default_effort) in [
13781 (json!({}), None),
13782 (json!({"reasoning_effort": "high"}), None),
13785 (json!({"reasoning": {"enabled": true}}), None),
13786 (json!({"enable_thinking": true}), None),
13787 (json!({}), Some("none")),
13788 (json!({}), Some("minimal")),
13789 (json!({}), Some("high")),
13790 ] {
13791 build(extra.clone(), &switchless, default_effort).unwrap_or_else(|e| {
13792 panic!("{extra} (default={default_effort:?}) must still be served: {e}")
13793 });
13794 }
13795 assert_eq!(
13798 build(json!({"enable_thinking": false}), &tool_caps(), None)
13799 .unwrap()
13800 .request
13801 .think,
13802 ThinkMode::NoThink
13803 );
13804 }
13805
13806 #[test]
13807 fn gemma4_default_think_on_renders_byte_identical_to_explicit_think_on() {
13808 let gemma_caps = ModelCaps {
13815 tools_branch: true,
13816 chat_ok: true,
13817 gemma_think: true,
13818 instruct_type: Some("gemma".into()),
13819 ..Default::default()
13820 };
13821 let render =
13822 |tmpl: &str, extra: serde_json::Value, default_effort: Option<&str>| -> String {
13823 let mut payload = serde_json::json!({
13824 "model": "google/gemma-4-31b-it",
13825 "messages": [{"role": "user", "content": "Weather in Paris?"}],
13826 });
13827 if let Some(obj) = extra.as_object() {
13828 for (k, v) in obj {
13829 payload[k] = v.clone();
13830 }
13831 }
13832 let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
13833 let (tx, _rx) = worker::event_channel();
13834 let plan = build_chat_request_with_trace(
13835 req,
13836 Some(&gemma_caps),
13837 tx,
13838 lanes::Lane::Interactive,
13839 None,
13840 None,
13841 default_effort,
13842 &ModelSamplingDefaults::default(),
13843 )
13844 .unwrap();
13845 chat::apply_chat_template_tools_ex(
13846 Some(tmpl),
13847 &plan.request.chat_turns,
13848 true,
13849 &plan.request.tools_json,
13850 &plan.request.tools_struct,
13851 plan.request.think,
13852 plan.request.reasoning_effort.as_deref(),
13853 None, )
13855 .unwrap()
13856 };
13857 let official = gemma_template("official");
13858 let unset_with_knob = render(&official, json!({}), Some("high"));
13859 let explicit_on = render(&official, json!({"reasoning_effort": "high"}), None);
13860 assert_eq!(
13861 unset_with_knob, explicit_on,
13862 "knob render must be byte-identical to the explicit think-on render"
13863 );
13864 assert!(
13865 unset_with_knob.starts_with("<|turn>system\n<|think|>\n"),
13866 "think-on injects the <|think|> system token: {unset_with_knob:?}"
13867 );
13868 assert!(
13869 unset_with_knob.ends_with("<|turn>model\n"),
13870 "think-on generation turn is OPEN: {unset_with_knob:?}"
13871 );
13872 let explicit_off_with_knob =
13876 render(&official, json!({"reasoning_effort": "none"}), Some("high"));
13877 let explicit_off = render(&official, json!({"reasoning_effort": "none"}), None);
13878 assert_eq!(explicit_off_with_knob, explicit_off);
13879 assert!(
13880 !explicit_off_with_knob.contains("<|think|>")
13881 && explicit_off_with_knob.ends_with("<|turn>model\n"),
13882 "explicit off keeps the official template's thinking-off bytes: \
13883 {explicit_off_with_knob:?}"
13884 );
13885 let unset_no_knob = render(&official, json!({}), None);
13887 assert_eq!(
13888 unset_no_knob, explicit_off,
13889 "knobless unset stays the template's own thinking-off default"
13890 );
13891 assert_ne!(unset_no_knob, unset_with_knob);
13892 let qat = gemma_template("qat");
13895 assert!(
13896 render(&qat, json!({}), None).ends_with("<|turn>model\n<|channel>thought\n<channel|>"),
13897 "QAT knobless unset keeps the closed-channel default"
13898 );
13899 assert_eq!(
13900 render(&qat, json!({}), Some("high")),
13901 render(&qat, json!({"reasoning_effort": "high"}), None),
13902 "QAT knob render must equal the explicit think-on render"
13903 );
13904 }
13905
13906 #[test]
13907 fn default_reasoning_effort_is_validated_at_metadata_load() {
13908 let parsed = OpenRouterMetadataFile::from_toml(
13910 r#"
13911[models.g]
13912default_reasoning_effort = "high"
13913"#,
13914 )
13915 .unwrap();
13916 assert_eq!(
13917 parsed.get("g").unwrap().default_reasoning_effort.as_deref(),
13918 Some("high")
13919 );
13920 let err = OpenRouterMetadataFile::from_toml(
13921 r#"
13922[models.g]
13923default_reasoning_effort = "always"
13924"#,
13925 )
13926 .unwrap_err();
13927 assert!(err.contains("default_reasoning_effort"), "{err}");
13928 }
13929
13930 #[test]
13931 fn reasoning_effort_maps_to_effort_level_on_step35_class_templates() {
13932 let effort_caps = ModelCaps {
13944 effort_levels: true,
13945 think_switch: false,
13946 ..tool_caps()
13947 };
13948 for (extra, want) in [
13949 (json!({}), None),
13950 (json!({"reasoning_effort": "low"}), Some("low")),
13951 (json!({"reasoning_effort": "medium"}), Some("medium")),
13952 (json!({"reasoning_effort": "high"}), Some("high")),
13953 (json!({"reasoning": {"effort": "high"}}), Some("high")),
13954 (json!({"reasoning_effort": "xhigh"}), Some("high")),
13956 (json!({"reasoning": {"effort": "max"}}), Some("high")),
13957 ] {
13958 let (tx, _rx) = worker::event_channel();
13959 let plan = build_chat_request(
13960 weather_request(extra.clone()),
13961 Some(&effort_caps),
13962 tx,
13963 lanes::Lane::Interactive,
13964 None,
13965 )
13966 .unwrap();
13967 assert_eq!(
13968 plan.request.reasoning_effort.as_deref(),
13969 want,
13970 "extra={extra}"
13971 );
13972 }
13973 for extra in [
13980 json!({"reasoning_effort": "none"}),
13981 json!({"reasoning_effort": "minimal"}),
13982 json!({"reasoning": {"enabled": false}}),
13983 json!({"enable_thinking": false}),
13984 json!({"include_reasoning": false}),
13985 ] {
13986 let (tx, _rx) = worker::event_channel();
13987 let err = build_chat_request(
13988 weather_request(extra.clone()),
13989 Some(&effort_caps),
13990 tx,
13991 lanes::Lane::Interactive,
13992 None,
13993 )
13994 .err()
13995 .unwrap_or_else(|| panic!("{extra} must not be clamped to a reasoning level"));
13996 assert!(
13997 err.contains("cannot disable reasoning"),
13998 "extra={extra}: {err}"
13999 );
14000 }
14001 for extra in [
14008 json!({"reasoning_effort": "high"}),
14009 json!({"reasoning": {"effort": "low"}}),
14010 ] {
14011 let (tx, _rx) = worker::event_channel();
14012 let plan = build_chat_request(
14013 weather_request(extra.clone()),
14014 Some(&tool_caps()),
14015 tx,
14016 lanes::Lane::Interactive,
14017 None,
14018 )
14019 .unwrap_or_else(|e| panic!("{extra} must translate, not refuse: {e}"));
14020 assert_eq!(plan.request.think, ThinkMode::Think, "extra={extra}");
14021 assert_eq!(plan.request.reasoning_effort, None, "extra={extra}");
14022 }
14023 let (tx, _rx) = worker::event_channel();
14025 let plan = build_chat_request(
14026 weather_request(json!({})),
14027 Some(&tool_caps()),
14028 tx,
14029 lanes::Lane::Interactive,
14030 None,
14031 )
14032 .unwrap();
14033 assert_eq!(plan.request.reasoning_effort, None);
14034 }
14035
14036 #[test]
14037 fn assistant_history_tool_calls_and_tool_role_render_into_turns() {
14038 let payload = serde_json::json!({
14039 "model": "m",
14040 "messages": [
14041 {"role": "user", "content": "Weather in Paris?"},
14042 {"role": "assistant", "content": null, "tool_calls": [
14043 {"id": "call_x", "type": "function", "function": {
14044 "name": "get_weather",
14045 "arguments": "{\"city\": \"Paris\", \"days\": 3}"}}]},
14046 {"role": "tool", "tool_call_id": "call_x", "content": "{\"temp_c\": 21}"}
14047 ],
14048 });
14049 let req: ChatCompletionReq = serde_json::from_value(payload).unwrap();
14050 let (tx, _rx) = worker::event_channel();
14051 let plan = build_chat_request(req, Some(&tool_caps()), tx, lanes::Lane::Interactive, None)
14052 .unwrap();
14053 let turns = &plan.request.chat_turns;
14054 assert_eq!(turns[1].tool_calls.len(), 1);
14055 assert_eq!(turns[1].tool_calls[0].name, "get_weather");
14056 assert_eq!(
14057 turns[1].tool_calls[0].params,
14058 vec![("city".into(), "Paris".into()), ("days".into(), "3".into())]
14059 );
14060 assert_eq!(turns[2].role, "tool");
14061 assert_eq!(turns[2].content, "{\"temp_c\": 21}");
14062 let mut p = plan
14065 .parser
14066 .expect("think-open chat arms the reasoning splitter");
14067 let pieces = p.push("thought</think>\n\nanswer <tool_call> is prose here");
14068 assert_eq!(
14069 pieces,
14070 vec![
14071 Piece::Reasoning("thought".into()),
14072 Piece::Content("answer <tool_call> is prose here".into()),
14073 ]
14074 );
14075 }
14076
14077 #[tokio::test]
14078 async fn blocking_tools_response_carries_tool_calls_and_finish_reason() {
14079 let (tx, rx) = worker::event_channel();
14080 tx.send(Event::Token {
14081 id: 1,
14082 text: "plan</think>\n\n".into(),
14083 })
14084 .unwrap();
14085 tx.send(Event::Token {
14086 id: 2,
14087 text: "<tool_call>\n<function=get_weather>\n\
14088<parameter=city>\nParis\n</parameter>\n</function>\n</tool_call>"
14089 .into(),
14090 })
14091 .unwrap();
14092 tx.send(Event::Done {
14093 stop_reason: "Eos".into(),
14094 n_tokens: 2,
14095 n_prompt: 40,
14096 n_cached: 0,
14097 elapsed_s: 0.5,
14098 spec: None,
14099 })
14100 .unwrap();
14101 drop(tx);
14102 let parser = ToolStreamParser::new(HashMap::new(), true);
14103 let response = blocking_response(
14104 rx,
14105 "m".into(),
14106 true,
14107 Vec::new(),
14108 Some(parser),
14109 Envelope::new(true),
14110 )
14111 .await;
14112 assert_eq!(response.status(), StatusCode::OK);
14113 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
14114 .await
14115 .unwrap();
14116 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
14117 assert_eq!(payload["choices"][0]["finish_reason"], "tool_calls");
14118 assert_eq!(
14121 payload["choices"][0]["message"]["content"],
14122 serde_json::Value::Null
14123 );
14124 assert_eq!(payload["choices"][0]["message"]["reasoning"], "plan");
14125 assert_eq!(
14126 payload["choices"][0]["message"]["reasoning_details"][0]["text"],
14127 "plan"
14128 );
14129 let call = &payload["choices"][0]["message"]["tool_calls"][0];
14130 assert_eq!(call["type"], "function");
14131 assert_eq!(call["function"]["name"], "get_weather");
14132 assert_eq!(call["function"]["arguments"], "{\"city\":\"Paris\"}");
14133 assert_eq!(payload["usage"]["prompt_tokens"], 40);
14136 assert_eq!(payload["usage"]["completion_tokens"], 2);
14137 assert_eq!(payload["usage"]["total_tokens"], 42);
14138 assert_eq!(
14139 payload["usage"]["prompt_tokens_details"]["cached_tokens"],
14140 0
14141 );
14142 }
14143
14144 #[test]
14145 fn cache_salt_plumbs_to_the_worker_namespace() {
14146 let req: CompletionReq = serde_json::from_value(serde_json::json!({
14148 "model": "m", "prompt": "task", "cache_salt": "tenant-a"
14149 }))
14150 .unwrap();
14151 let (tx, _rx) = worker::event_channel();
14152 assert_eq!(
14153 build_request(&req, tx, lanes::Lane::Interactive, None).cache_ns,
14154 "tenant-a"
14155 );
14156
14157 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
14158 "model": "m", "messages": [{"role": "user", "content": "task"}],
14159 "cache_salt": "tenant-b"
14160 }))
14161 .unwrap();
14162 let (tx, _rx) = worker::event_channel();
14163 assert_eq!(
14164 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
14165 .unwrap()
14166 .request
14167 .cache_ns,
14168 "tenant-b"
14169 );
14170
14171 let req: CompletionReq = serde_json::from_value(serde_json::json!({
14173 "model": "m", "prompt": "task"
14174 }))
14175 .unwrap();
14176 let (tx, _rx) = worker::event_channel();
14177 assert_eq!(
14178 build_request(&req, tx, lanes::Lane::Interactive, None).cache_ns,
14179 ""
14180 );
14181 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
14182 "model": "m", "messages": [{"role": "user", "content": "task"}]
14183 }))
14184 .unwrap();
14185 let (tx, _rx) = worker::event_channel();
14186 assert_eq!(
14187 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
14188 .unwrap()
14189 .request
14190 .cache_ns,
14191 ""
14192 );
14193 }
14194
14195 #[test]
14196 fn cache_salt_validation_rejects_oversized_value() {
14197 let salt = Some("a".repeat(CACHE_SALT_MAX_BYTES + 1));
14198 assert_eq!(
14199 validate_cache_namespace(&salt, false),
14200 Err("cache_salt must be at most 64 bytes")
14201 );
14202 }
14203
14204 #[test]
14205 fn cache_salt_validation_rejects_reserved_open_namespace() {
14206 let salt = Some("t:acme\u{1f}private".to_string());
14207 assert_eq!(
14208 validate_cache_namespace(&salt, false),
14209 Err("cache_salt must not use the reserved t: prefix without a keyring")
14210 );
14211 }
14212
14213 #[test]
14214 fn cache_salt_validation_accepts_normal_value() {
14215 let raw = "tenant-A_7.c2VjcmV0LXNjb3Bl+/=";
14216 let salt = Some(raw.to_string());
14217 assert_eq!(validate_cache_namespace(&salt, false).unwrap(), raw);
14218 assert_eq!(validate_cache_namespace(&None, false).unwrap(), "");
14219 let max_raw = "a".repeat(CACHE_SALT_MAX_BYTES);
14220 let max = Some(max_raw.clone());
14221 assert_eq!(validate_cache_namespace(&max, false).unwrap(), max_raw);
14222 }
14223
14224 #[test]
14225 fn cache_salt_validation_rejects_unsupported_characters() {
14226 let salt = Some("tenant salt".to_string());
14227 assert_eq!(
14228 validate_cache_namespace(&salt, false),
14229 Err("cache_salt contains unsupported characters")
14230 );
14231 }
14232
14233 #[test]
14234 fn affinity_key_honors_both_client_conventions_in_priority_order() {
14235 use axum::http::HeaderMap;
14236 let hdr = |v: &str| {
14237 let mut h = HeaderMap::new();
14238 h.insert("x-session-id", v.parse().unwrap());
14239 h
14240 };
14241 let empty = HeaderMap::new();
14242 let s = |v: &str| Some(v.to_string());
14243 assert_eq!(
14245 affinity_key(&s("explicit"), &None, &empty).unwrap(),
14246 s("explicit")
14247 );
14248 assert_eq!(
14249 affinity_key(&None, &s("openai-user"), &empty).unwrap(),
14250 s("openai-user")
14251 );
14252 assert_eq!(
14253 affinity_key(&None, &None, &hdr("hdr-id")).unwrap(),
14254 s("hdr-id")
14255 );
14256 assert_eq!(affinity_key(&s("a"), &s("b"), &hdr("c")).unwrap(), s("a"));
14259 assert_eq!(affinity_key(&None, &s("b"), &hdr("c")).unwrap(), s("b"));
14260 assert_eq!(affinity_key(&s(" "), &s(""), &hdr(" ")).unwrap(), None);
14263 assert_eq!(affinity_key(&s(""), &s("real"), &empty).unwrap(), s("real"));
14264 assert_eq!(
14266 affinity_key(&s(" padded "), &None, &empty).unwrap(),
14267 s("padded")
14268 );
14269 assert_eq!(affinity_key(&None, &None, &empty).unwrap(), None);
14271 assert!(
14272 affinity_key(
14273 &s(&"x".repeat(MAX_CLIENT_IDENTIFIER_BYTES + 1)),
14274 &None,
14275 &empty,
14276 )
14277 .unwrap_err()
14278 .contains("at most")
14279 );
14280 assert!(
14281 affinity_key(&s("forged\nlog"), &None, &empty)
14282 .unwrap_err()
14283 .contains("control")
14284 );
14285 }
14286
14287 #[test]
14288 fn affinity_key_plumbs_to_the_worker_request_on_both_bodies() {
14289 let req: CompletionReq = serde_json::from_value(serde_json::json!({
14290 "model": "m", "prompt": "task", "session_id": "conv-1"
14291 }))
14292 .unwrap();
14293 let (tx, _rx) = worker::event_channel();
14294 let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new()).unwrap();
14295 assert_eq!(
14296 build_request(&req, tx, lanes::Lane::Interactive, key)
14297 .affinity
14298 .as_deref(),
14299 Some("conv-1")
14300 );
14301 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
14303 "model": "m", "messages": [{"role": "user", "content": "task"}],
14304 "user": "conv-2"
14305 }))
14306 .unwrap();
14307 let (tx, _rx) = worker::event_channel();
14308 let key = affinity_key(&req.session_id, &req.user, &axum::http::HeaderMap::new()).unwrap();
14309 assert_eq!(
14310 build_chat_request(req, None, tx, lanes::Lane::Interactive, key)
14311 .unwrap()
14312 .request
14313 .affinity
14314 .as_deref(),
14315 Some("conv-2")
14316 );
14317 let req: CompletionReq = serde_json::from_value(serde_json::json!({
14319 "model": "m", "prompt": "task"
14320 }))
14321 .unwrap();
14322 let (tx, _rx) = worker::event_channel();
14323 assert!(
14324 build_request(&req, tx, lanes::Lane::Interactive, None)
14325 .affinity
14326 .is_none()
14327 );
14328 }
14329
14330 async fn sse_data_lines(resp: Response) -> Vec<String> {
14332 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
14333 .await
14334 .unwrap();
14335 String::from_utf8(bytes.to_vec())
14336 .unwrap()
14337 .lines()
14338 .filter_map(|l| l.strip_prefix("data: ").map(str::to_string))
14339 .collect()
14340 }
14341
14342 #[tokio::test]
14343 async fn chat_returns_reasoning_text_when_on_and_no_field_when_off() {
14344 let feed = |think: bool| {
14350 let (tx, rx) = worker::event_channel();
14351 let body = if think {
14352 "a plan</think>\n\nanswer"
14353 } else {
14354 "answer"
14355 };
14356 tx.send(Event::Token {
14357 id: 1,
14358 text: body.into(),
14359 })
14360 .unwrap();
14361 tx.send(Event::Done {
14362 stop_reason: "Eos".into(),
14363 n_tokens: 3,
14364 n_prompt: 10,
14365 n_cached: 0,
14366 elapsed_s: 0.1,
14367 spec: None,
14368 })
14369 .unwrap();
14370 drop(tx);
14371 rx
14372 };
14373 let resp = blocking_response(
14375 feed(true),
14376 "m".into(),
14377 true,
14378 Vec::new(),
14379 Some(ToolStreamParser::reasoning_only()),
14380 Envelope::new(true),
14381 )
14382 .await;
14383 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
14384 .await
14385 .unwrap();
14386 let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
14387 assert_eq!(v["choices"][0]["message"]["reasoning"], "a plan");
14388 assert_eq!(
14389 v["choices"][0]["message"]["reasoning_details"][0]["text"],
14390 "a plan"
14391 );
14392 assert_eq!(v["choices"][0]["message"]["content"], "answer");
14393 let resp = blocking_response(
14396 feed(false),
14397 "m".into(),
14398 true,
14399 Vec::new(),
14400 None,
14401 Envelope::new(true),
14402 )
14403 .await;
14404 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
14405 .await
14406 .unwrap();
14407 let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
14408 assert!(
14409 v["choices"][0]["message"].get("reasoning").is_none(),
14410 "a reasoning-off response must carry no reasoning field: {v}"
14411 );
14412 assert_eq!(v["choices"][0]["message"]["content"], "answer");
14413 let resp = sse_response(
14415 feed(true),
14416 "m".into(),
14417 true,
14418 Some(ToolStreamParser::reasoning_only()),
14419 Envelope::new(true),
14420 Vec::new(),
14421 None,
14422 )
14423 .into_response();
14424 let lines = sse_data_lines(resp).await;
14425 let chunks: Vec<serde_json::Value> = lines[..lines.len() - 1]
14426 .iter()
14427 .map(|l| serde_json::from_str(l).unwrap())
14428 .collect();
14429 let reasoning: String = chunks
14430 .iter()
14431 .filter_map(|c| c["choices"][0]["delta"]["reasoning"].as_str())
14432 .collect();
14433 assert_eq!(
14434 reasoning, "a plan",
14435 "think text must stream as delta.reasoning"
14436 );
14437 let content: String = chunks
14438 .iter()
14439 .filter_map(|c| c["choices"][0]["delta"]["content"].as_str())
14440 .collect();
14441 assert_eq!(content, "answer", "content must exclude the think segment");
14442 let resp = sse_response(
14444 feed(false),
14445 "m".into(),
14446 true,
14447 None,
14448 Envelope::new(true),
14449 Vec::new(),
14450 None,
14451 )
14452 .into_response();
14453 let lines = sse_data_lines(resp).await;
14454 for l in &lines[..lines.len() - 1] {
14455 let c: serde_json::Value = serde_json::from_str(l).unwrap();
14456 assert!(
14457 c["choices"][0]["delta"].get("reasoning").is_none(),
14458 "a reasoning-off stream must carry no reasoning deltas: {c}"
14459 );
14460 }
14461 }
14462
14463 #[tokio::test]
14464 async fn stream_chunks_carry_envelope_and_first_delta_role() {
14465 let (tx, rx) = worker::event_channel();
14466 tx.send(Event::Token {
14467 id: 1,
14468 text: "he".into(),
14469 })
14470 .unwrap();
14471 tx.send(Event::Token {
14472 id: 2,
14473 text: "llo".into(),
14474 })
14475 .unwrap();
14476 tx.send(Event::Done {
14477 stop_reason: "Eos".into(),
14478 n_tokens: 2,
14479 n_prompt: 10,
14480 n_cached: 0,
14481 elapsed_s: 0.1,
14482 spec: None,
14483 })
14484 .unwrap();
14485 drop(tx);
14486 let resp = sse_response(
14487 rx,
14488 "m".into(),
14489 true,
14490 None,
14491 Envelope::new(true),
14492 Vec::new(),
14493 None,
14494 )
14495 .into_response();
14496 let lines = sse_data_lines(resp).await;
14497 assert_eq!(lines.last().map(String::as_str), Some("[DONE]"));
14498 let chunks: Vec<serde_json::Value> = lines[..lines.len() - 1]
14499 .iter()
14500 .map(|l| serde_json::from_str(l).unwrap())
14501 .collect();
14502 let id = chunks[0]["id"].as_str().unwrap().to_string();
14504 assert!(id.starts_with("chatcmpl-"));
14505 for c in &chunks {
14506 assert_eq!(c["id"], id.as_str());
14507 assert!(c["created"].as_u64().unwrap() > 1_700_000_000);
14508 let fingerprint = c["system_fingerprint"].as_str().unwrap();
14509 assert!(
14510 build_id::fingerprint_is_well_formed(fingerprint),
14511 "chunk system_fingerprint {fingerprint:?} is not memra-<version>-<12 hex>"
14512 );
14513 assert_eq!(c["object"], "chat.completion.chunk");
14514 }
14515 assert_eq!(chunks[0]["choices"][0]["delta"]["role"], "assistant");
14517 assert_eq!(chunks[0]["choices"][0]["delta"]["content"], "he");
14518 assert!(chunks[1]["choices"][0]["delta"].get("role").is_none());
14519 let fin = chunks.last().unwrap();
14521 assert_eq!(fin["choices"][0]["finish_reason"], "stop");
14522 assert_eq!(fin["usage"]["prompt_tokens"], 10);
14523 }
14524
14525 #[tokio::test]
14526 async fn stream_token_events_equal_usage_on_every_finish_path() {
14527 for (stop_reason, expected_finish) in [
14528 ("Eos", "stop"),
14529 ("Callback", "stop"),
14530 ("MaxNew", "length"),
14531 ("ContextFull", "length"),
14532 ] {
14533 let (tx, rx) = worker::event_channel();
14534 tx.send(Event::Token {
14537 id: 248_046,
14538 text: String::new(),
14539 })
14540 .unwrap();
14541 tx.send(Event::Done {
14542 stop_reason: stop_reason.into(),
14543 n_tokens: 1,
14544 n_prompt: 8,
14545 n_cached: 8,
14546 elapsed_s: 0.1,
14547 spec: None,
14548 })
14549 .unwrap();
14550 drop(tx);
14551
14552 let resp = sse_response(
14553 rx,
14554 "m".into(),
14555 true,
14556 None,
14557 Envelope::new(true),
14558 Vec::new(),
14559 None,
14560 )
14561 .into_response();
14562 let lines = sse_data_lines(resp).await;
14563 assert_eq!(lines.last().map(String::as_str), Some("[DONE]"));
14564 let chunks: Vec<serde_json::Value> = lines[..lines.len() - 1]
14565 .iter()
14566 .map(|line| serde_json::from_str(line).unwrap())
14567 .collect();
14568 let token_events = chunks
14569 .iter()
14570 .filter(|chunk| chunk["choices"][0]["finish_reason"].is_null())
14571 .count();
14572 let terminal = chunks.last().unwrap();
14573 assert_eq!(token_events, 1, "{stop_reason} SSE token count");
14574 assert_eq!(terminal["usage"]["completion_tokens"], token_events);
14575 assert_eq!(terminal["choices"][0]["finish_reason"], expected_finish);
14576 }
14577 }
14578
14579 #[tokio::test]
14580 async fn stream_excludes_stop_text_like_non_stream_does() {
14581 let (tx, rx) = worker::event_channel();
14585 tx.send(Event::Token {
14586 id: 1,
14587 text: "answer\nPro".into(),
14588 })
14589 .unwrap();
14590 tx.send(Event::Token {
14591 id: 2,
14592 text: "blem: leaked prompt".into(),
14593 })
14594 .unwrap();
14595 tx.send(Event::Done {
14596 stop_reason: "Callback".into(),
14597 n_tokens: 2,
14598 n_prompt: 8,
14599 n_cached: 0,
14600 elapsed_s: 0.1,
14601 spec: None,
14602 })
14603 .unwrap();
14604 drop(tx);
14605 let resp = sse_response(
14606 rx,
14607 "m".into(),
14608 true,
14609 None,
14610 Envelope::new(true),
14611 vec!["Problem:".into()],
14612 None,
14613 )
14614 .into_response();
14615 let lines = sse_data_lines(resp).await;
14616 let content: String = lines
14617 .iter()
14618 .filter(|l| *l != "[DONE]")
14619 .filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
14620 .filter_map(|c| {
14621 c["choices"][0]["delta"]["content"]
14622 .as_str()
14623 .map(str::to_string)
14624 })
14625 .collect();
14626 assert_eq!(content, "answer\n");
14627
14628 let (tx, rx) = worker::event_channel();
14630 tx.send(Event::Token {
14631 id: 1,
14632 text: "ends in Pro".into(),
14633 })
14634 .unwrap();
14635 tx.send(Event::Done {
14636 stop_reason: "Eos".into(),
14637 n_tokens: 1,
14638 n_prompt: 8,
14639 n_cached: 0,
14640 elapsed_s: 0.1,
14641 spec: None,
14642 })
14643 .unwrap();
14644 drop(tx);
14645 let resp = sse_response(
14646 rx,
14647 "m".into(),
14648 true,
14649 None,
14650 Envelope::new(true),
14651 vec!["Problem:".into()],
14652 None,
14653 )
14654 .into_response();
14655 let lines = sse_data_lines(resp).await;
14656 let content: String = lines
14657 .iter()
14658 .filter(|l| *l != "[DONE]")
14659 .filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
14660 .filter_map(|c| {
14661 c["choices"][0]["delta"]["content"]
14662 .as_str()
14663 .map(str::to_string)
14664 })
14665 .collect();
14666 assert_eq!(content, "ends in Pro");
14667 }
14668
14669 #[tokio::test]
14670 async fn stream_worker_error_is_a_data_chunk_not_a_named_event() {
14671 let (tx, rx) = worker::event_channel();
14672 tx.send(Event::Error(worker::EngineError::engine("boom")))
14673 .unwrap();
14674 drop(tx);
14675 let resp = sse_response(
14676 rx,
14677 "m".into(),
14678 true,
14679 None,
14680 Envelope::new(true),
14681 Vec::new(),
14682 None,
14683 )
14684 .into_response();
14685 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
14686 .await
14687 .unwrap();
14688 let body = String::from_utf8(bytes.to_vec()).unwrap();
14689 assert!(
14691 !body.contains("event: error"),
14692 "named SSE event leaked: {body}"
14693 );
14694 let lines: Vec<&str> = body
14695 .lines()
14696 .filter_map(|l| l.strip_prefix("data: "))
14697 .collect();
14698 let err: serde_json::Value = serde_json::from_str(lines[0]).unwrap();
14699 assert_eq!(err["error"]["message"], "boom");
14700 assert_eq!(err["error"]["type"], "server_error");
14701 assert_eq!(err["error"]["code"], "engine_error");
14702 assert_eq!(lines.last(), Some(&"[DONE]"));
14703 }
14704
14705 #[test]
14706 fn ttft_sse_marker_ignores_keepalive_comments() {
14707 assert!(!is_sse_data_frame(b": keep-alive\n\n"));
14708 assert!(is_sse_data_frame(b"data: {\"choices\":[]}\n\n"));
14709 assert!(is_sse_data_frame(
14710 b"event: error\ndata: {\"error\":\"failed\"}\n\n"
14711 ));
14712 }
14713
14714 #[tokio::test]
14715 async fn error_bodies_use_the_openai_object_shape() {
14716 let (tx, rx) = worker::event_channel();
14717 tx.send(Event::Error(worker::EngineError::model_not_found(
14718 "unknown model \"x\"",
14719 )))
14720 .unwrap();
14721 drop(tx);
14722 let response =
14723 blocking_response(rx, "m".into(), true, Vec::new(), None, Envelope::new(true)).await;
14724 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
14725 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
14726 .await
14727 .unwrap();
14728 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
14729 assert_eq!(payload["error"]["message"], "unknown model \"x\"");
14731 assert_eq!(payload["error"]["type"], "invalid_request_error");
14732 assert_eq!(payload["error"]["param"], "model");
14733 assert_eq!(payload["error"]["code"], "model_not_found");
14734 }
14735
14736 fn retry_after(resp: &Response) -> Option<String> {
14743 resp.headers()
14744 .get(axum::http::header::RETRY_AFTER)
14745 .and_then(|v| v.to_str().ok())
14746 .map(str::to_string)
14747 }
14748
14749 async fn body_value(resp: Response) -> serde_json::Value {
14752 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
14753 .await
14754 .expect("body");
14755 serde_json::from_slice(&bytes).expect("json body")
14756 }
14757
14758 async fn chat_completion_admitted(st: &AppState, req: serde_json::Value) -> Response {
14774 let mut last_shed = serde_json::Value::Null;
14775 for _ in 0..50 {
14776 let resp = chat_completions(
14777 State(st.clone()),
14778 HeaderMap::new(),
14779 None,
14780 Json(serde_json::from_value(req.clone()).unwrap()),
14781 )
14782 .await;
14783 if resp.status() != StatusCode::TOO_MANY_REQUESTS {
14784 return resp;
14785 }
14786 let body = body_value(resp).await;
14787 let code = body["error"]["code"].as_str().unwrap_or_default();
14788 assert!(
14789 code.starts_with("shed_"),
14790 "only a contention shed may be retried; any other 429 is a finding: {body}"
14791 );
14792 last_shed = body;
14793 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
14794 }
14795 panic!(
14798 "still shed after 50 attempts — either load the retry budget cannot absorb \
14799 or a shed that no longer clears; last refusal: {last_shed}"
14800 );
14801 }
14802
14803 #[test]
14804 fn timeout_ms_parses_clamps_nothing_and_names_every_refusal() {
14805 assert_eq!(parse_timeout_ms(None).unwrap(), TIMEOUT_MS_DEFAULT);
14807 assert_eq!(
14808 parse_timeout_ms(Some(&serde_json::Value::Null)).unwrap(),
14809 TIMEOUT_MS_DEFAULT
14810 );
14811 for ms in [TIMEOUT_MS_MIN, 5_000, 45_000, TIMEOUT_MS_MAX] {
14815 assert_eq!(parse_timeout_ms(Some(&json!(ms))).unwrap(), ms);
14816 }
14817 for bad in [0u64, TIMEOUT_MS_MIN - 1, TIMEOUT_MS_MAX + 1, 600_000] {
14819 let err = parse_timeout_ms(Some(&json!(bad))).expect_err("out of range must refuse");
14820 assert!(err.contains("timeout_ms"), "{err}");
14821 assert!(
14822 err.contains(&TIMEOUT_MS_MIN.to_string())
14823 && err.contains(&TIMEOUT_MS_MAX.to_string()),
14824 "the message must state the range: {err}"
14825 );
14826 assert!(
14827 err.contains("stream"),
14828 "the message must point at streaming for longer work: {err}"
14829 );
14830 }
14831 for bad in [json!("30s"), json!(1.5), json!(true), json!({}), json!([])] {
14833 let err = parse_timeout_ms(Some(&bad)).expect_err("bad type must refuse");
14834 assert!(
14835 err.contains("timeout_ms") && err.contains("stream"),
14836 "{err}"
14837 );
14838 }
14839 assert!(parse_timeout_ms(Some(&json!(-1))).is_err());
14841 }
14842
14843 #[tokio::test]
14846 #[allow(clippy::await_holding_lock)] async fn a_bad_timeout_ms_is_the_same_named_400_on_every_surface() {
14848 let _l = drain_lock();
14849 let st = fake_worker_state();
14850
14851 let comp = completions(
14852 State(st.clone()),
14853 HeaderMap::new(),
14854 None,
14855 Json(
14856 serde_json::from_value(json!({
14857 "model": "m", "prompt": "t", "timeout_ms": 90_001}))
14858 .unwrap(),
14859 ),
14860 )
14861 .await;
14862 assert_eq!(comp.status(), StatusCode::BAD_REQUEST);
14863 let chat = chat_completions(
14864 State(st.clone()),
14865 HeaderMap::new(),
14866 None,
14867 Json(
14868 serde_json::from_value(json!({
14869 "model": "m", "messages": [{"role": "user", "content": "t"}],
14870 "timeout_ms": 90_001}))
14871 .unwrap(),
14872 ),
14873 )
14874 .await;
14875 assert_eq!(chat.status(), StatusCode::BAD_REQUEST);
14876 let resp_api = responses_api::responses(
14877 State(st.clone()),
14878 HeaderMap::new(),
14879 None,
14880 axum::body::Bytes::from(
14881 json!({"model": "m", "input": "t", "timeout_ms": 90_001}).to_string(),
14882 ),
14883 )
14884 .await;
14885 assert_eq!(resp_api.status(), StatusCode::BAD_REQUEST);
14886 let msgs = anthropic::messages(
14887 State(st.clone()),
14888 HeaderMap::new(),
14889 None,
14890 axum::body::Bytes::from(
14891 json!({"model": "m", "max_tokens": 16,
14892 "messages": [{"role": "user", "content": "t"}],
14893 "timeout_ms": 90_001})
14894 .to_string(),
14895 ),
14896 )
14897 .await;
14898 assert_eq!(msgs.status(), StatusCode::BAD_REQUEST);
14899
14900 for (surface, resp) in [
14902 ("/v1/completions", comp),
14903 ("/v1/chat/completions", chat),
14904 ("/v1/responses", resp_api),
14905 ] {
14906 let body = body_value(resp).await;
14907 assert_eq!(body["error"]["type"], "invalid_request_error", "{surface}");
14908 assert_eq!(body["error"]["param"], "timeout_ms", "{surface}");
14909 let m = body["error"]["message"].as_str().unwrap();
14910 assert!(
14911 m.contains("90000") && m.contains("stream"),
14912 "{surface}: {m}"
14913 );
14914 }
14915 let body = body_value(msgs).await;
14917 assert_eq!(body["error"]["type"], "invalid_request_error");
14918 let m = body["error"]["message"].as_str().unwrap();
14919 assert!(m.contains("timeout_ms") && m.contains("stream"), "{m}");
14920 }
14921
14922 #[tokio::test]
14925 #[allow(clippy::await_holding_lock)] async fn a_non_integer_timeout_ms_is_a_named_400() {
14927 let _l = drain_lock();
14928 let st = fake_worker_state();
14929 let resp = chat_completions(
14930 State(st),
14931 HeaderMap::new(),
14932 None,
14933 Json(
14934 serde_json::from_value(json!({
14935 "model": "m", "messages": [{"role": "user", "content": "t"}],
14936 "timeout_ms": "30s"}))
14937 .unwrap(),
14938 ),
14939 )
14940 .await;
14941 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
14942 let body = body_value(resp).await;
14943 assert_eq!(body["error"]["param"], "timeout_ms");
14944 }
14945
14946 #[tokio::test]
14952 #[allow(clippy::await_holding_lock)] async fn a_missed_non_stream_deadline_delivers_the_partial_bills_it_and_cancels_generation() {
14954 let _l = drain_lock();
14955 let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
14960 let cancel_seen = Arc::new(std::sync::atomic::AtomicBool::new(false));
14961 let worker_cancel = cancel_seen.clone();
14962 let health = health::WorkerHealth::new();
14963 let h = health.clone();
14964 std::thread::spawn(move || {
14965 h.mark_ready();
14966 while let Ok(Cmd::Generate(req)) = cmd_rx.recv() {
14967 worker::release_pending_admit();
14968 worker::release_admission_reservation(req.lane);
14969 let _ = req.tx.send(Event::PromptUsage {
14970 n_prompt: 1,
14971 n_cached: 0,
14972 });
14973 let _ = req.tx.send(Event::Token {
14974 id: 1,
14975 text: "partial".into(),
14976 });
14977 for _ in 0..5_000 {
14980 if req.tx.is_closed() {
14981 worker_cancel.store(true, std::sync::atomic::Ordering::SeqCst);
14982 break;
14983 }
14984 std::thread::sleep(std::time::Duration::from_millis(1));
14985 }
14986 }
14987 });
14988 for _ in 0..2_000 {
14989 if health.live().is_ok() {
14990 break;
14991 }
14992 std::thread::sleep(std::time::Duration::from_millis(1));
14993 }
14994 let mut st = fake_worker_state();
14995 st.cmd_tx = cmd_tx;
14996 st.health = health;
14997 let mock = MockMetering::admit_all();
14998 st.metering = Some(mock.clone());
14999
15000 let resp = chat_completion_admitted(
15001 &st,
15002 json!({
15003 "model": "m", "messages": [{"role": "user", "content": "t"}],
15004 "timeout_ms": 1_000}),
15005 )
15006 .await;
15007
15008 assert_eq!(resp.status(), StatusCode::OK);
15013 let body = body_value(resp).await;
15014 assert!(
15015 body["choices"][0]["message"]["content"]
15016 .as_str()
15017 .unwrap()
15018 .contains("partial"),
15019 "the tokens generated before the cut must be delivered: {body}"
15020 );
15021 assert_eq!(body["choices"][0]["finish_reason"], "error");
15025 assert_eq!(
15026 body["choices"][0]["native_finish_reason"],
15027 "deadline_exceeded"
15028 );
15029 assert_eq!(body["error"]["code"], "deadline_exceeded");
15030 assert_eq!(body["error"]["metadata"]["error_type"], "timeout");
15031 let message = body["error"]["message"].as_str().unwrap();
15032 assert!(
15033 message.contains("1000") && message.contains("stream"),
15034 "the partial must name the deadline and the streaming alternative: {message}"
15035 );
15036 assert_eq!(body["usage"]["completion_tokens"], 1);
15037
15038 let mut cancelled = false;
15043 for _ in 0..500 {
15044 if cancel_seen.load(std::sync::atomic::Ordering::SeqCst) {
15045 cancelled = true;
15046 break;
15047 }
15048 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
15049 }
15050 assert!(
15051 cancelled,
15052 "the deadline must CANCEL generation (worker's event channel closed)"
15053 );
15054
15055 let events = mock.events();
15060 assert!(
15061 events.contains(&MeterEvent::DeadlinePartial {
15062 prompt: 1,
15063 cached: 0,
15064 completion: 1,
15065 }),
15066 "the partial must settle as a deadline-partial with worker-truth counts: {events:?}"
15067 );
15068 assert!(
15069 !events
15070 .iter()
15071 .any(|e| matches!(e, MeterEvent::Complete { .. })),
15072 "a deadline cut must stay distinguishable from a full answer: {events:?}"
15073 );
15074 }
15075
15076 #[tokio::test]
15080 #[allow(clippy::await_holding_lock)] async fn a_deadline_missed_before_any_token_is_still_408_and_unbilled() {
15082 let _l = drain_lock();
15083 let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
15084 let health = health::WorkerHealth::new();
15085 let h = health.clone();
15086 std::thread::spawn(move || {
15087 h.mark_ready();
15088 while let Ok(Cmd::Generate(req)) = cmd_rx.recv() {
15091 worker::release_pending_admit();
15092 worker::release_admission_reservation(req.lane);
15093 let _ = req.tx.send(Event::PromptUsage {
15094 n_prompt: 1,
15095 n_cached: 0,
15096 });
15097 for _ in 0..5_000 {
15098 if req.tx.is_closed() {
15099 break;
15100 }
15101 std::thread::sleep(std::time::Duration::from_millis(1));
15102 }
15103 }
15104 });
15105 for _ in 0..2_000 {
15106 if health.live().is_ok() {
15107 break;
15108 }
15109 std::thread::sleep(std::time::Duration::from_millis(1));
15110 }
15111 let mut st = fake_worker_state();
15112 st.cmd_tx = cmd_tx;
15113 st.health = health;
15114 let mock = MockMetering::admit_all();
15115 st.metering = Some(mock.clone());
15116 let resp = chat_completion_admitted(
15117 &st,
15118 json!({
15119 "model": "m", "messages": [{"role": "user", "content": "t"}],
15120 "timeout_ms": 1_000}),
15121 )
15122 .await;
15123 assert_eq!(resp.status(), StatusCode::REQUEST_TIMEOUT);
15124 assert!(resp.headers().get("x-should-retry").is_none());
15126 assert_eq!(retry_after(&resp), None);
15127 let body = body_value(resp).await;
15128 assert_eq!(body["error"]["code"], "deadline_exceeded");
15129 assert!(
15130 body["error"]["message"]
15131 .as_str()
15132 .unwrap()
15133 .contains("not billed"),
15134 "the zero-token 408 keeps the billing promise: {body}"
15135 );
15136 let events = mock.events();
15137 assert!(
15138 events.contains(&MeterEvent::Unbilled {
15139 outcome: "deadline_exceeded",
15140 status: 408,
15141 code: "deadline_exceeded".into(),
15142 }),
15143 "the named zero-debit census outcome, not the generic reject — every sibling \
15144 deadline path settles this one: {events:?}"
15145 );
15146 }
15147
15148 #[tokio::test]
15151 #[allow(clippy::await_holding_lock)] async fn a_stream_that_misses_ttft_is_a_preheader_408_and_not_billed() {
15153 let _l = drain_lock();
15154 let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
15156 let health = health::WorkerHealth::new();
15157 let h = health.clone();
15158 std::thread::spawn(move || {
15159 h.mark_ready();
15160 while let Ok(Cmd::Generate(req)) = cmd_rx.recv() {
15161 worker::release_pending_admit();
15162 worker::release_admission_reservation(req.lane);
15163 let _ = req.tx.send(Event::PromptUsage {
15164 n_prompt: 1,
15165 n_cached: 0,
15166 });
15167 while !req.tx.is_closed() {
15168 std::thread::sleep(std::time::Duration::from_millis(1));
15169 }
15170 }
15171 });
15172 for _ in 0..2_000 {
15173 if health.live().is_ok() {
15174 break;
15175 }
15176 std::thread::sleep(std::time::Duration::from_millis(1));
15177 }
15178 let mut st = fake_worker_state();
15179 st.cmd_tx = cmd_tx;
15180 st.health = health;
15181 let mock = MockMetering::admit_all();
15182 st.metering = Some(mock.clone());
15183
15184 let resp = chat_completion_admitted(
15185 &st,
15186 json!({
15187 "model": "m", "messages": [{"role": "user", "content": "t"}],
15188 "stream": true, "timeout_ms": 1_000}),
15189 )
15190 .await;
15191 assert_eq!(resp.status(), StatusCode::REQUEST_TIMEOUT);
15194 let body = body_value(resp).await;
15195 assert_eq!(body["error"]["code"], "deadline_exceeded");
15196 assert!(
15197 body["error"]["message"]
15198 .as_str()
15199 .unwrap()
15200 .contains("first token"),
15201 "the streaming message must say the deadline bounded TIME TO FIRST TOKEN: {body}"
15202 );
15203 let events = mock.events();
15204 assert!(
15205 events.contains(&MeterEvent::Unbilled {
15206 outcome: "deadline_exceeded",
15207 status: 408,
15208 code: "deadline_exceeded".into(),
15209 }),
15210 "a TTFT miss must settle unbilled under the deadline outcome: {events:?}"
15211 );
15212 }
15213
15214 #[tokio::test]
15218 #[allow(clippy::await_holding_lock)] async fn a_stream_is_immune_to_the_deadline_after_its_first_token() {
15220 let _l = drain_lock();
15221 let mut st = fake_worker_state_with_steps(4, std::time::Duration::from_millis(400));
15224 let mock = MockMetering::admit_all();
15225 st.metering = Some(mock.clone());
15226 let resp = chat_completion_admitted(
15227 &st,
15228 json!({
15229 "model": "m", "messages": [{"role": "user", "content": "t"}],
15230 "stream": true, "timeout_ms": 1_000}),
15231 )
15232 .await;
15233 assert_eq!(
15234 resp.status(),
15235 StatusCode::OK,
15236 "TTFT was met — 200 is correct"
15237 );
15238 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
15239 .await
15240 .expect("the stream must run to completion past the deadline");
15241 let text = String::from_utf8(bytes.to_vec()).unwrap();
15242 assert!(text.contains("[DONE]"), "stream did not complete: {text}");
15243 let events = mock.events();
15244 assert!(
15245 events
15246 .iter()
15247 .any(|e| matches!(e, MeterEvent::Complete { completion: 4, .. })),
15248 "a stream past its deadline after first token still settles as COMPLETE with \
15249 all four tokens: {events:?}"
15250 );
15251 }
15252
15253 fn admission_counters_guard() -> std::sync::MutexGuard<'static, ()> {
15260 static COUNTERS: std::sync::Mutex<()> = std::sync::Mutex::new(());
15261 COUNTERS
15262 .lock()
15263 .unwrap_or_else(|poisoned| poisoned.into_inner())
15264 }
15265
15266 struct CounterRestore<'a>(&'a std::sync::atomic::AtomicUsize, usize);
15272 impl Drop for CounterRestore<'_> {
15273 fn drop(&mut self) {
15274 self.0.store(self.1, std::sync::atomic::Ordering::Release);
15275 }
15276 }
15277
15278 #[allow(clippy::result_large_err)] fn reserve_interactive_through_contention(
15289 st: &AppState,
15290 rl: &RateLimit,
15291 deadline_ms: u64,
15292 ) -> Result<PendingAdmissionGuard, (Response, &'static str)> {
15293 let reserve = || {
15294 reserve_pending_admit(
15295 st,
15296 lanes::Lane::Interactive,
15297 rl,
15298 RequestDeadline::starting_now(deadline_ms),
15299 )
15300 };
15301 let mut g = reserve();
15302 for _ in 0..50 {
15303 match &g {
15304 Ok(_) => break,
15305 Err((_, "shed_deadline")) => {
15306 std::thread::sleep(std::time::Duration::from_millis(10));
15307 g = reserve();
15308 }
15309 Err((_, outcome)) => panic!("unexpected refusal: {outcome}"),
15310 }
15311 }
15312 g
15313 }
15314
15315 #[test]
15318 fn the_queue_bound_sheds_with_429_retry_after_and_the_ratelimit_trio() {
15319 let _counters = admission_counters_guard();
15320 let st = fake_worker_state();
15321 let lane = lanes::Lane::Interactive;
15322 let cap = lane_cap(lane);
15323 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15324 let prev = counter.swap(max_queue_depth(cap), std::sync::atomic::Ordering::AcqRel);
15325 let _restore = CounterRestore(counter, prev);
15326 let rl = RateLimit {
15327 limit: cap,
15328 remaining: 0,
15329 reset_s: 1,
15330 };
15331 let (resp, outcome) = reserve_pending_admit(
15332 &st,
15333 lane,
15334 &rl,
15335 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15336 )
15337 .map(|_| ())
15338 .expect_err("a backlog at the bound must shed");
15339 assert_eq!(outcome, "shed_queue");
15340 assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
15341 assert!(
15342 retry_after(&resp).is_some(),
15343 "a shed must carry Retry-After so the router's spill can act on it"
15344 );
15345 let stamped = rl.attach(resp);
15347 for h in [
15348 "x-ratelimit-limit",
15349 "x-ratelimit-remaining",
15350 "x-ratelimit-reset",
15351 ] {
15352 assert!(stamped.headers().get(h).is_some(), "missing {h}");
15353 }
15354 }
15355
15356 #[test]
15360 fn admission_sheds_only_when_the_estimated_wait_cannot_fit_the_deadline() {
15361 let _counters = admission_counters_guard();
15362 let st = fake_worker_state();
15363 let lane = lanes::Lane::Interactive;
15364 let cap = lane_cap(lane);
15365 {
15366 let mut m = st.metrics.lock().unwrap();
15367 m.completed = 10;
15368 m.tokens_out = 1_000;
15369 m.step_p50_ms = 10.0; }
15371 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15372 let prev = counter.swap(cap, std::sync::atomic::Ordering::AcqRel); let _restore = CounterRestore(counter, prev);
15374 let rl = RateLimit {
15375 limit: cap,
15376 remaining: 0,
15377 reset_s: 1,
15378 };
15379 let admitted = reserve_pending_admit(
15381 &st,
15382 lane,
15383 &rl,
15384 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15385 );
15386 assert!(
15387 admitted.is_ok(),
15388 "a request whose deadline covers the estimate must be admitted"
15389 );
15390 drop(admitted); let (resp, outcome) = reserve_pending_admit(
15393 &st,
15394 lane,
15395 &rl,
15396 RequestDeadline::starting_now(TIMEOUT_MS_MIN),
15397 )
15398 .map(|_| ())
15399 .expect_err("a deadline shorter than the estimated wait must shed");
15400 assert_eq!(outcome, "shed_deadline");
15401 assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
15402 assert!(retry_after(&resp).is_some());
15403 }
15404
15405 #[test]
15408 fn deadline_shed_is_interactive_only_and_silent_with_free_slots() {
15409 let _counters = admission_counters_guard();
15410 let st = fake_worker_state();
15411 let cap = lane_cap(lanes::Lane::Interactive);
15412 {
15413 let mut m = st.metrics.lock().unwrap();
15414 m.completed = 10;
15415 m.tokens_out = 100_000; m.step_p50_ms = 100.0;
15417 }
15418 let free = RateLimit {
15420 limit: cap,
15421 remaining: 1,
15422 reset_s: 0,
15423 };
15424 let g = reserve_interactive_through_contention(&st, &free, TIMEOUT_MS_MIN);
15429 assert!(
15430 g.is_ok(),
15431 "free capacity must admit regardless of the estimate"
15432 );
15433 drop(g);
15434 let full = RateLimit {
15437 limit: cap,
15438 remaining: 0,
15439 reset_s: 5,
15440 };
15441 for lane in [lanes::Lane::Judge, lanes::Lane::Harvest] {
15442 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15443 let prev = counter.swap(1, std::sync::atomic::Ordering::AcqRel); let _restore = CounterRestore(counter, prev);
15445 let g = reserve_pending_admit(
15446 &st,
15447 lane,
15448 &full,
15449 RequestDeadline::starting_now(TIMEOUT_MS_MIN),
15450 );
15451 assert!(
15452 g.is_ok(),
15453 "{lane:?} must not be deadline-shed by the interactive gate"
15454 );
15455 drop(g);
15456 }
15457 }
15458
15459 #[test]
15466 fn a_saturated_queue_with_free_http_slots_queues_silently_without_a_ceiling() {
15467 let _counters = admission_counters_guard();
15468 let st = fake_worker_state();
15469 let lane = lanes::Lane::Interactive;
15470 let cap = lane_cap(lane);
15471 {
15472 let mut m = st.metrics.lock().unwrap();
15473 m.completed = 10;
15474 m.tokens_out = 1_000; m.step_p50_ms = 100.0; }
15477 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15478 let prev = counter.swap(cap, std::sync::atomic::Ordering::AcqRel); let _restore = CounterRestore(counter, prev);
15480 let free = RateLimit {
15483 limit: cap,
15484 remaining: 1,
15485 reset_s: 0,
15486 };
15487 let g = reserve_pending_admit(
15488 &st,
15489 lane,
15490 &free,
15491 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15492 );
15493 assert!(
15494 g.is_ok(),
15495 "flag off: a ~20 s projected wait whose deadline can absorb it queues \
15496 silently (no 429) - the darklanes#5 defect shape, preserved by default"
15497 );
15498 drop(g);
15499 }
15500
15501 #[test]
15506 fn the_queue_wait_ceiling_sheds_with_429_retry_after_and_the_ratelimit_trio() {
15507 let _counters = admission_counters_guard();
15508 let st = fake_worker_state();
15509 let lane = lanes::Lane::Interactive;
15510 let cap = lane_cap(lane);
15511 {
15512 let mut m = st.metrics.lock().unwrap();
15513 m.completed = 10;
15514 m.tokens_out = 1_000; m.step_p50_ms = 100.0; }
15517 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15518 let prev = counter.swap(cap, std::sync::atomic::Ordering::AcqRel); let _restore = CounterRestore(counter, prev);
15520 let free = RateLimit {
15521 limit: cap,
15522 remaining: 1,
15523 reset_s: 0,
15524 };
15525 let (resp, outcome) = reserve_pending_admit_with_ceiling(
15526 &st,
15527 lane,
15528 &free,
15529 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15530 5, )
15532 .map(|_| ())
15533 .expect_err("a projected wait past the ceiling must shed");
15534 assert_eq!(outcome, "shed_queue_wait");
15535 assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
15536 assert_eq!(
15537 retry_after(&resp).as_deref(),
15538 Some("20"),
15539 "Retry-After must carry the estimate (~10 s/wave x 2 waves)"
15540 );
15541 assert_eq!(
15542 resp.headers()
15543 .get("retry-after-ms")
15544 .and_then(|v| v.to_str().ok()),
15545 Some("20000"),
15546 "the ms twin must match"
15547 );
15548 let stamped = free.attach(resp);
15549 for h in [
15550 "x-ratelimit-limit",
15551 "x-ratelimit-remaining",
15552 "x-ratelimit-reset",
15553 ] {
15554 assert!(stamped.headers().get(h).is_some(), "missing {h}");
15555 }
15556 }
15557
15558 #[test]
15562 fn the_queue_wait_ceiling_admits_under_it_and_never_touches_dark_lanes() {
15563 let _counters = admission_counters_guard();
15564 let st = fake_worker_state();
15565 let lane = lanes::Lane::Interactive;
15566 let cap = lane_cap(lane);
15567 {
15568 let mut m = st.metrics.lock().unwrap();
15569 m.completed = 10;
15570 m.tokens_out = 1_000;
15571 m.step_p50_ms = 100.0; }
15573 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15574 let prev = counter.swap(cap, std::sync::atomic::Ordering::AcqRel);
15575 let _restore = CounterRestore(counter, prev);
15576 let free = RateLimit {
15577 limit: cap,
15578 remaining: 1,
15579 reset_s: 0,
15580 };
15581 let g = reserve_pending_admit_with_ceiling(
15582 &st,
15583 lane,
15584 &free,
15585 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15586 60, );
15588 assert!(
15589 g.is_ok(),
15590 "an estimate under the ceiling must admit and queue as before"
15591 );
15592 drop(g);
15593 let full = RateLimit {
15595 limit: cap,
15596 remaining: 0,
15597 reset_s: 5,
15598 };
15599 for dark in [lanes::Lane::Judge, lanes::Lane::Harvest] {
15600 let counter = &worker::ADMISSION_RESERVATIONS[dark.idx()];
15601 let prev = counter.swap(1, std::sync::atomic::Ordering::AcqRel); let _restore = CounterRestore(counter, prev);
15603 let g = reserve_pending_admit_with_ceiling(
15604 &st,
15605 dark,
15606 &full,
15607 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15608 1,
15609 );
15610 assert!(
15611 g.is_ok(),
15612 "{dark:?} must not be shed by the interactive queue-wait ceiling"
15613 );
15614 drop(g);
15615 }
15616 }
15617
15618 #[test]
15622 fn the_queue_wait_ceiling_leaves_the_existing_shed_arms_first_and_unchanged() {
15623 let _counters = admission_counters_guard();
15624 let st = fake_worker_state();
15625 let lane = lanes::Lane::Interactive;
15626 let cap = lane_cap(lane);
15627 {
15628 let mut m = st.metrics.lock().unwrap();
15629 m.completed = 10;
15630 m.tokens_out = 1_000;
15631 m.step_p50_ms = 100.0;
15632 }
15633 let rl = RateLimit {
15634 limit: cap,
15635 remaining: 0,
15636 reset_s: 1,
15637 };
15638 let counter = &worker::ADMISSION_RESERVATIONS[lane.idx()];
15639 let prev = counter.swap(max_queue_depth(cap), std::sync::atomic::Ordering::AcqRel);
15641 let _restore = CounterRestore(counter, prev);
15642 assert!(matches!(
15643 reserve_pending_admit_with_ceiling(
15644 &st,
15645 lane,
15646 &rl,
15647 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15648 1,
15649 ),
15650 Err((_, "shed_queue"))
15651 ));
15652 counter.store(cap, std::sync::atomic::Ordering::Release);
15654 assert!(matches!(
15655 reserve_pending_admit_with_ceiling(
15656 &st,
15657 lane,
15658 &rl,
15659 RequestDeadline::starting_now(TIMEOUT_MS_MIN),
15660 1,
15661 ),
15662 Err((_, "shed_deadline"))
15663 ));
15664 }
15665
15666 #[test]
15671 fn the_queue_wait_ceiling_is_wired_through_the_production_wrapper() {
15672 let src = include_str!("lib.rs");
15673 let code: String = src
15674 .lines()
15675 .map(|l| l.split("//").next().unwrap_or(""))
15676 .collect::<Vec<_>>()
15677 .join("\n");
15678 let start = code
15679 .find("pub(crate) fn reserve_pending_admit(")
15680 .expect("the production wrapper exists");
15681 let rest = &code[start..];
15682 let end = rest.find("\nfn ").unwrap_or(rest.len());
15683 let wrapper = &rest[..end];
15684 assert!(
15685 wrapper.contains(
15686 "reserve_pending_admit_with_ceiling(st, lane, rl, deadline, queue_wait_ceiling_s())"
15687 ),
15688 "every production ingress must judge the ceiling the env read armed"
15689 );
15690 }
15691
15692 #[test]
15693 fn pending_admission_reservation_is_atomic_and_rolls_back_on_drop() {
15694 let _counters = admission_counters_guard();
15695 let st = fake_worker_state();
15696 let cap = lane_cap(lanes::Lane::Interactive);
15697 let bound = max_queue_depth(cap);
15698 assert!(bound > 0, "the queue bound must admit at least one request");
15699 let rl = RateLimit {
15700 limit: cap,
15701 remaining: 0,
15702 reset_s: 1,
15703 };
15704 let _ = worker::PENDING_ADMITS.fetch_update(
15705 std::sync::atomic::Ordering::AcqRel,
15706 std::sync::atomic::Ordering::Acquire,
15707 |_| Some(0),
15708 );
15709 let counter = &worker::ADMISSION_RESERVATIONS[lanes::Lane::Interactive.idx()];
15710 let _restore = CounterRestore(counter, 0);
15711 counter.store(bound - 1, std::sync::atomic::Ordering::Release);
15712 let guard = reserve_pending_admit(
15713 &st,
15714 lanes::Lane::Interactive,
15715 &rl,
15716 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15717 )
15718 .expect("the final queue slot should be reservable");
15719 assert_eq!(
15720 worker::PENDING_ADMITS.load(std::sync::atomic::Ordering::Acquire),
15721 1
15722 );
15723 assert_eq!(counter.load(std::sync::atomic::Ordering::Acquire), bound);
15724 drop(guard);
15725 assert_eq!(
15726 worker::PENDING_ADMITS.load(std::sync::atomic::Ordering::Acquire),
15727 0
15728 );
15729 assert_eq!(
15730 counter.load(std::sync::atomic::Ordering::Acquire),
15731 bound - 1
15732 );
15733
15734 counter.store(bound, std::sync::atomic::Ordering::Release);
15735 let rejected = reserve_pending_admit(
15736 &st,
15737 lanes::Lane::Interactive,
15738 &rl,
15739 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15740 );
15741 assert!(matches!(rejected, Err((_, "shed_queue"))));
15742 }
15743
15744 #[test]
15745 fn admission_reservations_are_lane_scoped() {
15746 let _counters = admission_counters_guard();
15747 let st = fake_worker_state();
15748 let harvest = lanes::Lane::Harvest;
15749 let interactive = lanes::Lane::Interactive;
15750 let harvest_counter = &worker::ADMISSION_RESERVATIONS[harvest.idx()];
15751 let interactive_counter = &worker::ADMISSION_RESERVATIONS[interactive.idx()];
15752 let _restore = CounterRestore(harvest_counter, 0);
15753 harvest_counter.store(
15754 max_queue_depth(lane_cap(harvest)),
15755 std::sync::atomic::Ordering::Release,
15756 );
15757 interactive_counter.store(0, std::sync::atomic::Ordering::Release);
15758 let free = RateLimit {
15759 limit: lane_cap(interactive),
15760 remaining: 1,
15761 reset_s: 0,
15762 };
15763 let guard = reserve_pending_admit(
15774 &st,
15775 interactive,
15776 &free,
15777 RequestDeadline::starting_now(TIMEOUT_MS_MAX),
15778 )
15779 .expect("a full harvest queue must not consume interactive capacity");
15780 drop(guard);
15781 let tight = reserve_interactive_through_contention(&st, &free, TIMEOUT_MS_MIN);
15782 assert!(
15783 tight.is_ok(),
15784 "a full harvest queue must not deadline-shed a tight interactive request \
15785 (a backlog that outlasts the retry budget here is a cross-lane leak, not \
15786 contention)"
15787 );
15788 drop(tight);
15789 let harvest_rl = RateLimit {
15790 limit: lane_cap(harvest),
15791 remaining: 0,
15792 reset_s: 1,
15793 };
15794 assert!(matches!(
15795 reserve_pending_admit(
15796 &st,
15797 harvest,
15798 &harvest_rl,
15799 RequestDeadline::starting_now(TIMEOUT_MS_MAX)
15800 ),
15801 Err((_, "shed_queue"))
15802 ));
15803 }
15804
15805 #[test]
15806 fn taxonomy_maps_every_class_to_its_status_and_code() {
15807 use worker::{EngineError as E, ErrClass as C};
15808 let cases: Vec<(worker::EngineError, StatusCode, &str, &str)> = vec![
15809 (
15810 E::invalid_param("bad json", "response_format"),
15811 StatusCode::BAD_REQUEST,
15812 "invalid_request_error",
15813 "",
15814 ),
15815 (
15816 E::context_length("prompt (9000 tok) >= context cap (8192)"),
15817 StatusCode::BAD_REQUEST,
15818 "invalid_request_error",
15819 "context_length_exceeded",
15820 ),
15821 (
15822 E::model_not_found("unknown model \"nope\""),
15823 StatusCode::BAD_REQUEST,
15824 "invalid_request_error",
15825 "model_not_found",
15826 ),
15827 (
15828 E::rate_limit("lane judge is at capacity, retry"),
15829 StatusCode::TOO_MANY_REQUESTS,
15830 "rate_limit_error",
15831 "rate_limit_exceeded",
15832 ),
15833 (
15834 E::overloaded("no VRAM for a new session"),
15835 StatusCode::SERVICE_UNAVAILABLE,
15836 "server_error",
15837 "overloaded",
15838 ),
15839 (
15840 E::engine("graph step failed: launch error"),
15841 StatusCode::INTERNAL_SERVER_ERROR,
15842 "server_error",
15843 "engine_error",
15844 ),
15845 ];
15846 for (err, want_status, want_type, want_code) in cases {
15847 let (status, etype, code) = class_http(err.class);
15848 assert_eq!(status, want_status, "{:?}", err);
15849 assert_eq!(etype, want_type, "{:?}", err);
15850 if !want_code.is_empty() {
15851 assert_eq!(code, Some(want_code), "{:?}", err);
15852 }
15853 let body = engine_error_body(&err);
15855 assert_eq!(body["error"]["message"], err.message);
15856 assert_eq!(body["error"]["type"], want_type);
15857 }
15858 for c in [
15860 C::InvalidRequest,
15861 C::ContextLength,
15862 C::ModelNotFound,
15863 C::RateLimit,
15864 C::Overloaded,
15865 C::Engine,
15866 ] {
15867 let (s, t, _) = class_http(c);
15868 assert!(s.is_client_error() || s.is_server_error(), "{c:?} -> {s}");
15869 assert!(!t.is_empty());
15870 }
15871 }
15872
15873 #[test]
15874 fn a_cuda_oom_message_is_capacity_503_not_a_500() {
15875 let e = worker::EngineError::engine(
15880 "step error: DriverError(CUDA_ERROR_OUT_OF_MEMORY, \"out of memory\")",
15881 );
15882 let resp = engine_error_response(&e);
15883 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
15884 assert_eq!(retry_after(&resp).as_deref(), Some("5"));
15885 }
15886
15887 #[test]
15888 fn retry_headers_follow_the_sdk_contract() {
15889 for e in [
15893 worker::EngineError::rate_limit("shed"),
15894 worker::EngineError::overloaded("no VRAM"),
15895 ] {
15896 let resp = engine_error_response(&e);
15897 let ra = retry_after(&resp).expect("retryable class must carry Retry-After");
15898 let secs: u64 = ra
15899 .parse()
15900 .expect("Retry-After must be integer delay-seconds");
15901 assert!(
15902 secs > 0 && secs <= 60,
15903 "Retry-After {secs}s outside the honored window"
15904 );
15905 let ms = resp
15906 .headers()
15907 .get("retry-after-ms")
15908 .unwrap()
15909 .to_str()
15910 .unwrap();
15911 assert_eq!(
15912 ms.parse::<u64>().unwrap(),
15913 secs * 1000,
15914 "the two headers disagree"
15915 );
15916 assert!(
15917 resp.headers().get("x-should-retry").is_none(),
15918 "a retryable class must not say x-should-retry: false"
15919 );
15920 }
15921 }
15922
15923 #[tokio::test]
15930 async fn admit_predict_reject_matches_shed_contract() {
15931 let shed = retry_contract_response(
15933 (
15934 StatusCode::TOO_MANY_REQUESTS,
15935 Json(error_body(
15936 "interactive queue is at its bound",
15937 "rate_limit_error",
15938 None,
15939 Some("shed_queue"),
15940 )),
15941 )
15942 .into_response(),
15943 Some(7),
15944 );
15945 let predict = engine_error_response(&worker::EngineError::rate_limit_after(
15948 "predicted KV-to-completion exceeds the box budget; retry",
15949 7,
15950 ));
15951 assert_eq!(shed.status(), predict.status());
15952 for header in ["retry-after", "retry-after-ms"] {
15953 assert_eq!(
15954 shed.headers().get(header),
15955 predict.headers().get(header),
15956 "header {header} must be byte-identical to the shed contract"
15957 );
15958 }
15959 let shed_body: serde_json::Value = serde_json::from_slice(
15960 &axum::body::to_bytes(shed.into_body(), usize::MAX)
15961 .await
15962 .unwrap(),
15963 )
15964 .unwrap();
15965 let predict_body: serde_json::Value = serde_json::from_slice(
15966 &axum::body::to_bytes(predict.into_body(), usize::MAX)
15967 .await
15968 .unwrap(),
15969 )
15970 .unwrap();
15971 assert_eq!(shed_body["error"]["type"], predict_body["error"]["type"]);
15972 assert_eq!(predict_body["error"]["type"], "rate_limit_error");
15973 let shed_keys: Vec<&String> = shed_body["error"].as_object().unwrap().keys().collect();
15974 let predict_keys: Vec<&String> =
15975 predict_body["error"].as_object().unwrap().keys().collect();
15976 assert_eq!(shed_keys, predict_keys, "same body schema, key for key");
15977 assert_eq!(predict_body["error"]["code"], "rate_limit_exceeded");
15978
15979 let clamped = engine_error_response(&worker::EngineError::rate_limit_after("m", 400));
15981 assert_eq!(retry_after(&clamped).as_deref(), Some("60"));
15982 let plain = engine_error_response(&worker::EngineError::rate_limit("m"));
15984 assert_eq!(retry_after(&plain).as_deref(), Some("2"));
15985 }
15986
15987 #[tokio::test]
15988 #[allow(clippy::await_holding_lock)] async fn command_send_failure_obeys_the_retry_contract() {
15990 let _l = drain_lock();
15991 let mut st = fake_worker_state();
15992 let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
15993 drop(cmd_rx);
15994 st.cmd_tx = cmd_tx;
15995
15996 let completion = completions(
15997 State(st.clone()),
15998 axum::http::HeaderMap::new(),
15999 None,
16000 Json(
16001 serde_json::from_value(serde_json::json!({
16002 "model": "m", "prompt": "test"
16003 }))
16004 .unwrap(),
16005 ),
16006 )
16007 .await;
16008 let chat = chat_completions(
16009 State(st),
16010 axum::http::HeaderMap::new(),
16011 None,
16012 Json(
16013 serde_json::from_value(serde_json::json!({
16014 "model": "m", "messages": [{"role": "user", "content": "test"}]
16015 }))
16016 .unwrap(),
16017 ),
16018 )
16019 .await;
16020
16021 for resp in [completion, chat] {
16022 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
16023 assert_eq!(retry_after(&resp).as_deref(), Some("2"));
16024 assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
16025 assert_ne!(
16026 resp.headers()
16027 .get("x-should-retry")
16028 .and_then(|v| v.to_str().ok()),
16029 Some("false")
16030 );
16031 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
16032 .await
16033 .unwrap();
16034 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
16035 assert_eq!(payload["error"]["type"], "server_error");
16036 assert_eq!(payload["error"]["code"], "overloaded");
16037 }
16038 }
16039
16040 #[test]
16041 fn unfixable_client_errors_say_x_should_retry_false() {
16042 for e in [
16045 worker::EngineError::model_not_found("unknown model \"x\""),
16046 worker::EngineError::context_length("prompt too long"),
16047 worker::EngineError::invalid_param("bad", "messages"),
16048 ] {
16049 let resp = engine_error_response(&e);
16050 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
16051 assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
16052 assert!(
16053 retry_after(&resp).is_none(),
16054 "a 400 must not promise a retry window"
16055 );
16056 }
16057 }
16058
16059 #[tokio::test]
16060 async fn a_closed_worker_channel_is_503_not_500() {
16061 let (tx, rx) = worker::event_channel();
16065 drop(tx);
16066 let resp =
16067 blocking_response(rx, "m".into(), true, Vec::new(), None, Envelope::new(true)).await;
16068 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
16069 assert_eq!(retry_after(&resp).as_deref(), Some("5"));
16070 }
16071
16072 #[tokio::test]
16073 async fn a_dark_lane_shed_is_429_with_an_openai_object_body() {
16074 let (tx, rx) = worker::event_channel();
16077 tx.send(Event::Error(worker::EngineError::rate_limit(
16078 "lane judge shed: interactive p99 over budget, retry",
16079 )))
16080 .unwrap();
16081 let (resp, error_code) = peek_admission(rx)
16082 .await
16083 .expect_err("a shed must not be forwarded into the stream");
16084 assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
16085 assert_eq!(error_code, "rate_limit_exceeded");
16086 assert_eq!(retry_after(&resp).as_deref(), Some("2"));
16087 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
16088 .await
16089 .unwrap();
16090 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
16091 assert!(
16092 payload["error"].is_object(),
16093 "bare-string error body: {payload}"
16094 );
16095 assert_eq!(payload["error"]["type"], "rate_limit_error");
16096 assert!(
16097 payload["error"]["message"]
16098 .as_str()
16099 .unwrap()
16100 .contains("shed")
16101 );
16102 }
16103
16104 #[tokio::test]
16105 async fn interactive_admission_error_is_a_preheader_429() {
16106 let (tx, rx) = worker::event_channel();
16109 tx.send(Event::Error(worker::EngineError::rate_limit(
16110 "KV capacity unavailable",
16111 )))
16112 .unwrap();
16113 let (resp, error_code) = peek_admission(rx)
16114 .await
16115 .expect_err("admission error must stay pre-header");
16116 assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
16117 assert_eq!(error_code, "rate_limit_exceeded");
16118 }
16119
16120 #[tokio::test]
16121 async fn admission_peek_preserves_context_error_for_the_ledger() {
16122 let (tx, rx) = worker::event_channel();
16123 tx.send(Event::Error(worker::EngineError::context_length(
16124 "prompt exceeds configured model maximum",
16125 )))
16126 .unwrap();
16127 let (resp, error_code) = peek_admission(rx)
16128 .await
16129 .expect_err("context rejection must stay pre-header");
16130 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
16131 assert_eq!(error_code, "context_length_exceeded");
16132 }
16133
16134 #[tokio::test]
16135 async fn admission_peek_replays_prompt_usage_without_waiting_for_a_token() {
16136 let (tx, rx) = worker::event_channel();
16137 tx.send(Event::PromptUsage {
16138 n_prompt: 262_143,
16139 n_cached: 0,
16140 })
16141 .unwrap();
16142 let mut replay = peek_admission(rx).await.expect("successful admission");
16143 assert!(matches!(
16144 replay.recv().await,
16145 Some(Event::PromptUsage {
16146 n_prompt: 262_143,
16147 n_cached: 0
16148 }),
16149 ));
16150 }
16151
16152 #[test]
16153 fn penalties_plumb_from_http_to_sampler_config() {
16154 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
16157 "model": "m", "messages": [{"role": "user", "content": "task"}],
16158 "frequency_penalty": 0.5, "presence_penalty": 0.25, "repetition_penalty": 1.1
16159 }))
16160 .unwrap();
16161 let (tx, _rx) = worker::event_channel();
16162 let cfg = build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
16163 .unwrap()
16164 .request
16165 .sampler_cfg;
16166 assert_eq!(cfg.penalty_freq, 0.5);
16167 assert_eq!(cfg.penalty_present, 0.25);
16168 assert_eq!(cfg.penalty_repeat, 1.1);
16169 assert_eq!(cfg.penalty_last_n, memra_engine::spec::PEN_WINDOW_MAX);
16170
16171 let req: CompletionReq = serde_json::from_value(serde_json::json!({
16172 "model": "m", "prompt": "task", "frequency_penalty": 1.5
16173 }))
16174 .unwrap();
16175 let (tx, _rx) = worker::event_channel();
16176 let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
16177 assert_eq!(cfg.penalty_freq, 1.5);
16178 assert_eq!(cfg.penalty_last_n, memra_engine::spec::PEN_WINDOW_MAX);
16179
16180 let req: CompletionReq = serde_json::from_value(serde_json::json!({
16182 "model": "m", "prompt": "task"
16183 }))
16184 .unwrap();
16185 let (tx, _rx) = worker::event_channel();
16186 let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
16187 assert_eq!(cfg.penalty_last_n, 0);
16188 assert_eq!(cfg.penalty_repeat, 1.0);
16189 }
16190
16191 #[test]
16192 fn omitted_temperature_is_openai_default_not_greedy() {
16193 let chat_temp = |body: serde_json::Value| {
16209 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
16210 let (tx, _rx) = worker::event_channel();
16211 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
16212 .unwrap()
16213 .request
16214 .sampler_cfg
16215 .temperature
16216 };
16217 let comp_temp = |body: serde_json::Value| {
16218 let req: CompletionReq = serde_json::from_value(body).unwrap();
16219 let (tx, _rx) = worker::event_channel();
16220 build_request(&req, tx, lanes::Lane::Interactive, None)
16221 .sampler_cfg
16222 .temperature
16223 };
16224
16225 assert_eq!(
16227 chat_temp(serde_json::json!({
16228 "model": "m", "messages": [{"role": "user", "content": "t"}]})),
16229 1.0,
16230 "omitted chat temperature must be the OpenAI 1.0 default, not 0.0/greedy"
16231 );
16232 assert_eq!(
16233 comp_temp(serde_json::json!({
16234 "model": "m", "prompt": "t"})),
16235 1.0,
16236 "omitted completions temperature must be the OpenAI 1.0 default"
16237 );
16238
16239 assert_eq!(
16241 chat_temp(serde_json::json!({
16242 "model": "m", "messages": [{"role": "user", "content": "t"}],
16243 "temperature": 0.0})),
16244 0.0,
16245 "explicit temperature 0 must stay greedy"
16246 );
16247 assert_eq!(
16248 comp_temp(serde_json::json!({
16249 "model": "m", "prompt": "t", "temperature": 0})),
16250 0.0,
16251 "explicit temperature 0 must stay greedy"
16252 );
16253 assert!(
16255 memra_engine::sampler::Sampler::new(sampler_config(
16256 0.0,
16257 0,
16258 1.0,
16259 0.0,
16260 0.0,
16261 0.0,
16262 1.0,
16263 Some(0)
16264 ))
16265 .is_greedy()
16266 );
16267 assert!(
16268 !memra_engine::sampler::Sampler::new(sampler_config(
16269 1.0,
16270 0,
16271 1.0,
16272 0.0,
16273 0.0,
16274 0.0,
16275 1.0,
16276 Some(0)
16277 ))
16278 .is_greedy()
16279 );
16280
16281 assert_eq!(
16283 chat_temp(serde_json::json!({
16284 "model": "m", "messages": [{"role": "user", "content": "t"}],
16285 "temperature": 0.7})),
16286 0.7
16287 );
16288
16289 let req: CompletionReq = serde_json::from_value(serde_json::json!({
16293 "model": "m", "prompt": "t"}))
16294 .unwrap();
16295 let (tx, _rx) = worker::event_channel();
16296 let cfg = build_request(&req, tx, lanes::Lane::Interactive, None).sampler_cfg;
16297 assert_eq!(cfg.top_p, 1.0, "omitted top_p = OpenAI 1.0 = disabled");
16298 assert_eq!(cfg.top_k, 0, "omitted top_k = disabled");
16299 assert_eq!(cfg.min_p, 0.0, "omitted min_p = disabled");
16300 assert_eq!(cfg.penalty_last_n, 0, "omitted penalties = window off");
16301 assert!(
16306 memra_engine::sampler::Sampler::new(cfg).is_spec_sampling(),
16307 "the omitted-temperature default must ride sampled spec's pure-temp regime"
16308 );
16309 }
16310
16311 #[test]
16312 fn step35_chat_uses_published_sampling_defaults_only_when_omitted() {
16313 let caps = ModelCaps {
16314 chat_temperature_default: Some(0.5),
16315 chat_top_p_default: Some(0.9),
16316 chat_ok: true,
16317 ..Default::default()
16318 };
16319 let cfg = |extra: serde_json::Value| {
16320 let mut body = serde_json::json!({
16321 "model": "step35",
16322 "messages": [{"role": "user", "content": "task"}]
16323 });
16324 body.as_object_mut()
16325 .unwrap()
16326 .extend(extra.as_object().unwrap().clone());
16327 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
16328 let (tx, _rx) = worker::event_channel();
16329 build_chat_request(req, Some(&caps), tx, lanes::Lane::Interactive, None)
16330 .unwrap()
16331 .request
16332 .sampler_cfg
16333 };
16334
16335 let omitted = cfg(serde_json::json!({}));
16336 assert_eq!(omitted.temperature, 0.5);
16337 assert_eq!(omitted.top_p, 0.9);
16338
16339 let explicit_temp = cfg(serde_json::json!({"temperature": 0.7}));
16340 assert_eq!(explicit_temp.temperature, 0.7);
16341 assert_eq!(
16342 explicit_temp.top_p, 0.9,
16343 "omitting top_p must retain StepFun's nucleus default"
16344 );
16345
16346 let explicit = cfg(serde_json::json!({"temperature": 0.0, "top_p": 1.0}));
16347 assert_eq!(
16348 explicit.temperature, 0.0,
16349 "explicit greedy must remain authoritative"
16350 );
16351 assert_eq!(
16352 explicit.top_p, 1.0,
16353 "explicit untruncated sampling must remain authoritative"
16354 );
16355 }
16356
16357 fn qwen38_vendor_defaults() -> SamplingDefaults {
16361 SamplingDefaults {
16362 temperature: Some(1.0),
16363 top_p: Some(0.95),
16364 top_k: Some(20),
16365 min_p: Some(0.0),
16366 presence_penalty: Some(0.0),
16367 repetition_penalty: Some(1.0),
16368 frequency_penalty: None,
16369 }
16370 }
16371
16372 fn gemma4_vendor_defaults() -> SamplingDefaults {
16377 SamplingDefaults {
16378 temperature: Some(1.0),
16379 top_p: Some(0.95),
16380 top_k: Some(64),
16381 ..Default::default()
16382 }
16383 }
16384
16385 #[test]
16386 fn vendor_sampling_defaults_fill_only_the_omitted_fields() {
16387 let d = ModelSamplingDefaults::single(gemma4_vendor_defaults());
16392 let chat = |extra: serde_json::Value| {
16393 let mut body = serde_json::json!({
16394 "model": "google/gemma-4-31b-it",
16395 "messages": [{"role": "user", "content": "task"}],
16396 "seed": 7
16398 });
16399 body.as_object_mut()
16400 .unwrap()
16401 .extend(extra.as_object().unwrap().clone());
16402 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
16403 let (tx, _rx) = worker::event_channel();
16404 build_chat_request_with_trace(
16405 req,
16406 Some(&ModelCaps {
16407 chat_ok: true,
16408 ..Default::default()
16409 }),
16410 tx,
16411 lanes::Lane::Interactive,
16412 None,
16413 None,
16414 None,
16415 &d,
16416 )
16417 .unwrap()
16418 .request
16419 .sampler_cfg
16420 };
16421
16422 let omitted = chat(serde_json::json!({}));
16424 assert_eq!(omitted.temperature, 1.0, "gemma-4 card temperature");
16425 assert_eq!(omitted.top_p, 0.95, "gemma-4 card top_p");
16426 assert_eq!(omitted.top_k, 64, "gemma-4 card top_k");
16427 assert_eq!(omitted.min_p, 0.0, "undeclared min_p stays API-standard");
16429 assert_eq!(omitted.penalty_repeat, 1.0);
16430 assert_eq!(omitted.penalty_freq, 0.0);
16431 assert_eq!(omitted.penalty_present, 0.0);
16432 assert_eq!(omitted.penalty_last_n, 0, "no penalty => no history window");
16433 assert!(
16434 !memra_engine::sampler::Sampler::new(omitted).is_greedy(),
16435 "the vendor default must NOT be greedy — that is the whole point of the lane"
16436 );
16437
16438 let greedy = chat(serde_json::json!({"temperature": 0}));
16441 assert_eq!(
16442 greedy.temperature, 0.0,
16443 "explicit temperature 0 stays greedy"
16444 );
16445 assert!(
16446 memra_engine::sampler::Sampler::new(greedy).is_greedy(),
16447 "an explicit temperature 0 must satisfy the greedy predicate that gates the \
16448 spec/graph exactness arms"
16449 );
16450
16451 let one_field = chat(serde_json::json!({"top_k": 3}));
16453 assert_eq!(one_field.top_k, 3, "explicit top_k wins");
16454 assert_eq!(
16455 one_field.temperature, 1.0,
16456 "omitting temperature still takes the vendor value"
16457 );
16458 assert_eq!(one_field.top_p, 0.95, "omitting top_p still takes vendor");
16459
16460 let disabled = chat(serde_json::json!({"top_k": 0, "top_p": 1.0}));
16463 assert_eq!(
16464 disabled.top_k, 0,
16465 "an explicit top_k 0 means KEEP ALL, not 'unset'"
16466 );
16467 assert_eq!(
16468 disabled.top_p, 1.0,
16469 "an explicit top_p 1.0 means untruncated"
16470 );
16471
16472 let penal = chat(serde_json::json!({"presence_penalty": 1.5}));
16474 assert_eq!(penal.penalty_present, 1.5);
16475 assert_eq!(penal.penalty_last_n, memra_engine::spec::PEN_WINDOW_MAX);
16476 }
16477
16478 #[test]
16479 fn vendor_sampling_defaults_are_identical_on_every_surface() {
16480 let d = qwen38_vendor_defaults();
16492 let md = ModelSamplingDefaults::single(d);
16493 let comp = |extra: serde_json::Value| {
16494 let mut body = serde_json::json!({
16495 "model": "qwen/qwen3.8-27b", "prompt": "task", "seed": 11 });
16496 body.as_object_mut()
16497 .unwrap()
16498 .extend(extra.as_object().unwrap().clone());
16499 let req: CompletionReq = serde_json::from_value(body).unwrap();
16500 let (tx, _rx) = worker::event_channel();
16501 build_request_with_trace(&req, tx, lanes::Lane::Interactive, None, None, &d).sampler_cfg
16502 };
16503 let chat = |extra: serde_json::Value| {
16504 let mut body = serde_json::json!({
16505 "model": "qwen/qwen3.8-27b",
16506 "messages": [{"role": "user", "content": "task"}],
16507 "seed": 11 });
16508 body.as_object_mut()
16509 .unwrap()
16510 .extend(extra.as_object().unwrap().clone());
16511 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
16512 let (tx, _rx) = worker::event_channel();
16513 build_chat_request_with_trace(
16514 req,
16515 Some(&ModelCaps {
16516 chat_ok: true,
16517 ..Default::default()
16518 }),
16519 tx,
16520 lanes::Lane::Interactive,
16521 None,
16522 None,
16523 None,
16524 &md,
16525 )
16526 .unwrap()
16527 .request
16528 .sampler_cfg
16529 };
16530
16531 for extra in [
16532 serde_json::json!({}),
16533 serde_json::json!({"temperature": 0}),
16534 serde_json::json!({"temperature": 0.0}),
16535 serde_json::json!({"temperature": 0.7}),
16536 serde_json::json!({"top_p": 1.0}),
16537 serde_json::json!({"top_k": 0}),
16538 serde_json::json!({"min_p": 0.05}),
16539 serde_json::json!({"repetition_penalty": 1.1}),
16540 serde_json::json!({"frequency_penalty": 0.5}),
16541 serde_json::json!({"presence_penalty": 1.5}),
16542 serde_json::json!({
16543 "temperature": 0.3, "top_p": 0.5, "top_k": 7, "min_p": 0.02,
16544 "frequency_penalty": 0.1, "presence_penalty": 0.2,
16545 "repetition_penalty": 1.05 }),
16546 ] {
16547 let c = comp(extra.clone());
16548 let h = chat(extra.clone());
16549 assert_eq!(
16550 (
16551 c.temperature,
16552 c.top_p,
16553 c.top_k,
16554 c.min_p,
16555 c.penalty_repeat,
16556 c.penalty_freq,
16557 c.penalty_present,
16558 c.penalty_last_n,
16559 c.seed
16560 ),
16561 (
16562 h.temperature,
16563 h.top_p,
16564 h.top_k,
16565 h.min_p,
16566 h.penalty_repeat,
16567 h.penalty_freq,
16568 h.penalty_present,
16569 h.penalty_last_n,
16570 h.seed
16571 ),
16572 "/v1/completions and /v1/chat/completions disagree on {extra} — \
16573 standard-surface-law violation"
16574 );
16575 }
16576
16577 let omitted = comp(serde_json::json!({}));
16579 assert_eq!(
16580 omitted.temperature, 1.0,
16581 "qwen3.8 card thinking temperature"
16582 );
16583 assert_eq!(omitted.top_p, 0.95, "qwen3.8 card top_p");
16584 assert_eq!(omitted.top_k, 20, "qwen3.8 card top_k");
16585 assert!(
16587 memra_engine::sampler::Sampler::new(comp(serde_json::json!({"temperature": 0})))
16588 .is_greedy()
16589 );
16590 }
16591
16592 #[tokio::test]
16605 #[allow(clippy::await_holding_lock)] async fn same_omitted_request_resolves_identically_on_all_four_surfaces() {
16607 let _l = drain_lock();
16608 let step_caps = ModelCaps {
16609 chat_ok: true,
16610 chat_temperature_default: Some(0.5),
16611 chat_top_p_default: Some(0.9),
16612 ..Default::default()
16613 };
16614 let (cfg_tx, cfg_rx) = std::sync::mpsc::channel::<WorkerSaw>();
16615 let st = fake_worker_state_full(
16616 1,
16617 std::time::Duration::ZERO,
16618 HashMap::from([("m".to_string(), step_caps)]),
16619 Some(cfg_tx),
16620 );
16621 let fields = |saw: &WorkerSaw| {
16626 let c = &saw.sampler_cfg;
16627 (
16628 c.temperature,
16629 c.top_p,
16630 c.top_k,
16631 c.min_p,
16632 c.penalty_repeat,
16633 c.penalty_freq,
16634 c.penalty_present,
16635 c.penalty_last_n,
16636 )
16637 };
16638 let worker_saw = |surface: &str| {
16639 cfg_rx
16640 .recv_timeout(std::time::Duration::from_secs(10))
16641 .unwrap_or_else(|_| panic!("{surface}: request never reached the worker"))
16642 };
16643
16644 let resp = completions(
16645 State(st.clone()),
16646 axum::http::HeaderMap::new(),
16647 None,
16648 Json(serde_json::from_value(serde_json::json!({"model": "m", "prompt": "t"})).unwrap()),
16649 )
16650 .await;
16651 assert_eq!(
16652 resp.status(),
16653 StatusCode::OK,
16654 "/v1/completions rejected the omitted-sampling request"
16655 );
16656 let comp = worker_saw("/v1/completions");
16657
16658 let resp = chat_completions(
16659 State(st.clone()),
16660 axum::http::HeaderMap::new(),
16661 None,
16662 Json(
16663 serde_json::from_value(serde_json::json!({
16664 "model": "m", "messages": [{"role": "user", "content": "t"}]}))
16665 .unwrap(),
16666 ),
16667 )
16668 .await;
16669 assert_eq!(
16670 resp.status(),
16671 StatusCode::OK,
16672 "/v1/chat/completions rejected the omitted-sampling request"
16673 );
16674 let chat = worker_saw("/v1/chat/completions");
16675
16676 let resp = anthropic::messages(
16677 State(st.clone()),
16678 axum::http::HeaderMap::new(),
16679 None,
16680 axum::body::Bytes::from(
16681 serde_json::json!({
16682 "model": "m", "max_tokens": 16,
16683 "messages": [{"role": "user", "content": "t"}]})
16684 .to_string(),
16685 ),
16686 )
16687 .await;
16688 assert_eq!(
16689 resp.status(),
16690 StatusCode::OK,
16691 "/v1/messages rejected the omitted-sampling request"
16692 );
16693 let msg = worker_saw("/v1/messages");
16694
16695 let resp = responses_api::responses(
16696 State(st.clone()),
16697 axum::http::HeaderMap::new(),
16698 None,
16699 axum::body::Bytes::from(serde_json::json!({"model": "m", "input": "t"}).to_string()),
16700 )
16701 .await;
16702 assert_eq!(
16703 resp.status(),
16704 StatusCode::OK,
16705 "/v1/responses rejected the omitted-sampling request"
16706 );
16707 let rsp = worker_saw("/v1/responses");
16708
16709 for (surface, cfg) in [
16710 ("/v1/completions", &comp),
16711 ("/v1/messages", &msg),
16712 ("/v1/responses", &rsp),
16713 ] {
16714 assert_eq!(
16715 fields(cfg),
16716 fields(&chat),
16717 "{surface} resolved DIFFERENT effective sampling than /v1/chat/completions \
16718 for the same omitted-sampling request — standard-surface-law violation \
16719 (hermes d991b51699218285)"
16720 );
16721 }
16722 assert_eq!(
16725 (comp.sampler_cfg.temperature, comp.sampler_cfg.top_p),
16726 (0.5, 0.9),
16727 "an omitting client must get the model's vendor caps (Step-3.7: 0.5/0.9) on \
16728 EVERY surface, not the API-standard 1.0/1.0 (hermes d991b51699218285)"
16729 );
16730 }
16731
16732 #[tokio::test]
16742 #[allow(clippy::await_holding_lock)] async fn same_effort_value_resolves_identically_on_every_surface() {
16744 let _l = drain_lock();
16745 let caps = ModelCaps {
16748 chat_ok: true,
16749 effort_levels: true,
16750 ..Default::default()
16751 };
16752 let (saw_tx, saw_rx) = std::sync::mpsc::channel::<WorkerSaw>();
16753 let st = fake_worker_state_full(
16754 1,
16755 std::time::Duration::ZERO,
16756 HashMap::from([("m".to_string(), caps)]),
16757 Some(saw_tx),
16758 );
16759 let send = |st: AppState, surface: &'static str, effort: &'static str| async move {
16760 match surface {
16761 "/v1/chat/completions" => {
16762 chat_completions(
16763 State(st),
16764 axum::http::HeaderMap::new(),
16765 None,
16766 Json(
16767 serde_json::from_value(serde_json::json!({
16768 "model": "m", "max_tokens": 8,
16769 "reasoning_effort": effort,
16770 "messages": [{"role": "user", "content": "t"}]}))
16771 .unwrap(),
16772 ),
16773 )
16774 .await
16775 }
16776 "/v1/responses" => {
16777 responses_api::responses(
16778 State(st),
16779 axum::http::HeaderMap::new(),
16780 None,
16781 axum::body::Bytes::from(
16782 serde_json::json!({
16783 "model": "m", "max_output_tokens": 8, "input": "t",
16784 "reasoning": {"effort": effort}})
16785 .to_string(),
16786 ),
16787 )
16788 .await
16789 }
16790 "/v1/messages" => {
16791 anthropic::messages(
16792 State(st),
16793 axum::http::HeaderMap::new(),
16794 None,
16795 axum::body::Bytes::from(
16796 serde_json::json!({
16797 "model": "m", "max_tokens": 8,
16798 "messages": [{"role": "user", "content": "t"}],
16799 "output_config": {"effort": effort}})
16800 .to_string(),
16801 ),
16802 )
16803 .await
16804 }
16805 other => panic!("unknown surface {other}"),
16806 }
16807 };
16808 const SURFACES: [&str; 3] = ["/v1/chat/completions", "/v1/responses", "/v1/messages"];
16809
16810 for (effort, want_think, want_level) in [
16813 ("none", ThinkMode::NoThink, Some("low")),
16814 ("minimal", ThinkMode::NoThink, Some("low")),
16815 ("low", ThinkMode::Think, Some("low")),
16816 ("medium", ThinkMode::Think, Some("medium")),
16817 ("high", ThinkMode::Think, Some("high")),
16818 ("xhigh", ThinkMode::Think, Some("high")),
16820 ] {
16821 for surface in SURFACES {
16822 let resp = send(st.clone(), surface, effort).await;
16823 assert_eq!(
16824 resp.status(),
16825 StatusCode::OK,
16826 "{surface} rejected effort {effort:?} — the surfaces' allowlists \
16827 diverged again (issue #31)"
16828 );
16829 let saw = saw_rx
16830 .recv_timeout(std::time::Duration::from_secs(10))
16831 .unwrap_or_else(|_| {
16832 panic!("{surface}: effort {effort:?} request never reached the worker")
16833 });
16834 assert_eq!(
16835 (saw.think, saw.reasoning_effort.as_deref()),
16836 (want_think, want_level),
16837 "{surface} resolved effort {effort:?} to a DIFFERENT worker-truth \
16838 reasoning surface — the parameter was dropped or remapped before \
16839 parse_think (issue #31 regression)"
16840 );
16841 }
16842 }
16843
16844 for effort in ["bogus", "banana", ""] {
16847 for surface in SURFACES {
16848 let resp = send(st.clone(), surface, effort).await;
16849 assert_eq!(
16850 resp.status(),
16851 StatusCode::BAD_REQUEST,
16852 "{surface} accepted effort {effort:?} — silent-accept regression \
16853 (issue #31: the value never reached parse_think's allowlist)"
16854 );
16855 let body = axum::body::to_bytes(resp.into_body(), 1 << 20)
16857 .await
16858 .unwrap();
16859 let v: serde_json::Value = serde_json::from_slice(&body)
16860 .unwrap_or_else(|_| panic!("{surface}: non-JSON 400 body for {effort:?}"));
16861 match surface {
16862 "/v1/messages" => {
16863 assert_eq!(v["type"], "error", "{surface} error envelope");
16864 assert_eq!(
16865 v["error"]["type"], "invalid_request_error",
16866 "{surface} error type"
16867 );
16868 }
16869 _ => {
16870 assert!(
16871 v["error"]["message"].is_string(),
16872 "{surface} OpenAI-shaped error body: {v}"
16873 );
16874 }
16875 }
16876 }
16877 }
16878
16879 let resp = anthropic::messages(
16883 State(st.clone()),
16884 axum::http::HeaderMap::new(),
16885 None,
16886 axum::body::Bytes::from(
16887 serde_json::json!({
16888 "model": "m", "max_tokens": 8,
16889 "messages": [{"role": "user", "content": "t"}],
16890 "thinking": {"type": "enabled"},
16891 "output_config": {"effort": "none"}})
16892 .to_string(),
16893 ),
16894 )
16895 .await;
16896 assert_eq!(resp.status(), StatusCode::OK);
16897 let saw = saw_rx
16898 .recv_timeout(std::time::Duration::from_secs(10))
16899 .expect("thinking+effort request never reached the worker");
16900 assert_eq!(
16901 saw.think,
16902 ThinkMode::Think,
16903 "thinking.type (the documented Anthropic lever) must win the switch over \
16904 output_config.effort"
16905 );
16906 let resp = anthropic::messages(
16907 State(st.clone()),
16908 axum::http::HeaderMap::new(),
16909 None,
16910 axum::body::Bytes::from(
16911 serde_json::json!({
16912 "model": "m", "max_tokens": 8,
16913 "messages": [{"role": "user", "content": "t"}],
16914 "thinking": {"type": "enabled"},
16915 "output_config": {"effort": "banana"}})
16916 .to_string(),
16917 ),
16918 )
16919 .await;
16920 assert_eq!(
16921 resp.status(),
16922 StatusCode::BAD_REQUEST,
16923 "an invalid effort must 400 even next to an explicit thinking.type — \
16924 precedence must not re-open the silent-accept hole"
16925 );
16926 }
16927
16928 #[test]
16929 fn vendor_sampling_defaults_are_boot_validated() {
16930 let parsed = OpenRouterMetadataFile::from_toml(
16933 r#"
16934[models.g]
16935default_temperature = 1.0
16936default_top_p = 0.95
16937default_top_k = 64
16938default_min_p = 0.0
16939default_presence_penalty = 0.0
16940default_frequency_penalty = 0.0
16941default_repetition_penalty = 1.0
16942"#,
16943 )
16944 .unwrap();
16945 let g = parsed.get("g").unwrap();
16946 assert_eq!(g.default_temperature, Some(1.0));
16947 assert_eq!(g.default_top_p, Some(0.95));
16948 assert_eq!(g.default_top_k, Some(64));
16949
16950 let err = OpenRouterMetadataFile::from_toml(
16954 r#"
16955[models.g]
16956default_temperature = 0.0
16957"#,
16958 )
16959 .unwrap_err();
16960 assert!(err.contains("default_temperature"), "{err}");
16961 assert!(
16962 err.contains("greedy"),
16963 "the refusal must say WHY a zero default is refused: {err}"
16964 );
16965
16966 for bad in [
16967 "default_temperature = 2.5",
16968 "default_temperature = -1.0",
16969 "default_top_p = 0.0",
16970 "default_top_p = 1.5",
16971 "default_min_p = 1.0",
16972 "default_min_p = -0.1",
16973 "default_presence_penalty = 3.0",
16974 "default_frequency_penalty = -2.5",
16975 "default_repetition_penalty = 0.0",
16976 ] {
16977 let err =
16978 OpenRouterMetadataFile::from_toml(&format!("[models.g]\n{bad}\n")).unwrap_err();
16979 let key = bad.split(' ').next().unwrap();
16980 assert!(err.contains(key), "{bad} must be refused by name: {err}");
16981 }
16982
16983 let err = OpenRouterMetadataFile::from_toml(
16987 r#"
16988[models.g]
16989default_temperture = 1.0
16990"#,
16991 )
16992 .unwrap_err();
16993 assert!(
16994 err.contains("unknown field"),
16995 "an unknown key must be fatal, which is what makes binary-first ordering \
16996 mandatory: {err}"
16997 );
16998 }
16999
17000 #[test]
17001 fn non_thinking_sampling_arm_is_boot_validated() {
17002 let parsed = OpenRouterMetadataFile::from_toml(
17006 r#"
17007[models.q]
17008default_temperature = 1.0
17009default_top_p = 0.95
17010default_top_k = 20
17011
17012[models.q.non_thinking_sampling]
17013temperature = 0.7
17014top_p = 0.8
17015top_k = 20
17016presence_penalty = 1.5
17017"#,
17018 )
17019 .unwrap();
17020 let arm = parsed
17021 .get("q")
17022 .unwrap()
17023 .non_thinking_sampling
17024 .as_ref()
17025 .unwrap();
17026 assert_eq!(arm.temperature, Some(0.7));
17027 assert_eq!(arm.top_p, Some(0.8));
17028 assert_eq!(arm.top_k, Some(20));
17029 assert_eq!(arm.presence_penalty, Some(1.5));
17030 assert_eq!(
17031 arm.min_p, None,
17032 "undeclared arm fields stay undeclared, never invented"
17033 );
17034
17035 let err = OpenRouterMetadataFile::from_toml(
17039 r#"
17040[models.q]
17041[models.q.non_thinking_sampling]
17042temperature = 0.0
17043"#,
17044 )
17045 .unwrap_err();
17046 assert!(err.contains("non_thinking_sampling.temperature"), "{err}");
17047 assert!(err.contains("greedy"), "{err}");
17048
17049 let err = OpenRouterMetadataFile::from_toml(
17052 r#"
17053[models.q]
17054[models.q.non_thinking_sampling]
17055"#,
17056 )
17057 .unwrap_err();
17058 assert!(err.contains("non_thinking_sampling"), "{err}");
17059 assert!(err.contains("declare"), "{err}");
17060
17061 for bad in [
17063 "temperature = 2.5",
17064 "top_p = 0.0",
17065 "top_p = 1.5",
17066 "min_p = 1.0",
17067 "presence_penalty = 3.0",
17068 "frequency_penalty = -2.5",
17069 "repetition_penalty = 0.0",
17070 ] {
17071 let err = OpenRouterMetadataFile::from_toml(&format!(
17072 "[models.q]\n[models.q.non_thinking_sampling]\n{bad}\n"
17073 ))
17074 .unwrap_err();
17075 let key = bad.split(' ').next().unwrap();
17076 assert!(
17077 err.contains(&format!("non_thinking_sampling.{key}")),
17078 "the refusal for {bad:?} must name the nested key: {err}"
17079 );
17080 }
17081
17082 let err = OpenRouterMetadataFile::from_toml(
17086 r#"
17087[models.q]
17088[models.q.non_thinking_sampling]
17089temperture = 0.7
17090"#,
17091 )
17092 .unwrap_err();
17093 assert!(err.contains("unknown field"), "{err}");
17094 }
17095
17096 fn qwen38_non_thinking_defaults() -> SamplingDefaults {
17101 SamplingDefaults {
17102 temperature: Some(0.7),
17103 top_p: Some(0.8),
17104 top_k: Some(20),
17105 presence_penalty: Some(1.5),
17106 ..Default::default()
17107 }
17108 }
17109
17110 fn qwen38_two_arm_defaults() -> ModelSamplingDefaults {
17111 ModelSamplingDefaults {
17112 thinking: qwen38_vendor_defaults(),
17113 non_thinking: Some(qwen38_non_thinking_defaults()),
17114 }
17115 }
17116
17117 fn qwen38_caps() -> ModelCaps {
17121 ModelCaps {
17122 chat_ok: true,
17123 qwen_think: true,
17124 think_switch: true,
17125 ..Default::default()
17126 }
17127 }
17128
17129 fn sampler_key(c: &SamplerConfig) -> (f32, f32, usize, f32, f32, f32, f32, usize, u64) {
17132 (
17133 c.temperature,
17134 c.top_p,
17135 c.top_k,
17136 c.min_p,
17137 c.penalty_present,
17138 c.penalty_freq,
17139 c.penalty_repeat,
17140 c.penalty_last_n,
17141 c.seed,
17142 )
17143 }
17144
17145 fn build_with_arms(
17146 defaults: &ModelSamplingDefaults,
17147 caps: &ModelCaps,
17148 default_effort: Option<&str>,
17149 extra: serde_json::Value,
17150 ) -> Request {
17151 let mut body = serde_json::json!({
17152 "model": "m",
17153 "messages": [{"role": "user", "content": "task"}],
17154 "seed": 3
17156 });
17157 body.as_object_mut()
17158 .unwrap()
17159 .extend(extra.as_object().unwrap().clone());
17160 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
17161 let (tx, _rx) = worker::event_channel();
17162 build_chat_request_with_trace(
17163 req,
17164 Some(caps),
17165 tx,
17166 lanes::Lane::Interactive,
17167 None,
17168 None,
17169 default_effort,
17170 defaults,
17171 )
17172 .unwrap()
17173 .request
17174 }
17175
17176 #[test]
17177 fn resolved_thinking_mode_picks_the_vendor_sampling_arm() {
17178 let two_arm = qwen38_two_arm_defaults();
17183 let single_arm = ModelSamplingDefaults::single(qwen38_vendor_defaults());
17184 let caps = qwen38_caps();
17185
17186 let off_spellings = [
17188 serde_json::json!({"reasoning_effort": "none"}),
17189 serde_json::json!({"enable_thinking": false}),
17190 serde_json::json!({"chat_template_kwargs": {"enable_thinking": false}}),
17191 serde_json::json!({"reasoning": {"enabled": false}}),
17192 ];
17193 for extra in &off_spellings {
17194 let r = build_with_arms(&two_arm, &caps, None, extra.clone());
17195 assert_eq!(r.think, ThinkMode::NoThink, "{extra}");
17196 let c = &r.sampler_cfg;
17197 assert_eq!(c.temperature, 0.7, "{extra}: non-thinking card temperature");
17198 assert_eq!(c.top_p, 0.8, "{extra}: non-thinking card top_p");
17199 assert_eq!(c.top_k, 20, "{extra}: non-thinking card top_k");
17200 assert_eq!(
17201 c.penalty_present, 1.5,
17202 "{extra}: non-thinking presence_penalty"
17203 );
17204 assert_eq!(
17205 c.penalty_last_n,
17206 memra_engine::spec::PEN_WINDOW_MAX,
17207 "{extra}: the arm's presence penalty uses the cross-path history window"
17208 );
17209 assert_eq!(
17210 c.min_p, 0.0,
17211 "{extra}: the arm recommends no min_p — API standard, never the other arm's"
17212 );
17213
17214 let s = build_with_arms(&single_arm, &caps, None, extra.clone());
17217 assert_eq!(s.think, ThinkMode::NoThink, "{extra}");
17218 assert_eq!(s.sampler_cfg.temperature, 1.0, "{extra}: single-arm model");
17219 assert_eq!(s.sampler_cfg.top_p, 0.95, "{extra}: single-arm model");
17220 assert_eq!(
17221 s.sampler_cfg.penalty_present, 0.0,
17222 "{extra}: single-arm model"
17223 );
17224 }
17225
17226 for extra in [
17229 serde_json::json!({}),
17230 serde_json::json!({"enable_thinking": true}),
17231 serde_json::json!({"reasoning_effort": "high"}),
17232 serde_json::json!({"reasoning": {"enabled": true}}),
17233 ] {
17234 for defaults in [&two_arm, &single_arm] {
17235 let c = build_with_arms(defaults, &caps, None, extra.clone()).sampler_cfg;
17236 assert_eq!(c.temperature, 1.0, "{extra}: thinking card temperature");
17237 assert_eq!(c.top_p, 0.95, "{extra}: thinking card top_p");
17238 assert_eq!(c.top_k, 20, "{extra}: thinking card top_k");
17239 assert_eq!(
17240 c.penalty_present, 0.0,
17241 "{extra}: thinking arm has no presence"
17242 );
17243 }
17244 }
17245
17246 let c = build_with_arms(&two_arm, &caps, Some("none"), serde_json::json!({})).sampler_cfg;
17249 assert_eq!(
17250 c.temperature, 0.7,
17251 "deployment-default off = non-thinking arm"
17252 );
17253 let c = build_with_arms(
17255 &two_arm,
17256 &caps,
17257 Some("none"),
17258 serde_json::json!({"enable_thinking": true}),
17259 )
17260 .sampler_cfg;
17261 assert_eq!(
17262 c.temperature, 1.0,
17263 "explicit ON beats the deployment default"
17264 );
17265
17266 let c = build_with_arms(
17268 &two_arm,
17269 &caps,
17270 None,
17271 serde_json::json!({"enable_thinking": false, "temperature": 0.55}),
17272 )
17273 .sampler_cfg;
17274 assert_eq!(c.temperature, 0.55, "explicit temperature survives the arm");
17275 assert_eq!(c.top_p, 0.8, "unset top_p still takes the non-thinking arm");
17276 let c = build_with_arms(
17277 &two_arm,
17278 &caps,
17279 None,
17280 serde_json::json!({
17281 "reasoning_effort": "none", "top_p": 0.99, "presence_penalty": 0.0}),
17282 )
17283 .sampler_cfg;
17284 assert_eq!(c.top_p, 0.99, "explicit top_p wins");
17285 assert_eq!(
17286 c.penalty_present, 0.0,
17287 "an explicit presence_penalty 0.0 wins over the arm's 1.5 — a disabling value \
17288 is a value, not an absence"
17289 );
17290 assert_eq!(
17291 c.penalty_last_n, 0,
17292 "all penalties off => no history window"
17293 );
17294 assert_eq!(c.top_k, 20, "unset top_k still takes the arm");
17295
17296 let c = build_with_arms(
17299 &two_arm,
17300 &caps,
17301 None,
17302 serde_json::json!({"enable_thinking": false, "temperature": 0}),
17303 )
17304 .sampler_cfg;
17305 assert!(
17306 memra_engine::sampler::Sampler::new(c).is_greedy(),
17307 "explicit temperature 0 must stay greedy on the non-thinking arm"
17308 );
17309
17310 let c = build_with_arms(
17313 &single_arm,
17314 &caps,
17315 None,
17316 serde_json::json!({"enable_thinking": false, "temperature": 0.55}),
17317 )
17318 .sampler_cfg;
17319 assert_eq!(c.temperature, 0.55);
17320 assert_eq!(
17321 c.top_p, 0.95,
17322 "single-arm model: unset top_p takes its one arm"
17323 );
17324 }
17325
17326 #[test]
17327 fn sampling_arms_never_blend_field_by_field() {
17328 let parsed = OpenRouterMetadataFile::from_toml(
17333 r#"
17334[models.m]
17335default_temperature = 1.0
17336default_min_p = 0.05
17337
17338[models.m.non_thinking_sampling]
17339temperature = 0.6
17340"#,
17341 )
17342 .unwrap();
17343 let caps = ModelCaps {
17344 chat_temperature_default: Some(0.5),
17345 chat_top_p_default: Some(0.9),
17346 ..Default::default()
17347 };
17348 let d = ModelSamplingDefaults::resolve(parsed.get("m"), Some(&caps));
17349 let client = ClientSampling {
17350 seed: Some(1),
17351 ..Default::default()
17352 };
17353
17354 let off = resolve_sampler_config(client, d.for_mode(ThinkMode::NoThink));
17355 assert_eq!(off.temperature, 0.6, "the arm's own field applies");
17356 assert_eq!(
17357 off.min_p, 0.0,
17358 "min_p undeclared on the arm = API standard, NOT the thinking arm's 0.05"
17359 );
17360 assert_eq!(
17361 off.top_p, 1.0,
17362 "top_p undeclared on the arm = API standard, NOT the arch cap's 0.9"
17363 );
17364
17365 for mode in [ThinkMode::Default, ThinkMode::Think] {
17367 let on = resolve_sampler_config(client, d.for_mode(mode));
17368 assert_eq!(on.temperature, 1.0);
17369 assert_eq!(on.min_p, 0.05);
17370 assert_eq!(on.top_p, 0.9, "primary arm keeps the arch-cap fallback");
17371 }
17372 }
17373
17374 #[test]
17375 fn single_arm_models_and_thinking_on_requests_match_the_pre_arm_law_exactly() {
17376 let caps = qwen38_caps();
17385 let single_arm = ModelSamplingDefaults::single(qwen38_vendor_defaults());
17386 let two_arm = qwen38_two_arm_defaults();
17387
17388 let bodies = [
17389 serde_json::json!({}),
17390 serde_json::json!({"enable_thinking": true}),
17391 serde_json::json!({"reasoning_effort": "high"}),
17392 serde_json::json!({"reasoning_effort": "none"}),
17393 serde_json::json!({"enable_thinking": false}),
17394 serde_json::json!({"chat_template_kwargs": {"enable_thinking": false}}),
17395 serde_json::json!({"temperature": 0.3, "top_p": 0.5}),
17396 serde_json::json!({"enable_thinking": false, "temperature": 0}),
17397 ];
17398 for extra in &bodies {
17399 let r = build_with_arms(&single_arm, &caps, None, extra.clone());
17401 let mut client = ClientSampling {
17402 seed: Some(3),
17403 ..Default::default()
17404 };
17405 if let Some(t) = extra.get("temperature").and_then(|v| v.as_f64()) {
17406 client.temperature = Some(t as f32);
17407 }
17408 if let Some(p) = extra.get("top_p").and_then(|v| v.as_f64()) {
17409 client.top_p = Some(p as f32);
17410 }
17411 let pre_arm = resolve_sampler_config(client, &qwen38_vendor_defaults());
17412 assert_eq!(
17413 sampler_key(&r.sampler_cfg),
17414 sampler_key(&pre_arm),
17415 "{extra}: single-arm model diverged from the pre-arm resolution law"
17416 );
17417
17418 if r.think != ThinkMode::NoThink {
17421 let t = build_with_arms(&two_arm, &caps, None, extra.clone());
17422 assert_eq!(t.think, r.think, "{extra}");
17423 assert_eq!(t.reasoning_effort, r.reasoning_effort, "{extra}");
17424 assert_eq!(
17425 sampler_key(&t.sampler_cfg),
17426 sampler_key(&r.sampler_cfg),
17427 "{extra}: a thinking-on request must not feel the non-thinking arm"
17428 );
17429 }
17430 }
17431 }
17432
17433 #[test]
17434 fn constraint_forced_nothink_takes_the_non_thinking_arm() {
17435 let r = build_with_arms(
17441 &qwen38_two_arm_defaults(),
17442 &qwen38_caps(),
17443 None,
17444 serde_json::json!({"response_format": {"type": "json_object"}}),
17445 );
17446 assert_eq!(
17447 r.think,
17448 ThinkMode::NoThink,
17449 "constraint forces the switch off"
17450 );
17451 assert_eq!(
17452 r.sampler_cfg.temperature, 0.7,
17453 "and the arm follows the real mode"
17454 );
17455 assert_eq!(r.sampler_cfg.penalty_present, 1.5);
17456 }
17457
17458 #[test]
17459 fn metadata_sampling_defaults_outrank_arch_caps_but_never_the_client() {
17460 let caps = ModelCaps {
17465 chat_temperature_default: Some(0.5),
17466 chat_top_p_default: Some(0.9),
17467 chat_ok: true,
17468 ..Default::default()
17469 };
17470 let metadata = OpenRouterModelMetadata {
17471 default_temperature: Some(1.0),
17472 default_top_p: Some(0.95),
17473 default_top_k: Some(64),
17474 ..Default::default()
17475 };
17476
17477 let caps_only = SamplingDefaults::resolve(None, Some(&caps));
17478 assert_eq!(caps_only.temperature, Some(0.5), "arch cap is the fallback");
17479 assert_eq!(caps_only.top_p, Some(0.9));
17480 assert_eq!(caps_only.top_k, None, "caps declare no top_k");
17481
17482 let both = SamplingDefaults::resolve(Some(&metadata), Some(&caps));
17483 assert_eq!(
17484 both.temperature,
17485 Some(1.0),
17486 "metadata outranks the arch cap"
17487 );
17488 assert_eq!(both.top_p, Some(0.95));
17489 assert_eq!(both.top_k, Some(64));
17490
17491 let partial = SamplingDefaults::resolve(
17493 Some(&OpenRouterModelMetadata {
17494 default_temperature: Some(0.7),
17495 ..Default::default()
17496 }),
17497 Some(&caps),
17498 );
17499 assert_eq!(partial.temperature, Some(0.7));
17500 assert_eq!(
17501 partial.top_p,
17502 Some(0.9),
17503 "an undeclared metadata field must fall through to the cap, not to 1.0"
17504 );
17505
17506 assert_eq!(
17508 SamplingDefaults::resolve(None, None),
17509 SamplingDefaults::default()
17510 );
17511 }
17512
17513 #[test]
17514 fn vendor_defaults_leave_the_pure_temp_sampled_spec_regime() {
17515 let resolved = |d: &SamplingDefaults| {
17528 resolve_sampler_config(
17529 ClientSampling {
17530 seed: Some(1),
17531 ..Default::default()
17532 },
17533 d,
17534 )
17535 };
17536
17537 assert!(
17539 memra_engine::sampler::Sampler::new(resolved(&SamplingDefaults::default()))
17540 .is_spec_sampling(),
17541 "the API-standard default must stay in the fast pure-temp regime"
17542 );
17543
17544 for (name, d) in [
17545 ("qwen/qwen3.8-27b", qwen38_vendor_defaults()),
17546 ("google/gemma-4-31b-it", gemma4_vendor_defaults()),
17547 ] {
17548 let sampler = memra_engine::sampler::Sampler::new(resolved(&d));
17549 assert!(
17550 !sampler.is_greedy(),
17551 "{name}: vendor default must not be greedy"
17552 );
17553 assert!(
17554 !sampler.is_spec_sampling(),
17555 "{name}: vendor top_p/top_k DO leave the pure-temp regime — if this ever \
17556 starts passing, either the vendor numbers changed or the in-graph draft \
17557 learned filters, and the perf note in docs/SERVING.md needs revisiting"
17558 );
17559 }
17560
17561 let opted_out = resolve_sampler_config(
17563 ClientSampling {
17564 top_p: Some(1.0),
17565 top_k: Some(0),
17566 seed: Some(1),
17567 ..Default::default()
17568 },
17569 &qwen38_vendor_defaults(),
17570 );
17571 assert!(
17572 memra_engine::sampler::Sampler::new(opted_out).is_spec_sampling(),
17573 "explicitly disabling the filters must restore the pure-temp regime"
17574 );
17575 }
17576
17577 #[test]
17578 fn omitted_seed_is_fresh_entropy_not_a_pinned_zero() {
17579 let comp_seed = |body: serde_json::Value| {
17586 let req: CompletionReq = serde_json::from_value(body).unwrap();
17587 let (tx, _rx) = worker::event_channel();
17588 build_request(&req, tx, lanes::Lane::Interactive, None)
17589 .sampler_cfg
17590 .seed
17591 };
17592 let chat_seed = |body: serde_json::Value| {
17593 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
17594 let (tx, _rx) = worker::event_channel();
17595 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
17596 .unwrap()
17597 .request
17598 .sampler_cfg
17599 .seed
17600 };
17601
17602 let a = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
17605 let b = comp_seed(serde_json::json!({"model": "m", "prompt": "t"}));
17606 let c = chat_seed(serde_json::json!({
17607 "model": "m", "messages": [{"role": "user", "content": "t"}]}));
17608 assert_ne!(
17609 a, 0,
17610 "omitted seed must not be the pinned 0 that caused the loop"
17611 );
17612 assert_ne!(b, 0);
17613 assert_ne!(c, 0);
17614 assert_ne!(
17615 a, b,
17616 "two seed-omitting requests must get DIFFERENT streams"
17617 );
17618 assert_ne!(a, c);
17619
17620 assert_eq!(
17623 comp_seed(serde_json::json!({
17624 "model": "m", "prompt": "t", "seed": 0})),
17625 0,
17626 "explicit seed 0 must stay 0 — the determinism gates depend on it"
17627 );
17628 assert_eq!(
17629 comp_seed(serde_json::json!({
17630 "model": "m", "prompt": "t", "seed": 12345})),
17631 12345
17632 );
17633 assert_eq!(
17634 chat_seed(serde_json::json!({
17635 "model": "m", "messages": [{"role": "user", "content": "t"}],
17636 "seed": 777})),
17637 777
17638 );
17639 assert_eq!(
17641 comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42})),
17642 comp_seed(serde_json::json!({"model": "m", "prompt": "t", "seed": 42}))
17643 );
17644
17645 let seeds: std::collections::HashSet<u64> = (0..256).map(|_| fresh_seed()).collect();
17648 assert_eq!(
17649 seeds.len(),
17650 256,
17651 "fresh_seed must not collide across rapid calls"
17652 );
17653 assert!(!seeds.contains(&0));
17654 }
17655
17656 #[test]
17657 fn response_format_builds_grammar_only_when_present() {
17658 let mk = |rf: Option<serde_json::Value>| {
17662 let mut body = serde_json::json!({
17663 "model": "m", "messages": [{"role": "user", "content": "t"}]});
17664 if let Some(rf) = rf {
17665 body["response_format"] = rf;
17666 }
17667 let req: ChatCompletionReq = serde_json::from_value(body).unwrap();
17668 let (tx, _rx) = worker::event_channel();
17669 build_chat_request(req, None, tx, lanes::Lane::Interactive, None)
17670 };
17671 assert!(mk(None).unwrap().request.grammar.is_none());
17672 assert!(
17673 mk(Some(serde_json::json!({"type": "text"})))
17674 .unwrap()
17675 .request
17676 .grammar
17677 .is_none()
17678 );
17679 assert!(matches!(
17680 mk(Some(serde_json::json!({"type": "json_object"})))
17681 .unwrap()
17682 .request
17683 .grammar,
17684 Some(constrained::GrammarSpec::JsonObject)
17685 ));
17686 assert!(matches!(
17687 mk(Some(serde_json::json!({"type": "json_schema",
17688 "json_schema": {"schema": {"type": "object"}}})))
17689 .unwrap()
17690 .request
17691 .grammar,
17692 Some(constrained::GrammarSpec::JsonSchema(_))
17693 ));
17694 assert!(mk(Some(serde_json::json!({"type": "yaml"}))).is_err());
17696 }
17697
17698 #[test]
17708 fn response_format_think_table_switch_postthink_refusal() {
17709 let mk = |caps: &ModelCaps| {
17710 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
17711 "model": "m", "messages": [{"role": "user", "content": "t"}],
17712 "response_format": {"type": "json_object"}}))
17713 .unwrap();
17714 let (tx, _rx) = worker::event_channel();
17715 build_chat_request(req, Some(caps), tx, lanes::Lane::Interactive, None)
17716 };
17717 let switch = ModelCaps {
17719 chat_ok: true,
17720 qwen_think: true,
17721 think_switch: true,
17722 ..Default::default()
17723 };
17724 let plan = mk(&switch).unwrap();
17725 assert_eq!(
17726 plan.request.think,
17727 memra_tokenizer::chat::ThinkMode::NoThink,
17728 "switch-carrying template must keep the grammar-from-token-1 path"
17729 );
17730 assert!(plan.request.grammar.is_some());
17731
17732 let postthink = ModelCaps {
17734 chat_ok: true,
17735 qwen_think: true,
17736 think_switch: false,
17737 think_close: vec![128799],
17738 ..Default::default()
17739 };
17740 let plan = mk(&postthink).unwrap();
17741 assert_ne!(
17742 plan.request.think,
17743 memra_tokenizer::chat::ThinkMode::NoThink,
17744 "post-think constrained request must keep the think channel ON"
17745 );
17746 assert!(plan.request.grammar.is_some());
17747
17748 let no_contract = ModelCaps {
17750 chat_ok: true,
17751 qwen_think: true,
17752 think_switch: false,
17753 think_close: Vec::new(),
17754 ..Default::default()
17755 };
17756 let err = match mk(&no_contract) {
17757 Err(err) => err,
17758 Ok(_) => panic!("think-forced template with no close contract must refuse"),
17759 };
17760 assert!(
17761 err.contains("think-close"),
17762 "refusal must name the missing close contract: {err}"
17763 );
17764 }
17765
17766 #[test]
17767 fn unsupported_semantic_params_are_named_rejections() {
17768 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
17770 "model": "m", "messages": [{"role": "user", "content": "t"}],
17771 "response_format": {"type": "json_object"}
17772 }))
17773 .unwrap();
17774 assert!(req.response_format.is_some());
17775 let req: ChatCompletionReq = serde_json::from_value(serde_json::json!({
17776 "model": "m", "messages": [{"role": "user", "content": "t"}],
17777 "response_format": {"type": "text"}, "logprobs": false, "n": 1,
17778 "user": "u-1", "stream_options": {"include_usage": true}
17779 }))
17780 .unwrap();
17781 assert_eq!(req.response_format.as_ref().unwrap()["type"], "text");
17783 assert_eq!(req.logprobs.as_ref().unwrap().as_bool(), Some(false));
17784 assert_eq!(req.n, Some(1));
17785 assert!(reject_unsupported(&[("logit_bias", false, "")]).is_ok());
17787 let (msg, param) = reject_unsupported(&[("logit_bias", true, " (why)")]).unwrap_err();
17788 assert_eq!(param, "logit_bias");
17789 assert_eq!(msg, "logit_bias is not supported (why)");
17790 }
17791
17792 #[test]
17793 fn completions_accept_openai_stop_forms() {
17794 for (value, expected) in [
17795 (serde_json::json!("Problem:"), vec!["Problem:"]),
17796 (
17797 serde_json::json!(["Question:", "Problem:"]),
17798 vec!["Question:", "Problem:"],
17799 ),
17800 (serde_json::Value::Null, Vec::<&str>::new()),
17801 ] {
17802 let req: CompletionReq = serde_json::from_value(serde_json::json!({
17803 "model": "plain_quant", "prompt": "task", "stop": value
17804 }))
17805 .unwrap();
17806 assert_eq!(req.stop.into_vec(), expected);
17807 }
17808 }
17809
17810 fn fake_worker_state() -> AppState {
17817 fake_worker_state_with_steps(1, std::time::Duration::ZERO)
17818 }
17819
17820 fn fake_worker_state_with_steps(steps: usize, step_delay: std::time::Duration) -> AppState {
17821 fake_worker_state_full(steps, step_delay, HashMap::new(), None)
17822 }
17823
17824 struct WorkerSaw {
17829 sampler_cfg: SamplerConfig,
17830 think: ThinkMode,
17831 reasoning_effort: Option<String>,
17832 }
17833
17834 fn fake_worker_state_full(
17841 steps: usize,
17842 step_delay: std::time::Duration,
17843 caps: HashMap<String, ModelCaps>,
17844 saw_tx: Option<std::sync::mpsc::Sender<WorkerSaw>>,
17845 ) -> AppState {
17846 let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
17847 let health = health::WorkerHealth::new();
17848 let h = health.clone();
17849 std::thread::spawn(move || {
17850 h.mark_ready();
17851 while let Ok(Cmd::Generate(mut req)) = cmd_rx.recv() {
17852 if let Some(tx) = &saw_tx {
17853 let _ = tx.send(WorkerSaw {
17854 sampler_cfg: req.sampler_cfg.clone(),
17855 think: req.think,
17856 reasoning_effort: req.reasoning_effort.clone(),
17857 });
17858 }
17859 worker::release_pending_admit();
17863 worker::release_admission_reservation(req.lane);
17864 h.beat_busy();
17865 if let Some(ready) = req.constraint_ready.take() {
17866 let _ = ready.send(Ok(()));
17867 }
17868 let _ = req.tx.send(Event::PromptUsage {
17869 n_prompt: 1,
17870 n_cached: 0,
17871 });
17872 if let Some(spec) = req.capture.as_ref() {
17877 let _ = req.tx.send(Event::PromptCapture {
17878 hidden: spec.hidden.then(|| vec![1.0, 0.0]),
17879 logits: if spec.logit_pieces.is_empty() {
17880 Vec::new()
17881 } else {
17882 vec![2.0, 0.0]
17883 },
17884 });
17885 }
17886 for step in 0..steps {
17887 h.beat_busy();
17888 let text = if steps == 1 { "ok" } else { "x" };
17889 let _ = req.tx.send(Event::Token {
17890 id: step as u32 + 1,
17891 text: text.into(),
17892 });
17893 if !step_delay.is_zero() {
17894 std::thread::sleep(step_delay);
17895 }
17896 }
17897 let _ = req.tx.send(Event::Done {
17898 stop_reason: "Eos".into(),
17899 n_tokens: steps,
17900 n_prompt: 1,
17901 n_cached: 0,
17902 elapsed_s: 0.01,
17903 spec: None,
17904 });
17905 h.set_phase(health::PHASE_IDLE);
17906 }
17907 });
17908 for _ in 0..2000 {
17911 if health.live().is_ok() {
17912 break;
17913 }
17914 std::thread::sleep(std::time::Duration::from_millis(1));
17915 }
17916 AppState {
17917 cmd_tx,
17918 models: Arc::new(vec!["m".into()]),
17919 caps: Arc::new(caps),
17920 openrouter_metadata: Arc::new(HashMap::new()),
17921 provider_metadata: Arc::new(None),
17922 metering: None,
17923
17924 budget_tokenizers: None,
17925 api_auth: ApiAuth::default(),
17926 metrics_auth: MetricsAuth::default(),
17927 metrics: SharedMetrics::default(),
17928 inflight: Arc::new(Default::default()),
17929 tenant_inflight: Arc::new(Default::default()),
17930 health,
17931 bg: None,
17932 }
17933 }
17934
17935 #[tokio::test]
17936 #[allow(clippy::await_holding_lock)] async fn deep_schema_fails_while_normal_decode_keeps_stepping() {
17938 let _l = drain_lock();
17939 let st = fake_worker_state_with_steps(64, std::time::Duration::from_millis(5));
17940 let normal_state = st.clone();
17941 let normal = tokio::spawn(async move {
17942 chat_completions(
17943 State(normal_state),
17944 axum::http::HeaderMap::new(),
17945 None,
17946 Json(
17947 serde_json::from_value(serde_json::json!({
17948 "model": "m",
17949 "messages": [{"role": "user", "content": "keep decoding"}],
17950 }))
17951 .unwrap(),
17952 ),
17953 )
17954 .await
17955 });
17956 tokio::time::sleep(std::time::Duration::from_millis(15)).await;
17957
17958 let mut deep = serde_json::json!({"type": "string"});
17959 for _ in 0..(constrained::MAX_SCHEMA_DEPTH / 2 + 1) {
17960 deep = serde_json::json!({"allOf": [deep]});
17961 }
17962 let bad = chat_completions(
17963 State(st.clone()),
17964 axum::http::HeaderMap::new(),
17965 None,
17966 Json(
17967 serde_json::from_value(serde_json::json!({
17968 "model": "m",
17969 "messages": [{"role": "user", "content": "bad schema"}],
17970 "response_format": {
17971 "type": "json_schema",
17972 "json_schema": {"schema": deep},
17973 },
17974 }))
17975 .unwrap(),
17976 ),
17977 )
17978 .await;
17979 assert_eq!(bad.status(), StatusCode::BAD_REQUEST);
17980 assert_eq!(bad.headers().get("x-should-retry").unwrap(), "false");
17981 let bytes = axum::body::to_bytes(bad.into_body(), usize::MAX)
17982 .await
17983 .unwrap();
17984 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
17985 assert!(
17986 payload["error"]["message"]
17987 .as_str()
17988 .unwrap()
17989 .contains("maximum nesting depth")
17990 );
17991 assert!(
17992 !normal.is_finished(),
17993 "bad schema stalled or replaced the normal decode"
17994 );
17995
17996 let normal_response = normal.await.unwrap();
17997 assert_eq!(normal_response.status(), StatusCode::OK);
17998 let snapshot = st.health.snapshot();
17999 assert!(
18000 st.health.live().is_ok(),
18001 "normal decode left health stalled"
18002 );
18003 assert!(snapshot.beat_age_ms < snapshot.stall_threshold_ms);
18004 }
18005
18006 #[tokio::test]
18007 #[allow(clippy::await_holding_lock)] async fn valid_response_format_preflight_preserves_generation() {
18009 let _l = drain_lock();
18010 let response = chat_completions(
18011 State(fake_worker_state()),
18012 axum::http::HeaderMap::new(),
18013 None,
18014 Json(
18015 serde_json::from_value(serde_json::json!({
18016 "model": "m",
18017 "messages": [{"role": "user", "content": "valid schema"}],
18018 "response_format": {"type": "json_object"},
18019 }))
18020 .unwrap(),
18021 ),
18022 )
18023 .await;
18024 assert_eq!(response.status(), StatusCode::OK);
18025 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18026 .await
18027 .unwrap();
18028 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
18029 assert_eq!(payload["choices"][0]["message"]["content"], "ok");
18030 }
18031
18032 #[tokio::test]
18033 #[allow(clippy::await_holding_lock)] async fn unknown_model_refuses_model_not_found_before_admission() {
18035 let _l = drain_lock();
18036 let response = chat_completions(
18041 State(fake_worker_state()),
18042 axum::http::HeaderMap::new(),
18043 None,
18044 Json(
18045 serde_json::from_value(serde_json::json!({
18046 "model": "qwen/qwen3.8-27b-typo",
18047 "messages": [{"role": "user", "content": "hi"}],
18048 }))
18049 .unwrap(),
18050 ),
18051 )
18052 .await;
18053 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
18054 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18055 .await
18056 .unwrap();
18057 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
18058 assert_eq!(payload["error"]["code"], "model_not_found");
18059 assert_eq!(payload["error"]["type"], "invalid_request_error");
18060
18061 let response = completions(
18063 State(fake_worker_state()),
18064 axum::http::HeaderMap::new(),
18065 None,
18066 Json(
18067 serde_json::from_value(serde_json::json!({
18068 "model": "nope",
18069 "prompt": "hi",
18070 }))
18071 .unwrap(),
18072 ),
18073 )
18074 .await;
18075 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
18076 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18077 .await
18078 .unwrap();
18079 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
18080 assert_eq!(payload["error"]["code"], "model_not_found");
18081 }
18082
18083 const METRICS_KEY_ACME: &str = "completion-acme-secret";
18084 const METRICS_KEY_BLUE: &str = "completion-blue-secret";
18085
18086 fn multi_key_metrics_state(metrics_token: Option<&str>) -> AppState {
18087 let spec = format!(
18088 "acme:{},blue:{}",
18089 auth::sha256_hex(METRICS_KEY_ACME),
18090 auth::sha256_hex(METRICS_KEY_BLUE),
18091 );
18092 let keyring = Box::leak(Box::new(auth::KeyStore::from_spec(&spec).unwrap()));
18093 let mut st = fake_worker_state();
18094 st.api_auth.keyring = Some(keyring);
18095 st.metrics_auth = MetricsAuth::new(
18096 true,
18097 st.api_auth.configured(),
18098 metrics_token.map(str::to_string),
18099 );
18100 {
18101 let mut metrics = st.metrics.lock().unwrap();
18102 metrics.admitted = 17;
18103 metrics.prompt_tokens_in = 400;
18104 metrics.cached_tokens_in = 60;
18105 metrics.prefix_hits = 2;
18106 metrics.prefix_misses = 3;
18107 metrics.prefix_inserts = 5;
18108 metrics.prefix_evictions = 7;
18109 metrics.prefix_skips_budget = 9;
18110 metrics.prefix_skips_pinned = 10;
18111 metrics.prefix_hit_tokens = 11;
18112 metrics.lcp_hist[4] = 13;
18113 metrics.ns_tokens.insert("t:acme".into(), [100, 40]);
18114 metrics.ns_tokens.insert("t:blue".into(), [300, 20]);
18115 metrics.adsd_suspect_total.insert("t:acme".into(), 1);
18116 metrics.adsd_suspect_total.insert("t:blue".into(), 2);
18117 metrics.prefix_entries = 29;
18118 metrics.prefix_bytes = 31;
18119 metrics.active_sessions = 3;
18120 metrics.queued_requests = 5;
18121 metrics.admission_inflight.insert("m".into(), 4);
18122 metrics
18123 .admission_booked_bytes
18124 .insert("m".into(), 41_000_000);
18125 metrics.continuation_pool_entries = 7;
18126 metrics.spec_pool_entries = 11;
18127 metrics.cuda_driver_free_bytes = 13;
18128 metrics.cuda_pool_reserved_bytes = 17;
18129 metrics.cuda_pool_used_bytes = 19;
18130 metrics.cuda_pool_cached_bytes = 23;
18131 metrics.batch_size_last = 37;
18132 metrics.spec.insert(
18133 "m".into(),
18134 memra_engine::spec::SpecTelemetry {
18135 rounds: 2,
18136 drafted: 6,
18137 accepted: 4,
18138 ..Default::default()
18139 },
18140 );
18141 let mut spec_window = memra_engine::spec::SpecTelemetry {
18142 rounds: 4,
18143 drafted: 12,
18144 accepted: 6,
18145 ..Default::default()
18146 };
18147 spec_window.pos_drafted[..3].copy_from_slice(&[4, 4, 4]);
18148 spec_window.pos_accepted[..3].copy_from_slice(&[3, 2, 1]);
18149 metrics.spec_window.insert("m".into(), spec_window);
18150 metrics.constraint_compiler_fail_closed.insert(
18151 "m".into(),
18152 Arc::new(std::sync::atomic::AtomicBool::new(true)),
18153 );
18154 }
18155 st
18156 }
18157
18158 async fn metrics_json(st: AppState, bearer: &str) -> serde_json::Value {
18159 let mut headers = HeaderMap::new();
18160 headers.insert("authorization", format!("Bearer {bearer}").parse().unwrap());
18161 let response = get_metrics(State(st), headers).await;
18162 assert_eq!(response.status(), StatusCode::OK);
18163 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18164 .await
18165 .unwrap();
18166 serde_json::from_slice(&bytes).unwrap()
18167 }
18168
18169 async fn yield_metrics_json(st: AppState, bearer: &str) -> serde_json::Value {
18170 let mut headers = HeaderMap::new();
18171 headers.insert("authorization", format!("Bearer {bearer}").parse().unwrap());
18172 let response = yield_metrics(State(st), headers).await;
18173 assert_eq!(response.status(), StatusCode::OK);
18174 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18175 .await
18176 .unwrap();
18177 serde_json::from_slice(&bytes).unwrap()
18178 }
18179
18180 #[test]
18181 fn exposed_open_bind_is_refused_before_server_start() {
18182 assert!(validate_bind_security("127.0.0.1:8080", false, false).unwrap());
18183 assert!(validate_bind_security("[::1]:8080", false, false).unwrap());
18184
18185 let err = validate_bind_security("0.0.0.0:8000", false, false).unwrap_err();
18186 assert!(err.contains("refusing unauthenticated non-loopback bind"));
18187 assert!(err.contains("MEMRA_API_KEY"));
18188 assert!(err.contains("MEMRA_ALLOW_OPEN_BIND=1"));
18189 assert!(validate_bind_security("[::]:8000", false, false).is_err());
18190
18191 assert!(!validate_bind_security("0.0.0.0:8000", true, false).unwrap());
18192 assert!(!validate_bind_security("0.0.0.0:8000", false, true).unwrap());
18193 }
18194
18195 #[tokio::test]
18196 async fn keyed_metrics_require_and_accept_api_bearer() {
18197 let mut st = fake_worker_state();
18198 st.api_auth.single_key = Some(Arc::from("completion-secret"));
18199 st.metrics_auth = MetricsAuth::new(true, st.api_auth.configured(), None);
18200
18201 let response = get_metrics(State(st.clone()), HeaderMap::new()).await;
18202 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
18203 let response = yield_metrics(State(st.clone()), HeaderMap::new()).await;
18204 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
18205
18206 let mut headers = HeaderMap::new();
18207 headers.insert("authorization", "Bearer completion-secret".parse().unwrap());
18208 assert_eq!(
18209 get_metrics(State(st.clone()), headers.clone())
18210 .await
18211 .status(),
18212 StatusCode::OK,
18213 );
18214 let body = metrics_json(st.clone(), "completion-secret").await;
18215 assert!(
18216 body.get("admitted").is_some(),
18217 "the legacy single-key domain keeps cumulative counters",
18218 );
18219 assert!(
18220 body.get("active_sessions").is_none(),
18221 "a static completion key is not an operator metrics principal",
18222 );
18223 assert_eq!(
18224 yield_metrics(State(st), headers).await.status(),
18225 StatusCode::OK
18226 );
18227 }
18228
18229 #[tokio::test]
18230 async fn keyring_metrics_bearer_sees_only_its_tenant_rows() {
18231 let st = multi_key_metrics_state(None);
18232 let body = metrics_json(st.clone(), METRICS_KEY_ACME).await;
18233 assert_eq!(
18234 body.as_object().unwrap().len(),
18235 2,
18236 "completion metrics must contain only tenant-scoped rows",
18237 );
18238 let tenants = body["tenants"].as_object().unwrap();
18239 assert_eq!(tenants.len(), 1);
18240 assert_eq!(tenants["t:acme"]["prompt_tokens_in"], 100);
18241 assert!(!tenants.contains_key("t:blue"));
18242 let adsd = body["adsd_suspect_total"].as_object().unwrap();
18243 assert_eq!(adsd.len(), 1);
18244 assert_eq!(adsd["t:acme"], 1);
18245 assert!(!adsd.contains_key("t:blue"));
18246
18247 let mut headers = HeaderMap::new();
18248 headers.insert(
18249 "authorization",
18250 format!("Bearer {METRICS_KEY_ACME}").parse().unwrap(),
18251 );
18252 assert_eq!(
18253 yield_metrics(State(st), headers).await.status(),
18254 StatusCode::FORBIDDEN,
18255 "the process-wide yield view requires an operator metrics token",
18256 );
18257 }
18258
18259 #[tokio::test]
18260 async fn tenant_metrics_hide_capacity_and_aggregate_spec() {
18261 let body = metrics_json(multi_key_metrics_state(None), METRICS_KEY_ACME).await;
18262 for operator_only in [
18263 "prefix_cache_entries",
18264 "prefix_cache_bytes",
18265 "prefix_cache_skips_budget",
18266 "prefix_cache_skips_pinned",
18267 "active_sessions",
18268 "queued_requests",
18269 "admission_inflight",
18270 "admission_booked_bytes",
18271 "continuation_pool_entries",
18272 "spec_pool_entries",
18273 "cuda_driver_free_bytes",
18274 "cuda_pool_reserved_bytes",
18275 "cuda_pool_used_bytes",
18276 "cuda_pool_cached_bytes",
18277 "constraint_compiler_fail_closed",
18278 "serve_idle_seconds",
18279 "spec",
18280 "spec_tau",
18281 "spec_accept_by_position",
18282 "dual_pp",
18283 "pp_wave",
18284 "peer_probe_bypassed",
18285 "peer_probe_boundary_copies",
18286 "peer_probe_runtime_reprobes",
18287 "peer_probe_runtime_failures",
18288 "peer_probe_deferred_total",
18289 "peer_probe_integrity_degraded",
18290 "peer_probe_degraded_to_host_bounce",
18291 ] {
18292 assert!(
18293 body.get(operator_only).is_none(),
18294 "tenant metrics must not expose operator field {operator_only}",
18295 );
18296 }
18297 }
18298
18299 #[test]
18300 fn populated_spec_acceptance_metrics_are_operator_only() {
18301 for scope in [
18302 MetricsScope::CompletionDomain,
18303 MetricsScope::Tenant("t:acme".into()),
18304 ] {
18305 let mut body = json!({});
18306 insert_spec_acceptance_metrics(&mut body, &scope, || {
18307 panic!("tenant scope evaluated the process-wide spec snapshot")
18308 });
18309 assert!(body.get("spec_tau").is_none(), "{scope:?} leaked spec tau");
18310 assert!(
18311 body.get("spec_accept_by_position").is_none(),
18312 "{scope:?} leaked the accept histogram"
18313 );
18314 }
18315
18316 let mut telemetry = memra_engine::spec::SpecTelemetry {
18317 rounds: 4,
18318 drafted: 12,
18319 accepted: 6,
18320 ..Default::default()
18321 };
18322 telemetry.pos_drafted[..3].copy_from_slice(&[4, 4, 4]);
18323 telemetry.pos_accepted[..3].copy_from_slice(&[3, 2, 1]);
18324 let mut body = json!({});
18325 insert_spec_acceptance_metrics(&mut body, &MetricsScope::All, || {
18326 HashMap::from([("model-a".to_string(), telemetry)])
18327 });
18328 assert_eq!(body["spec_tau"]["model-a"], 1.5);
18329 let histogram = &body["spec_accept_by_position"]["model-a"];
18330 assert_eq!(histogram["window_seconds"], worker::SPEC_METRICS_WINDOW_S);
18331 assert_eq!(histogram["rounds"], 4);
18332 assert_eq!(histogram["offered"], json!([4, 4, 4]));
18333 assert_eq!(histogram["accepted"], json!([3, 2, 1]));
18334 assert_eq!(histogram["accept_rate"], json!([0.75, 0.5, 0.25]));
18335 }
18336
18337 #[test]
18338 fn populated_dual_pp_metrics_are_operator_only() {
18339 let populated = DualPpMetricsSnapshot {
18340 stage_ns: [1_000_000, 2_000_000, 3_000_000, 4_000_000],
18341 stage_samples: [1, 1, 1, 1],
18342 dropped_timing_samples: 0,
18343 overlaps: 17,
18344 slot_pairs: 19,
18345 slot_uses: [19, 19],
18346 slot_collisions: 0,
18347 };
18348 for scope in [
18349 MetricsScope::CompletionDomain,
18350 MetricsScope::Tenant("t:acme".into()),
18351 ] {
18352 let mut body = json!({});
18353 insert_dual_pp_metrics(&mut body, &scope, || populated);
18354 assert!(
18355 body.get("dual_pp").is_none(),
18356 "{scope:?} leaked dual PP topology"
18357 );
18358 }
18359
18360 let mut body = json!({});
18361 insert_dual_pp_metrics(&mut body, &MetricsScope::All, || populated);
18362 assert_eq!(body["dual_pp"]["overlaps"], 17);
18363 assert_eq!(body["dual_pp"]["slot_pairs"], 19);
18364 assert_eq!(body["dual_pp"]["slot_uses"], json!([19, 19]));
18365 assert_eq!(body["dual_pp"]["slot_collisions"], 0);
18366 assert_eq!(
18367 body["dual_pp"]["cuda_event_spans"]["wave_a_stage0"]["mean_ms"],
18368 1.0
18369 );
18370 }
18371
18372 #[test]
18373 fn populated_pp_wave_metrics_are_operator_only() {
18374 let populated = PpWaveMetricsSnapshot {
18375 ticks: 11,
18376 cells: 96,
18377 overlaps: 37,
18378 };
18379 for scope in [
18380 MetricsScope::CompletionDomain,
18381 MetricsScope::Tenant("t:acme".into()),
18382 ] {
18383 let mut body = json!({});
18384 insert_pp_wave_metrics(&mut body, &scope, || populated);
18385 assert!(
18386 body.get("pp_wave").is_none(),
18387 "{scope:?} leaked PP wave topology"
18388 );
18389 }
18390
18391 let mut body = json!({});
18392 insert_pp_wave_metrics(&mut body, &MetricsScope::All, || populated);
18393 assert_eq!(body["pp_wave"]["ticks"], 11);
18394 assert_eq!(body["pp_wave"]["cells"], 96);
18395 assert_eq!(body["pp_wave"]["overlaps"], 37);
18396 }
18397
18398 #[test]
18399 fn peer_probe_metrics_are_operator_only() {
18400 let populated = memra_engine::pp::PeerProbeMetrics {
18401 bypassed: 1,
18402 boundary_copies: 8_192,
18403 runtime_probes: 1,
18404 runtime_failures: 0,
18405 deferred_total: 4,
18406 integrity_degraded: true,
18407 degraded_to_host_bounce: true,
18408 };
18409 for scope in [
18410 MetricsScope::CompletionDomain,
18411 MetricsScope::Tenant("t:acme".into()),
18412 ] {
18413 let mut body = json!({});
18414 insert_peer_probe_metrics(&mut body, &scope, || populated);
18415 assert!(body.get("peer_probe_bypassed").is_none());
18416 }
18417
18418 let mut body = json!({});
18419 insert_peer_probe_metrics(&mut body, &MetricsScope::All, || populated);
18420 assert_eq!(body["peer_probe_bypassed"], 1);
18421 assert_eq!(body["peer_probe_boundary_copies"], 8_192);
18422 assert_eq!(body["peer_probe_runtime_reprobes"], 1);
18423 assert_eq!(body["peer_probe_runtime_failures"], 0);
18424 assert_eq!(body["peer_probe_deferred_total"], 4);
18425 assert_eq!(body["peer_probe_integrity_degraded"], true);
18426 assert_eq!(body["peer_probe_degraded_to_host_bounce"], true);
18427 }
18428
18429 #[tokio::test]
18430 async fn prefix_aggregate_metrics_are_operator_only_but_tenant_ratio_remains() {
18431 let tenant_body = metrics_json(multi_key_metrics_state(None), METRICS_KEY_ACME).await;
18432 for operator_only in [
18433 "lcp_histogram",
18434 "cache_hit_token_ratio",
18435 "prefix_cache_hits",
18436 "prefix_cache_misses",
18437 "prefix_cache_inserts",
18438 "prefix_cache_evictions",
18439 "prefix_cache_skips_budget",
18440 "prefix_cache_skips_pinned",
18441 "prefix_cache_hit_tokens",
18442 ] {
18443 assert!(
18444 tenant_body.get(operator_only).is_none(),
18445 "tenant metrics must not expose global prefix field {operator_only}",
18446 );
18447 }
18448 assert_eq!(tenant_body["tenants"].as_object().unwrap().len(), 1);
18449 assert_eq!(tenant_body["tenants"]["t:acme"]["prompt_tokens_in"], 100);
18450 assert_eq!(tenant_body["tenants"]["t:acme"]["cached_tokens_in"], 40);
18451 assert_eq!(
18452 tenant_body["tenants"]["t:acme"]["cache_hit_token_ratio"],
18453 0.4
18454 );
18455
18456 let operator_body = metrics_json(
18457 multi_key_metrics_state(Some("scrape-secret")),
18458 "scrape-secret",
18459 )
18460 .await;
18461 assert_eq!(operator_body["prefix_cache_hits"], 2);
18462 assert_eq!(operator_body["prefix_cache_misses"], 3);
18463 assert_eq!(operator_body["prefix_cache_inserts"], 5);
18464 assert_eq!(operator_body["prefix_cache_evictions"], 7);
18465 assert_eq!(operator_body["prefix_cache_skips_budget"], 9);
18466 assert_eq!(operator_body["prefix_cache_skips_pinned"], 10);
18467 assert_eq!(operator_body["prefix_cache_hit_tokens"], 11);
18468 assert_eq!(operator_body["cache_hit_token_ratio"], 0.15);
18469 assert_eq!(operator_body["lcp_histogram"]["counts"][4], 13);
18470 }
18471
18472 #[tokio::test]
18473 async fn configured_metrics_token_is_exclusive_and_sees_all_tenants() {
18474 let st = multi_key_metrics_state(Some("scrape-secret"));
18475 let mut completion_headers = HeaderMap::new();
18476 completion_headers.insert(
18477 "authorization",
18478 format!("Bearer {METRICS_KEY_ACME}").parse().unwrap(),
18479 );
18480 assert_eq!(
18481 get_metrics(State(st.clone()), completion_headers.clone())
18482 .await
18483 .status(),
18484 StatusCode::FORBIDDEN,
18485 );
18486 assert_eq!(
18487 yield_metrics(State(st.clone()), completion_headers)
18488 .await
18489 .status(),
18490 StatusCode::FORBIDDEN,
18491 );
18492
18493 let body = metrics_json(st.clone(), "scrape-secret").await;
18494 let tenants = body["tenants"].as_object().unwrap();
18495 assert_eq!(tenants.len(), 2);
18496 assert!(tenants.contains_key("t:acme"));
18497 assert!(tenants.contains_key("t:blue"));
18498 assert_eq!(body["adsd_suspect_total"]["t:acme"], 1);
18499 assert_eq!(body["adsd_suspect_total"]["t:blue"], 2);
18500 assert_eq!(body["active_sessions"], 3);
18501 assert_eq!(body["queued_requests"], 5);
18502 assert_eq!(body["admission_inflight"]["m"], 4);
18504 assert_eq!(body["admission_booked_bytes"]["m"], 41_000_000);
18505 assert_eq!(body["prefix_cache_bytes"], 31);
18506 assert_eq!(body["cuda_driver_free_bytes"], 13);
18507 assert_eq!(body["constraint_compiler_fail_closed"]["m"], 1);
18508 assert_eq!(body["spec"]["m"]["drafted"], 6);
18509 assert_eq!(body["spec_tau"]["m"], 1.5);
18510 assert_eq!(
18511 body["spec_accept_by_position"]["m"]["accepted"],
18512 json!([3, 2, 1])
18513 );
18514 let yield_body = yield_metrics_json(st, "scrape-secret").await;
18515 assert_eq!(yield_body["batch_size_last"], 37);
18516 }
18517
18518 #[tokio::test]
18519 async fn metrics_token_protects_public_override_without_api_keys() {
18520 let mut st = fake_worker_state();
18521 st.metrics_auth = MetricsAuth::new(false, false, Some("scrape-secret".into()));
18522
18523 assert_eq!(
18524 get_metrics(State(st.clone()), HeaderMap::new())
18525 .await
18526 .status(),
18527 StatusCode::UNAUTHORIZED,
18528 );
18529 let mut headers = HeaderMap::new();
18530 headers.insert("authorization", "Bearer scrape-secret".parse().unwrap());
18531 assert_eq!(
18532 get_metrics(State(st.clone()), headers.clone())
18533 .await
18534 .status(),
18535 StatusCode::OK,
18536 );
18537 assert_eq!(
18538 yield_metrics(State(st), headers).await.status(),
18539 StatusCode::OK
18540 );
18541 }
18542
18543 #[tokio::test]
18544 async fn no_key_loopback_metrics_remain_open_for_development() {
18545 let mut st = fake_worker_state();
18546 st.metrics_auth = MetricsAuth::new(true, false, None);
18547 let response = get_metrics(State(st.clone()), HeaderMap::new()).await;
18548 assert_eq!(response.status(), StatusCode::OK);
18549 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18550 .await
18551 .unwrap();
18552 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
18553 assert!(
18554 body.get("active_sessions").is_some(),
18555 "no-key loopback development keeps full operator visibility",
18556 );
18557 assert_eq!(
18558 yield_metrics(State(st), HeaderMap::new()).await.status(),
18559 StatusCode::OK,
18560 );
18561 }
18562
18563 #[test]
18564 fn rate_limit_math_remaining_hits_zero_at_cap_and_reset_arms() {
18565 let metrics = SharedMetrics::default();
18566 let rl = RateLimit::compute(4, 1, &metrics);
18568 assert_eq!((rl.limit, rl.remaining, rl.reset_s), (4, 3, 0));
18569 let rl = RateLimit::compute(4, 3, &metrics);
18570 assert_eq!(rl.remaining, 1);
18571 let rl = RateLimit::compute(4, 4, &metrics);
18573 assert_eq!(rl.remaining, 0);
18574 assert!(rl.reset_s > 0, "reset must arm when no slots are free");
18575 assert_eq!(RateLimit::compute(4, 9, &metrics).remaining, 0);
18577 let m = worker::Metrics {
18579 completed: 2,
18580 tokens_out: 200,
18581 step_p50_ms: 20.0,
18582 ..Default::default()
18583 };
18584 assert_eq!(reset_estimate_s(&m), 2); }
18586
18587 #[test]
18588 fn inflight_guard_counts_up_and_frees_on_drop() {
18589 let counts: InflightCounts = Arc::new(Default::default());
18590 let tenants: TenantGauge = Arc::new(Default::default());
18591 let (g1, n1, t1) = InflightGuard::try_acquire(
18592 counts.clone(),
18593 lanes::Lane::Interactive,
18594 tenants.clone(),
18595 "acme",
18596 None,
18597 )
18598 .unwrap();
18599 let (g2, n2, t2) = InflightGuard::try_acquire(
18600 counts.clone(),
18601 lanes::Lane::Interactive,
18602 tenants.clone(),
18603 "acme",
18604 None,
18605 )
18606 .unwrap();
18607 assert_eq!((n1, n2), (1, 2));
18608 assert_eq!((t1, t2), (1, 2));
18610 let (gj, nj, tj) = InflightGuard::try_acquire(
18612 counts.clone(),
18613 lanes::Lane::Judge,
18614 tenants.clone(),
18615 "blue",
18616 None,
18617 )
18618 .unwrap();
18619 assert_eq!((nj, tj), (1, 1));
18620 drop(g1);
18621 drop(gj);
18622 assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 1);
18623 assert_eq!(counts[1].load(std::sync::atomic::Ordering::SeqCst), 0);
18624 assert_eq!(tenants.lock().unwrap().get("acme"), Some(&1));
18625 assert!(tenants.lock().unwrap().get("blue").is_none());
18627 drop(g2);
18628 assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 0);
18629 assert!(tenants.lock().unwrap().is_empty());
18630 }
18631
18632 #[test]
18633 fn tenant_concurrency_cap_is_atomic_across_arrivals() {
18634 let counts: InflightCounts = Arc::new(Default::default());
18635 let tenants: TenantGauge = Arc::new(Default::default());
18636 let start = Arc::new(std::sync::Barrier::new(3));
18637 let attempted = Arc::new(std::sync::Barrier::new(3));
18638 let mut joins = Vec::new();
18639 for _ in 0..2 {
18640 let counts = counts.clone();
18641 let tenants = tenants.clone();
18642 let start = start.clone();
18643 let attempted = attempted.clone();
18644 joins.push(std::thread::spawn(move || {
18645 start.wait();
18646 let result = InflightGuard::try_acquire(
18647 counts,
18648 lanes::Lane::Interactive,
18649 tenants,
18650 "preview_001",
18651 Some(1),
18652 );
18653 let won = result.is_ok();
18654 attempted.wait(); drop(result);
18656 won
18657 }));
18658 }
18659 start.wait();
18660 attempted.wait();
18661 let wins = joins
18662 .into_iter()
18663 .map(|join| join.join().unwrap())
18664 .filter(|won| *won)
18665 .count();
18666 assert_eq!(wins, 1, "exactly one simultaneous request may pass cap=1");
18667 assert_eq!(counts[0].load(std::sync::atomic::Ordering::SeqCst), 0);
18668 assert!(tenants.lock().unwrap().is_empty());
18669 }
18670
18671 #[tokio::test]
18672 async fn tenant_concurrency_cap_rejects_before_worker_admission() {
18673 let st = fake_worker_state();
18674 let tenant = auth::TenantCtx {
18675 tenant: "preview_001".into(),
18676 lane_class: auth::LaneClass::Interactive,
18677 rate_limit: Some(1),
18678 key_prefix: None,
18679 };
18680 let first_env = Envelope::new(true);
18681 let (guard, first_rl) =
18682 match acquire_request_slot(&st, lanes::Lane::Interactive, &tenant, &first_env) {
18683 Ok(slot) => slot,
18684 Err(_) => panic!("the first request must acquire the tenant slot"),
18685 };
18686 assert_eq!((first_rl.limit, first_rl.remaining), (1, 0));
18687
18688 let second_env = Envelope::new(true);
18689 let response =
18690 match acquire_request_slot(&st, lanes::Lane::Interactive, &tenant, &second_env) {
18691 Err(response) => response,
18692 Ok(_) => panic!("the second request must be rejected at the tenant cap"),
18693 };
18694 assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
18695 assert_eq!(response.headers()["retry-after"], "2");
18696 assert_eq!(response.headers()["retry-after-ms"], "2000");
18697 assert_eq!(response.headers()["x-ratelimit-limit"], "1");
18698 assert_eq!(response.headers()["x-ratelimit-remaining"], "0");
18699 assert_eq!(response.headers()["x-request-id"], second_env.id);
18700 assert_eq!(
18701 st.inflight[0].load(std::sync::atomic::Ordering::SeqCst),
18702 1,
18703 "rejected request must not consume a lane slot"
18704 );
18705 assert_eq!(
18706 st.tenant_inflight
18707 .lock()
18708 .unwrap()
18709 .get("preview_001")
18710 .copied(),
18711 Some(1),
18712 "rejected request must not increment the tenant gauge"
18713 );
18714 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
18715 .await
18716 .unwrap();
18717 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
18718 assert_eq!(payload["error"]["type"], "rate_limit_error");
18719 assert_eq!(payload["error"]["code"], "rate_limit_exceeded");
18720 assert!(
18721 payload["error"]["message"]
18722 .as_str()
18723 .unwrap()
18724 .contains("concurrent request limit")
18725 );
18726
18727 drop(guard);
18728 let _ = InflightGuard::try_acquire(
18729 st.inflight.clone(),
18730 lanes::Lane::Interactive,
18731 st.tenant_inflight.clone(),
18732 "preview_001",
18733 Some(1),
18734 )
18735 .expect("slot must reopen after the in-flight request completes");
18736 }
18737
18738 #[test]
18739 fn tenant_rate_limit_override_is_min_with_global_cap() {
18740 let metrics = SharedMetrics::default();
18741 let unlimited = auth::TenantCtx::default_tenant();
18742 let capped = auth::TenantCtx {
18743 tenant: "acme".into(),
18744 lane_class: auth::LaneClass::Interactive,
18745 rate_limit: Some(2),
18746 key_prefix: None,
18747 };
18748 let global = lane_cap(lanes::Lane::Interactive);
18749 let rl = RateLimit::at_admit(lanes::Lane::Interactive, 1, &metrics, &unlimited, 1);
18751 assert_eq!((rl.limit, rl.remaining), (global, global - 1));
18752 let rl = RateLimit::at_admit(lanes::Lane::Interactive, 5, &metrics, &capped, 1);
18754 assert_eq!((rl.limit, rl.remaining), (2, 1));
18755 let rl = RateLimit::at_admit(lanes::Lane::Interactive, 5, &metrics, &capped, 2);
18756 assert_eq!(rl.remaining, 0);
18757 assert!(rl.reset_s > 0, "reset must arm at the tenant cap too");
18758 let rl = RateLimit::at_admit(lanes::Lane::Interactive, global, &metrics, &capped, 0);
18762 assert_eq!(rl.remaining, 0);
18763 let wide = auth::TenantCtx {
18764 rate_limit: Some(global + 100),
18765 ..capped.clone()
18766 };
18767 let rl = RateLimit::at_admit(lanes::Lane::Interactive, 1, &metrics, &wide, 1);
18768 assert_eq!((rl.limit, rl.remaining), (global, global - 1));
18769 }
18770
18771 #[test]
18772 fn batch_class_keys_default_to_harvest_and_cannot_claim_interactive() {
18773 let batch = auth::TenantCtx {
18774 tenant: "bulk".into(),
18775 lane_class: auth::LaneClass::Batch,
18776 rate_limit: None,
18777 key_prefix: None,
18778 };
18779 let interactive = auth::TenantCtx::default_tenant();
18780 let hdr = |v: Option<&str>| {
18781 let mut h = axum::http::HeaderMap::new();
18782 if let Some(v) = v {
18783 h.insert("x-lane", axum::http::HeaderValue::from_str(v).unwrap());
18784 }
18785 h
18786 };
18787 assert_eq!(
18789 lane_for_tenant(&hdr(None), &interactive).unwrap(),
18790 lanes::Lane::Interactive
18791 );
18792 assert_eq!(
18793 lane_for_tenant(&hdr(Some("judge")), &interactive).unwrap(),
18794 lanes::Lane::Judge
18795 );
18796 assert_eq!(
18798 lane_for_tenant(&hdr(None), &batch).unwrap(),
18799 lanes::Lane::Harvest
18800 );
18801 assert_eq!(
18802 lane_for_tenant(&hdr(Some("judge")), &batch).unwrap(),
18803 lanes::Lane::Judge
18804 );
18805 let resp = lane_for_tenant(&hdr(Some("interactive")), &batch).unwrap_err();
18806 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
18807 let resp = lane_for_tenant(&hdr(Some("turbo")), &interactive).unwrap_err();
18809 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
18810 }
18811
18812 #[tokio::test]
18813 async fn handler_layer_refusals_are_openai_objects_with_x_should_retry() {
18814 let hdr = |v: &str| {
18819 let mut h = axum::http::HeaderMap::new();
18820 h.insert("x-lane", axum::http::HeaderValue::from_str(v).unwrap());
18821 h
18822 };
18823 let body = |resp: Response| async move {
18824 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
18825 .await
18826 .unwrap();
18827 serde_json::from_slice::<serde_json::Value>(&bytes).unwrap()
18828 };
18829
18830 let resp = lane_for_tenant(&hdr("turbo"), &auth::TenantCtx::default_tenant()).unwrap_err();
18831 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
18832 assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
18833 let payload = body(resp).await;
18834 assert!(
18835 payload["error"].is_object(),
18836 "bare-string error body: {payload}"
18837 );
18838 assert_eq!(payload["error"]["type"], "invalid_request_error");
18839 assert_eq!(payload["error"]["param"], "x-lane");
18840 assert_eq!(payload["error"]["code"], "invalid_lane");
18841
18842 let batch = auth::TenantCtx {
18843 tenant: "bulk".into(),
18844 lane_class: auth::LaneClass::Batch,
18845 rate_limit: None,
18846 key_prefix: None,
18847 };
18848 let resp = lane_for_tenant(&hdr("interactive"), &batch).unwrap_err();
18849 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
18850 assert_eq!(resp.headers().get("x-should-retry").unwrap(), "false");
18851 let payload = body(resp).await;
18852 assert_eq!(payload["error"]["type"], "authentication_error");
18853 assert_eq!(payload["error"]["param"], "x-lane");
18854 }
18855
18856 static DRAIN_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
18859
18860 fn drain_lock() -> std::sync::MutexGuard<'static, ()> {
18873 let guard = DRAIN_LOCK.lock().unwrap_or_else(|poisoned| {
18874 DRAIN_LOCK.clear_poison();
18877 poisoned.into_inner()
18878 });
18879 DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
18880 guard
18881 }
18882
18883 struct DrainingRestore;
18889 impl Drop for DrainingRestore {
18890 fn drop(&mut self) {
18891 DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
18892 }
18893 }
18894
18895 #[tokio::test]
18896 #[allow(clippy::await_holding_lock)] async fn responses_carry_rate_limit_headers_and_slot_frees() {
18898 let _l = drain_lock();
18899 let st = fake_worker_state();
18900 let resp = chat_completions(
18903 State(st.clone()),
18904 axum::http::HeaderMap::new(),
18905 None,
18906 Json(
18907 serde_json::from_value(serde_json::json!({
18908 "model": "m", "messages": [{"role": "user", "content": "t"}]
18909 }))
18910 .unwrap(),
18911 ),
18912 )
18913 .await;
18914 assert_eq!(resp.status(), StatusCode::OK);
18915 let h = resp.headers();
18916 let limit: usize = h["x-ratelimit-limit"].to_str().unwrap().parse().unwrap();
18917 let remaining: usize = h["x-ratelimit-remaining"]
18918 .to_str()
18919 .unwrap()
18920 .parse()
18921 .unwrap();
18922 assert_eq!(remaining, limit - 1);
18923 assert_eq!(h["x-ratelimit-reset"], "0");
18924 assert_eq!(
18925 st.inflight[0].load(std::sync::atomic::Ordering::SeqCst),
18926 0,
18927 "slot must free at completion"
18928 );
18929 let resp = completions(
18932 State(st.clone()),
18933 axum::http::HeaderMap::new(),
18934 None,
18935 Json(
18936 serde_json::from_value(serde_json::json!({
18937 "model": "m", "prompt": "t", "stream": true
18938 }))
18939 .unwrap(),
18940 ),
18941 )
18942 .await;
18943 assert_eq!(resp.status(), StatusCode::OK);
18944 assert!(resp.headers().contains_key("x-ratelimit-limit"));
18945 assert!(resp.headers().contains_key("x-ratelimit-remaining"));
18946 assert!(resp.headers().contains_key("x-ratelimit-reset"));
18947 assert_eq!(
18948 st.inflight[0].load(std::sync::atomic::Ordering::SeqCst),
18949 1,
18950 "stream in flight holds the slot"
18951 );
18952 let _ = axum::body::to_bytes(resp.into_body(), usize::MAX)
18953 .await
18954 .unwrap();
18955 assert_eq!(
18956 st.inflight[0].load(std::sync::atomic::Ordering::SeqCst),
18957 0,
18958 "slot must free when the stream completes"
18959 );
18960 }
18961
18962 #[tokio::test]
18968 #[allow(clippy::await_holding_lock)] async fn multi_item_capture_requests_open_one_receipt_per_item_under_child_ids() {
18970 let _l = drain_lock();
18971 let mut st = fake_worker_state();
18972 let mock = MockMetering::admit_all();
18973 st.metering = Some(mock.clone());
18974
18975 let resp = embed_api::embeddings_admitted(
18976 State(st.clone()),
18977 HeaderMap::new(),
18978 AdmittedJson(
18979 serde_json::from_value(json!({"model": "m", "input": ["a", "bb", "ccc"]})).unwrap(),
18980 BodyAdmissionLease(None),
18981 ),
18982 )
18983 .await;
18984 assert_eq!(resp.status(), StatusCode::OK);
18985 let parent = resp.headers()["x-request-id"].to_str().unwrap().to_string();
18986 assert!(
18987 !parent.contains('.'),
18988 "the caller sees the parent id: {parent}"
18989 );
18990 let body: serde_json::Value = serde_json::from_slice(
18991 &axum::body::to_bytes(resp.into_body(), usize::MAX)
18992 .await
18993 .unwrap(),
18994 )
18995 .unwrap();
18996 assert_eq!(body["data"].as_array().map(Vec::len), Some(3));
18997 let events = mock.events();
18998 let opened: Vec<(String, &'static str)> = events
18999 .iter()
19000 .filter_map(|e| match e {
19001 MeterEvent::Open {
19002 request_id, route, ..
19003 } => Some((request_id.clone(), *route)),
19004 _ => None,
19005 })
19006 .collect();
19007 assert_eq!(
19008 opened,
19009 vec![
19010 (format!("{parent}.0"), "/v1/embeddings"),
19011 (format!("{parent}.1"), "/v1/embeddings"),
19012 (format!("{parent}.2"), "/v1/embeddings"),
19013 ],
19014 "one receipt per input, each under its own child id: {events:?}"
19015 );
19016 assert_eq!(
19017 events
19018 .iter()
19019 .filter(|e| matches!(e, MeterEvent::Complete { .. }))
19020 .count(),
19021 3,
19022 "every input settles its own receipt: {events:?}"
19023 );
19024
19025 let resp = embed_api::rerank_admitted(
19026 State(st),
19027 HeaderMap::new(),
19028 AdmittedJson(
19029 serde_json::from_value(
19030 json!({"model": "m", "query": "q", "documents": ["d0", "d1"]}),
19031 )
19032 .unwrap(),
19033 BodyAdmissionLease(None),
19034 ),
19035 )
19036 .await;
19037 assert_eq!(resp.status(), StatusCode::OK);
19038 let parent = resp.headers()["x-request-id"].to_str().unwrap().to_string();
19039 let opened: Vec<String> = mock
19040 .events()
19041 .into_iter()
19042 .skip(events.len())
19043 .filter_map(|e| match e {
19044 MeterEvent::Open {
19045 request_id,
19046 route: "/v1/rerank",
19047 ..
19048 } => Some(request_id),
19049 _ => None,
19050 })
19051 .collect();
19052 assert_eq!(opened, vec![format!("{parent}.0"), format!("{parent}.1")]);
19053 }
19054
19055 #[tokio::test]
19056 #[allow(clippy::await_holding_lock)] async fn handlers_sync_worker_truth_usage_and_cost_before_terminal_response() {
19058 let _l = drain_lock();
19059 let mut st = fake_worker_state();
19060 let mock = MockMetering::admit_all();
19061 st.metering = Some(mock.clone());
19062
19063 let nonstream = chat_completions(
19064 State(st.clone()),
19065 HeaderMap::new(),
19066 None,
19067 Json(
19068 serde_json::from_value(json!({
19069 "model": "m",
19070 "messages": [{"role": "user", "content": "t"}],
19071 }))
19072 .unwrap(),
19073 ),
19074 )
19075 .await;
19076 assert_eq!(nonstream.status(), StatusCode::OK);
19077 let nonstream_id = nonstream.headers()["x-request-id"]
19078 .to_str()
19079 .unwrap()
19080 .to_string();
19081
19082 let stream = completions(
19083 State(st),
19084 HeaderMap::new(),
19085 None,
19086 Json(
19087 serde_json::from_value(json!({
19088 "model": "m",
19089 "prompt": "t",
19090 "stream": true,
19091 }))
19092 .unwrap(),
19093 ),
19094 )
19095 .await;
19096 assert_eq!(stream.status(), StatusCode::OK);
19097 let stream_id = stream.headers()["x-request-id"]
19098 .to_str()
19099 .unwrap()
19100 .to_string();
19101 let _ = axum::body::to_bytes(stream.into_body(), usize::MAX)
19102 .await
19103 .unwrap();
19104
19105 let events = mock.events();
19109 let opened: Vec<&str> = events
19110 .iter()
19111 .filter_map(|e| match e {
19112 MeterEvent::Open { request_id, .. } => Some(request_id.as_str()),
19113 _ => None,
19114 })
19115 .collect();
19116 assert_eq!(opened, vec![nonstream_id.as_str(), stream_id.as_str()]);
19117 let completes = events
19118 .iter()
19119 .filter(|e| {
19120 matches!(
19121 e,
19122 MeterEvent::Complete {
19123 prompt: 1,
19124 cached: 0,
19125 completion: 1,
19126 }
19127 )
19128 })
19129 .count();
19130 assert_eq!(
19131 completes, 2,
19132 "both surfaces settle complete with worker-truth usage: {events:?}"
19133 );
19134 }
19135
19136 #[tokio::test]
19137 #[allow(clippy::await_holding_lock)] async fn completion_admission_supports_metered_blocked_and_paid_transitions() {
19139 let _l = drain_lock();
19140 let mock = MockMetering::with_limits(vec![
19146 ReserveScript::Insufficient,
19147 ReserveScript::Admit { with_permit: false },
19148 ReserveScript::Blocked,
19149 ReserveScript::Admit { with_permit: true },
19150 ]);
19151 let mut st = fake_worker_state();
19152 st.metering = Some(mock.clone());
19153
19154 let metrics = get_metrics(State(st.clone()), HeaderMap::new()).await;
19156 assert_eq!(metrics.status(), StatusCode::OK);
19157 let metrics_body = axum::body::to_bytes(metrics.into_body(), usize::MAX)
19158 .await
19159 .unwrap();
19160 let metrics_body: serde_json::Value = serde_json::from_slice(&metrics_body).unwrap();
19161 assert_eq!(metrics_body["budget_source_reload_failed"], 0);
19162 assert_eq!(metrics_body["budget_source_reload_consecutive"], 0);
19163 assert_eq!(metrics_body["budget_source_available"], true);
19164
19165 let request = || {
19166 Json(
19167 serde_json::from_value::<CompletionReq>(json!({
19168 "model": "m",
19169 "prompt_ids": [1],
19170 "max_tokens": 1,
19171 }))
19172 .unwrap(),
19173 )
19174 };
19175
19176 let denied = completions(State(st.clone()), HeaderMap::new(), None, request()).await;
19177 assert_eq!(denied.status(), StatusCode::PAYMENT_REQUIRED);
19178 let denied_body = axum::body::to_bytes(denied.into_body(), usize::MAX)
19179 .await
19180 .unwrap();
19181 let denied_body: serde_json::Value = serde_json::from_slice(&denied_body).unwrap();
19182 assert_eq!(denied_body["error"]["type"], "insufficient_balance");
19183 assert_eq!(denied_body["error"]["code"], "insufficient_balance");
19184
19185 let included = completions(State(st.clone()), HeaderMap::new(), None, request()).await;
19186 assert_eq!(included.status(), StatusCode::OK);
19187
19188 let blocked = completions(State(st.clone()), HeaderMap::new(), None, request()).await;
19191 assert_eq!(blocked.status(), StatusCode::PAYMENT_REQUIRED);
19192
19193 let admitted = completions(State(st.clone()), HeaderMap::new(), None, request()).await;
19194 assert_eq!(admitted.status(), StatusCode::OK);
19195
19196 let events = mock.events();
19197 let terminal: Vec<&MeterEvent> = events
19198 .iter()
19199 .filter(|e| matches!(e, MeterEvent::Reject { .. } | MeterEvent::Complete { .. }))
19200 .collect();
19201 assert_eq!(
19202 terminal.len(),
19203 4,
19204 "four requests, four terminal settles: {events:?}"
19205 );
19206 assert!(matches!(
19207 terminal[0],
19208 MeterEvent::Reject { status: 402, .. }
19209 ));
19210 assert!(matches!(terminal[1], MeterEvent::Complete { .. }));
19211 assert!(matches!(
19212 terminal[2],
19213 MeterEvent::Reject { status: 402, .. }
19214 ));
19215 assert!(matches!(terminal[3], MeterEvent::Complete { .. }));
19216 let permits: Vec<bool> = events
19218 .iter()
19219 .filter_map(|e| match e {
19220 MeterEvent::Open { with_permit, .. } => Some(*with_permit),
19221 _ => None,
19222 })
19223 .collect();
19224 assert_eq!(
19225 permits,
19226 vec![false, false, false, true],
19227 "the permit rides the receipt exactly when reserve minted one: {events:?}"
19228 );
19229 }
19230
19231 #[tokio::test]
19235 async fn a_capped_key_answers_its_own_402_and_the_principal_crosses_the_seam() {
19236 let mock = MockMetering::with_limits(vec![ReserveScript::PrincipalCapped]);
19237 let mut st = fake_worker_state();
19238 st.metering = Some(mock.clone());
19239 let tenant = auth::TenantCtx {
19240 tenant: "acme".into(),
19241 lane_class: auth::LaneClass::Interactive,
19242 rate_limit: None,
19243 key_prefix: Some("mk-acme-testprefix00".into()),
19244 };
19245 let mut request = gate_request(1, 1);
19246 let rejection = admit_tenant_budget(&st, &tenant, &mut request)
19247 .expect_err("a capped key must be refused at admission");
19248 assert!(matches!(rejection, BudgetRejection::PrincipalCapped));
19249 let (response, outcome) = rejection.into_response();
19250 assert_eq!(outcome, "key_spend_cap_reached");
19251 assert_eq!(response.status(), StatusCode::PAYMENT_REQUIRED);
19252 let body = body_value(response).await;
19253 assert_eq!(body["error"]["code"], "key_spend_cap_reached");
19254 assert!(
19255 body["error"]["message"].as_str().unwrap().contains("cap"),
19256 "the 402 must point at the KEY's cap, not tenant credit: {body}"
19257 );
19258 let events = mock.events();
19259 assert!(
19260 events.contains(&MeterEvent::Reserve {
19261 tenant: "acme".into(),
19262 principal: Some("mk-acme-testprefix00".into()),
19263 model: "qwen/qwen3.8-27b".into(),
19264 }),
19265 "the key prefix must reach reserve: {events:?}"
19266 );
19267 }
19268
19269 #[tokio::test]
19270 #[allow(clippy::await_holding_lock)] async fn streaming_client_disconnect_records_partial_usage_and_cost() {
19272 let _l = drain_lock();
19273 let mut st = fake_worker_state_with_steps(4, std::time::Duration::from_millis(100));
19274 let mock = MockMetering::admit_all();
19275 st.metering = Some(mock.clone());
19276
19277 let response = completions(
19278 State(st),
19279 HeaderMap::new(),
19280 None,
19281 Json(
19282 serde_json::from_value(json!({
19283 "model": "m",
19284 "prompt": "disconnect after one delta",
19285 "stream": true,
19286 }))
19287 .unwrap(),
19288 ),
19289 )
19290 .await;
19291 assert_eq!(response.status(), StatusCode::OK);
19292 let request_id = response.headers()["x-request-id"]
19293 .to_str()
19294 .unwrap()
19295 .to_string();
19296 let mut body = Box::pin(response.into_body().into_data_stream());
19297 let first = std::future::poll_fn(|cx| body.as_mut().poll_next(cx))
19298 .await
19299 .expect("stream ended before first delta")
19300 .expect("stream body failed");
19301 assert!(
19302 is_sse_data_frame(&first),
19303 "first frame was not SSE data: {first:?}"
19304 );
19305 drop(body);
19306
19307 let mut dropped = None;
19310 for _ in 0..500 {
19311 if let Some(event) = mock
19312 .events()
19313 .into_iter()
19314 .find(|e| matches!(e, MeterEvent::Dropped { .. }))
19315 {
19316 dropped = Some(event);
19317 break;
19318 }
19319 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
19320 }
19321 let events = mock.events();
19322 assert!(
19323 events
19324 .iter()
19325 .any(|e| matches!(e, MeterEvent::Open { request_id: id, .. } if id == &request_id)),
19326 "the receipt was opened under the caller-visible request id: {events:?}"
19327 );
19328 assert_eq!(
19329 dropped,
19330 Some(MeterEvent::Dropped {
19331 prompt: 1,
19332 cached: 0,
19333 completion: 1,
19334 }),
19335 "a client disconnect must leave the partial counts on the dropped receipt \
19336 (the implementation prices that drop): {events:?}"
19337 );
19338 }
19339
19340 #[tokio::test]
19341 #[allow(clippy::await_holding_lock)] async fn draining_rejects_new_requests_with_503_and_retry_after() {
19343 let _l = drain_lock();
19344 let st = fake_worker_state();
19345 let _down = DrainingRestore;
19348 DRAINING.store(true, std::sync::atomic::Ordering::SeqCst);
19349 let resp = chat_completions(
19351 State(st.clone()),
19352 axum::http::HeaderMap::new(),
19353 None,
19354 Json(
19355 serde_json::from_value(serde_json::json!({
19356 "model": "m", "messages": [{"role": "user", "content": "t"}]
19357 }))
19358 .unwrap(),
19359 ),
19360 )
19361 .await;
19362 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19363 let ra = resp.headers()["retry-after"].to_str().unwrap().to_string();
19368 let ra_s: u64 = ra
19369 .parse()
19370 .expect("Retry-After must be integer delay-seconds");
19371 assert!(
19372 ra_s > 0 && ra_s <= 60,
19373 "Retry-After {ra_s}s is outside the honored window"
19374 );
19375 let ra_ms: u64 = resp.headers()["retry-after-ms"]
19376 .to_str()
19377 .unwrap()
19378 .parse()
19379 .unwrap();
19380 assert_eq!(ra_ms, ra_s * 1000, "the two retry headers must agree");
19381 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19382 .await
19383 .unwrap();
19384 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19385 assert!(
19386 payload["error"]["message"]
19387 .as_str()
19388 .unwrap()
19389 .contains("draining")
19390 );
19391 assert_eq!(payload["error"]["type"], "server_error");
19392 assert_eq!(payload["error"]["code"], "draining");
19393 let resp = completions(
19394 State(st.clone()),
19395 axum::http::HeaderMap::new(),
19396 None,
19397 Json(
19398 serde_json::from_value(serde_json::json!({
19399 "model": "m", "prompt": "t"
19400 }))
19401 .unwrap(),
19402 ),
19403 )
19404 .await;
19405 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19406 assert!(resp.headers().contains_key("retry-after"));
19407 assert_eq!(
19408 st.inflight[0].load(std::sync::atomic::Ordering::SeqCst),
19409 0,
19410 "rejected requests must not hold slots"
19411 );
19412 let resp = health_live(State(st.clone())).await.into_response();
19415 assert_eq!(
19416 resp.status(),
19417 StatusCode::OK,
19418 "a drain must not look like a liveness fault"
19419 );
19420 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19421 .await
19422 .unwrap();
19423 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19424 assert_eq!(payload["status"], "draining");
19425 let resp = health_ready(State(st.clone())).await.into_response();
19427 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19428 let retry_s = drain_deadline_s().clamp(1, 60);
19429 let retry_s_text = retry_s.to_string();
19430 let retry_ms_text = (retry_s * 1000).to_string();
19431 assert_eq!(retry_after(&resp).as_deref(), Some(retry_s_text.as_str()));
19432 assert_eq!(
19433 resp.headers().get("retry-after-ms").unwrap(),
19434 retry_ms_text.as_str()
19435 );
19436 assert_ne!(
19437 resp.headers()
19438 .get("x-should-retry")
19439 .and_then(|v| v.to_str().ok()),
19440 Some("false")
19441 );
19442 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19443 .await
19444 .unwrap();
19445 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19446 assert_eq!(payload["status"], "not_ready");
19447 assert!(payload["detail"].as_str().unwrap().contains("draining"));
19448 DRAINING.store(false, std::sync::atomic::Ordering::SeqCst);
19449 let resp = chat_completions(
19451 State(st.clone()),
19452 axum::http::HeaderMap::new(),
19453 None,
19454 Json(
19455 serde_json::from_value(serde_json::json!({
19456 "model": "m", "messages": [{"role": "user", "content": "t"}]
19457 }))
19458 .unwrap(),
19459 ),
19460 )
19461 .await;
19462 assert_eq!(resp.status(), StatusCode::OK);
19463 }
19464
19465 #[tokio::test]
19468 #[allow(clippy::await_holding_lock)] async fn health_is_green_only_while_the_worker_is_alive() {
19470 let _l = drain_lock();
19473 let st = fake_worker_state();
19474 let resp = health_live(State(st.clone())).await.into_response();
19477 assert_eq!(resp.status(), StatusCode::OK);
19478 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19479 .await
19480 .unwrap();
19481 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19482 assert_eq!(payload["status"], "ok");
19483 assert_eq!(payload["worker"]["phase"], "idle");
19484 assert!(payload["worker"]["stall_threshold_ms"].as_u64().unwrap() > 0);
19485 let ready = health_ready(State(st.clone())).await.into_response();
19486 assert_eq!(ready.status(), StatusCode::OK);
19487
19488 st.health.mark_dead("worker thread panicked: test-injected");
19492 let resp = health_live(State(st.clone())).await.into_response();
19493 assert_eq!(
19494 resp.status(),
19495 StatusCode::SERVICE_UNAVAILABLE,
19496 "a dead worker MUST NOT report a healthy liveness"
19497 );
19498 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19499 .await
19500 .unwrap();
19501 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19502 assert_eq!(payload["status"], "unhealthy");
19503 assert!(
19505 payload["detail"]
19506 .as_str()
19507 .unwrap()
19508 .contains("test-injected"),
19509 "cause not surfaced: {payload}"
19510 );
19511 let ready = health_ready(State(st.clone())).await.into_response();
19512 assert_eq!(
19513 ready.status(),
19514 StatusCode::SERVICE_UNAVAILABLE,
19515 "dead is also not ready"
19516 );
19517
19518 st.health.mark_ready();
19521 assert_eq!(
19522 health_live(State(st.clone()))
19523 .await
19524 .into_response()
19525 .status(),
19526 StatusCode::OK,
19527 "mark_ready must clear the latch (a successful respawn)"
19528 );
19529 }
19530
19531 #[tokio::test]
19532 #[allow(clippy::await_holding_lock)] async fn readyz_peer_probe_integrity_is_present_and_advisory() {
19534 let _l = drain_lock();
19535 let st = fake_worker_state();
19536
19537 let ready = health_ready(State(st.clone())).await.into_response();
19538 assert_eq!(ready.status(), StatusCode::OK);
19539 let bytes = axum::body::to_bytes(ready.into_body(), usize::MAX)
19540 .await
19541 .unwrap();
19542 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19543 assert_eq!(payload["peer_probe_integrity"], "ok");
19544
19545 st.health.note_peer_probe_deferral(2, false);
19546 let deferred = health_ready(State(st.clone())).await.into_response();
19547 assert_eq!(deferred.status(), StatusCode::OK);
19548 let bytes = axum::body::to_bytes(deferred.into_body(), usize::MAX)
19549 .await
19550 .unwrap();
19551 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19552 assert_eq!(payload["peer_probe_integrity"], "deferred_2");
19553
19554 st.health.note_peer_probe_deferral(4, true);
19555 let degraded = health_ready(State(st.clone())).await.into_response();
19556 assert_eq!(
19557 degraded.status(),
19558 StatusCode::OK,
19559 "peer degradation is advisory while plain serving remains healthy"
19560 );
19561 let bytes = axum::body::to_bytes(degraded.into_body(), usize::MAX)
19562 .await
19563 .unwrap();
19564 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19565 assert_eq!(payload["peer_probe_integrity"], "degraded");
19566
19567 st.health.mark_dead("test-injected worker failure");
19568 let unready = health_ready(State(st)).await.into_response();
19569 assert_eq!(unready.status(), StatusCode::SERVICE_UNAVAILABLE);
19570 let bytes = axum::body::to_bytes(unready.into_body(), usize::MAX)
19571 .await
19572 .unwrap();
19573 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19574 assert_eq!(
19575 payload["peer_probe_integrity"], "degraded",
19576 "the advisory field must also survive an unrelated readiness failure"
19577 );
19578 }
19579
19580 #[tokio::test]
19581 #[allow(clippy::await_holding_lock)] async fn liveness_failure_obeys_the_retry_contract() {
19583 let _l = drain_lock();
19588 let st = fake_worker_state();
19589 st.health
19590 .mark_dead("worker thread panicked: retry-contract-test");
19591
19592 let resp = health_live(State(st)).await.into_response();
19593 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19594 assert_eq!(retry_after(&resp).as_deref(), Some("2"));
19595 assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
19596 assert_ne!(
19597 resp.headers()
19598 .get("x-should-retry")
19599 .and_then(|v| v.to_str().ok()),
19600 Some("false")
19601 );
19602 }
19603
19604 #[tokio::test]
19605 #[allow(clippy::await_holding_lock)] async fn readiness_failure_obeys_the_retry_contract() {
19607 let _l = drain_lock();
19608 let st = fake_worker_state();
19609 st.health
19610 .mark_dead("worker thread panicked: retry-contract-test");
19611
19612 let resp = health_ready(State(st)).await.into_response();
19613 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19614 assert_eq!(retry_after(&resp).as_deref(), Some("2"));
19615 assert_eq!(resp.headers().get("retry-after-ms").unwrap(), "2000");
19616 assert_ne!(
19617 resp.headers()
19618 .get("x-should-retry")
19619 .and_then(|v| v.to_str().ok()),
19620 Some("false")
19621 );
19622 }
19623
19624 #[tokio::test]
19625 #[allow(clippy::await_holding_lock)] async fn a_wedged_gpu_flips_health_even_though_the_worker_thread_is_fine() {
19627 let _l = drain_lock();
19637 let st = fake_worker_state();
19638 assert_eq!(
19639 health_live(State(st.clone()))
19640 .await
19641 .into_response()
19642 .status(),
19643 StatusCode::OK
19644 );
19645 st.health
19646 .mark_gpu_fault("nvidia-smi probe exceeded 10s deadline (GSP hang class)");
19647 let resp = health_live(State(st.clone())).await.into_response();
19648 assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
19649 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
19650 .await
19651 .unwrap();
19652 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
19653 assert!(
19654 payload["detail"]
19655 .as_str()
19656 .unwrap()
19657 .contains("probe exceeded")
19658 );
19659 st.health.mark_ready();
19662 assert_eq!(
19663 health_live(State(st.clone()))
19664 .await
19665 .into_response()
19666 .status(),
19667 StatusCode::SERVICE_UNAVAILABLE,
19668 "a GPU fault must not be cleared by an in-process respawn"
19669 );
19670 }
19671
19672 #[test]
19673 fn v1_models_entry_keeps_catalog_shape_with_honest_nulls() {
19674 let caps = ModelCaps {
19676 tools_branch: true,
19677 hy3: false,
19678 qwen_think: true,
19679 think_switch: true,
19680 chat_ok: true,
19681 context_length: 262144,
19682 tokenizer: "qwen2".into(),
19683 instruct_type: Some("chatml".into()),
19684 effort_levels: false,
19685 qwen_effort: false,
19686 gemma_think: false,
19687 dsv4: false,
19688 glm5: false,
19689 chat_temperature_default: None,
19690 chat_top_p_default: None,
19691 n_vocab: 151_936,
19692 think_close: Vec::new(),
19693 };
19694 let e = model_entry_v1("main", Some(&caps), None);
19695 assert_eq!(e["id"], "main");
19696 assert_eq!(e["name"], "main");
19697 assert_eq!(e["object"], "model");
19698 assert_eq!(e["context_length"], 262144);
19699 assert!(e["pricing"]["input"].is_null());
19701 assert!(e["pricing"]["output"].is_null());
19702
19703 let meta = OpenRouterModelMetadata {
19706 pricing: OpenRouterPricing {
19707 prompt: Some("0.00000038".into()),
19708 cached_prompt: Some("0.0000002".into()),
19709 completion: Some("0.0000026".into()),
19710 ..Default::default()
19711 },
19712 input_modalities: vec!["image".into(), "video".into()],
19713 max_output_length: Some(32768),
19714 ..Default::default()
19715 };
19716 let e = model_entry_v1("main", Some(&caps), Some(&meta));
19717 assert_eq!(e["pricing"]["currency"], "USD");
19720 assert_eq!(e["pricing"]["unit"], "per_1m_tokens");
19721 assert_eq!(e["pricing"]["input"], "0.38");
19722 assert_eq!(e["pricing"]["output"], "2.60");
19723 assert_eq!(e["pricing"]["cached_input"], "0.20");
19724 assert!(e["pricing"]["cache_write"].is_null());
19725 assert_eq!(e["pricing"]["minimum_request"], "0");
19726 assert_eq!(e["owned_by"], "main");
19727 assert_eq!(e["type"], "chat");
19728 assert_eq!(e["max_output_tokens"], 32768);
19729 assert_eq!(e["endpoints"], json!(["chat/completions"]));
19730 assert_eq!(e["input_modalities"], json!(["text", "image", "video"]));
19731 assert_eq!(e["output_modalities"], json!(["text"]));
19732 assert_eq!(e["capabilities"]["streaming"], true);
19733 assert_eq!(e["capabilities"]["tools"], true);
19734 assert_eq!(e["lifecycle"]["status"], "active");
19735 assert!(e["lifecycle"]["deprecation_at"].is_null());
19736 assert_eq!(e["reliability"]["first_token_timeout_seconds"], 120);
19737 assert_eq!(e["reliability"]["capacity_scope"], "model_region");
19738 let mut keys: Vec<&str> = e.as_object().unwrap().keys().map(String::as_str).collect();
19742 keys.sort_unstable();
19743 assert_eq!(
19744 keys,
19745 [
19746 "capabilities",
19747 "context_length",
19748 "endpoints",
19749 "id",
19750 "input_modalities",
19751 "lifecycle",
19752 "max_output_tokens",
19753 "name",
19754 "object",
19755 "output_modalities",
19756 "owned_by",
19757 "pricing",
19758 "reliability",
19759 "type",
19760 ],
19761 "unexpected /v1/models entry keys"
19762 );
19763 let mut price_keys: Vec<&str> = e["pricing"]
19764 .as_object()
19765 .unwrap()
19766 .keys()
19767 .map(String::as_str)
19768 .collect();
19769 price_keys.sort_unstable();
19770 assert_eq!(
19771 price_keys,
19772 [
19773 "cache_write",
19774 "cached_input",
19775 "currency",
19776 "input",
19777 "minimum_request",
19778 "output",
19779 "unit",
19780 ],
19781 "unexpected /v1/models pricing keys"
19782 );
19783
19784 let e = model_entry_v1("m", None, None);
19786 assert!(e["context_length"].is_null());
19787 assert!(e["max_output_tokens"].is_null());
19788 let bare = ModelCaps::default(); let e = model_entry_v1("m", Some(&bare), None);
19790 assert!(e["context_length"].is_null());
19791 }
19792
19793 #[test]
19799 fn catalog_row_follows_the_declared_surface() {
19800 let caps = ModelCaps {
19801 tools_branch: true,
19802 ..Default::default()
19803 };
19804
19805 let embed = OpenRouterModelMetadata {
19806 surface: Some("embedding".into()),
19807 max_output_length: Some(1),
19808 ..Default::default()
19809 };
19810 let e = model_entry_v1("qwen", Some(&caps), Some(&embed));
19811 assert_eq!(e["type"], "embedding");
19812 assert_eq!(e["endpoints"], json!(["embeddings"]));
19813 assert_eq!(e["output_modalities"], json!(["embeddings"]));
19814 assert_eq!(e["capabilities"]["streaming"], false);
19815 assert_eq!(
19816 e["capabilities"]["tools"], false,
19817 "an embedder has no tools"
19818 );
19819 assert_eq!(e["capabilities"]["reasoning"], false);
19820 assert_eq!(e["capabilities"]["structured_output"], false);
19821 assert_eq!(e["capabilities"]["prompt_caching"], false);
19822 assert!(
19823 e["max_output_tokens"].is_null(),
19824 "a surface that emits no completion tokens must not advertise a ceiling"
19825 );
19826
19827 let rerank = OpenRouterModelMetadata {
19828 surface: Some("rerank".into()),
19829 ..Default::default()
19830 };
19831 let r = model_entry_v1("qwen", Some(&caps), Some(&rerank));
19832 assert_eq!(r["type"], "rerank");
19833 assert_eq!(r["endpoints"], json!(["rerank"]));
19834 assert_eq!(r["output_modalities"], json!(["rerank"]));
19835 assert_eq!(r["capabilities"]["tools"], false);
19836 assert_eq!(r["capabilities"]["reasoning"], false);
19837
19838 let chat = OpenRouterModelMetadata {
19841 max_output_length: Some(32768),
19842 ..Default::default()
19843 };
19844 let c = model_entry_v1("main", Some(&caps), Some(&chat));
19845 assert_eq!(c["type"], "chat");
19846 assert_eq!(c["endpoints"], json!(["chat/completions"]));
19847 assert_eq!(c["output_modalities"], json!(["text"]));
19848 assert_eq!(c["capabilities"]["tools"], true);
19849 assert_eq!(c["max_output_tokens"], 32768);
19850 }
19851
19852 #[test]
19855 fn unknown_surface_is_rejected_at_config_load() {
19856 let bad = OpenRouterModelMetadata {
19857 surface: Some("embeddings".into()), ..Default::default()
19859 };
19860 let err = validate_openrouter_metadata("qwen/qwen3-embedding-8b", &bad)
19861 .expect_err("an unknown surface must not load");
19862 assert!(err.contains("surface"), "{err}");
19863
19864 for good in ["chat", "embedding", "rerank"] {
19865 let ok = OpenRouterModelMetadata {
19866 surface: Some(good.into()),
19867 ..Default::default()
19868 };
19869 assert!(
19870 validate_openrouter_metadata("m", &ok).is_ok(),
19871 "{good} must load"
19872 );
19873 }
19874 }
19875
19876 #[test]
19877 fn per_million_price_is_exact_decimal_shift() {
19878 assert_eq!(per_million_price("0.00000038").as_deref(), Some("0.38"));
19880 assert_eq!(per_million_price("0.0000026").as_deref(), Some("2.60"));
19881 assert_eq!(per_million_price("0.0000002").as_deref(), Some("0.20"));
19882 assert_eq!(per_million_price("0").as_deref(), Some("0.00"));
19883 assert_eq!(per_million_price("1.5").as_deref(), Some("1500000.00"));
19884 assert_eq!(per_million_price("0.000000125").as_deref(), Some("0.125"));
19885 assert_eq!(per_million_price("not-a-price"), None);
19886 assert_eq!(per_million_price(""), None);
19887 }
19888
19889 #[test]
19890 fn metadata_provider_block_parses_and_validates() {
19891 let (_, provider) = OpenRouterMetadataFile::parse(
19892 r#"
19893 [provider]
19894 id = "tiyuvta"
19895 status_url = "https://status.tiyuvta.ai"
19896 support_contact = "mailto:support@tiyuvta.ai"
19897 incident_contact = "mailto:incidents@tiyuvta.ai"
19898 regions = ["eu-central"]
19899 "#,
19900 )
19901 .unwrap();
19902 let provider = provider.unwrap();
19903 assert_eq!(provider.id, "tiyuvta");
19904 assert_eq!(provider.regions, vec!["eu-central"]);
19905 let err = OpenRouterMetadataFile::parse("[provider]\nid = \"\"\n").unwrap_err();
19907 assert!(err.contains("provider.id"), "{err}");
19908 let err = OpenRouterMetadataFile::parse(
19910 "[provider]\nid = \"x\"\nsupport_contact = \"ops@example.com\"\n",
19911 )
19912 .unwrap_err();
19913 assert!(err.contains("must be a URI"), "{err}");
19914 let (_, provider) = OpenRouterMetadataFile::parse("").unwrap();
19916 assert!(provider.is_none());
19917 }
19918
19919 #[test]
19920 fn models_openai_default_body_stays_byte_identical() {
19921 let body = models_openai_body(&["main".into(), "judge".into()]);
19922 let bytes = serde_json::to_vec(&body).unwrap();
19923 assert_eq!(
19924 bytes,
19925 br#"{"object":"list","data":[{"id":"main","object":"model"},{"id":"judge","object":"model"}]}"#
19926 );
19927 }
19928
19929 #[test]
19930 fn canonical_model_id_tolerates_a_marketplace_stripping_the_vendor_prefix() {
19931 let loaded = vec![
19933 "qwen/qwen3.6-27b".to_string(),
19934 "qwen/qwen3.6-35b-a3b".to_string(),
19935 ];
19936 assert_eq!(
19937 canonical_model_id(&loaded, "qwen3.6-35b-a3b").as_deref(),
19938 Some("qwen/qwen3.6-35b-a3b"),
19939 );
19940 assert_eq!(
19941 canonical_model_id(&loaded, "qwen3.6-27b").as_deref(),
19942 Some("qwen/qwen3.6-27b"),
19943 );
19944 assert_eq!(
19946 canonical_model_id(&loaded, "qwen/qwen3.6-35b-a3b").as_deref(),
19947 Some("qwen/qwen3.6-35b-a3b"),
19948 );
19949 assert_eq!(canonical_model_id(&loaded, "gpt-4o"), None);
19951 assert_eq!(canonical_model_id(&loaded, "vendor/qwen3.6-35b-a3b"), None);
19952 assert_eq!(canonical_model_id(&loaded, ""), None);
19953 }
19954
19955 #[test]
19956 fn canonical_model_id_refuses_an_ambiguous_suffix_rather_than_guessing() {
19957 let loaded = vec!["a/shared-name".to_string(), "b/shared-name".to_string()];
19960 assert_eq!(canonical_model_id(&loaded, "shared-name"), None);
19961 assert_eq!(
19963 canonical_model_id(&loaded, "a/shared-name").as_deref(),
19964 Some("a/shared-name")
19965 );
19966 assert_eq!(
19967 canonical_model_id(&loaded, "b/shared-name").as_deref(),
19968 Some("b/shared-name")
19969 );
19970 let bare = vec!["solo".to_string()];
19972 assert_eq!(canonical_model_id(&bare, "solo").as_deref(), Some("solo"));
19973 }
19974
19975 #[test]
19976 fn openrouter_models_entry_serializes_complete_metadata() {
19977 let metadata = OpenRouterMetadataFile::from_toml(
19978 r#"
19979[models.main]
19980hugging_face_id = "Qwen/Qwen3.6-27B"
19981created = 1786032000
19982quantization = "nvfp4"
19983description = "Qwen3.6 27B served by memra."
19984max_prompt_length = 245760
19985max_output_length = 16384
19986default_output_length = 4096
19987is_ready = true
19988is_free = false
19989discount_to_user = 0.1
19990openrouter_slug = "qwen/qwen3.6-27b"
19991datacenters = [{ country_code = "US", region = "us-east" }]
19992zdr = true
19993hipaa = false
19994
19995[models.main.pricing]
19996prompt = "0.000000234"
19997cached_prompt = "0.0000000585"
19998cache_write = "0.000000234"
19999completion = "0.000001872"
20000internal_reasoning = "0.000001872"
20001request = "0.01"
20002
20003[models.main.capacity]
20004prompt_tpm = 1000000
20005cached_prompt_tpm = 2000000
20006completion_tpm = 500000
20007request_rpm = 1000
20008concurrency = 64
20009"#,
20010 )
20011 .unwrap();
20012 let caps = ModelCaps {
20013 tools_branch: true,
20014 qwen_think: true,
20015 think_switch: true,
20016 chat_ok: true,
20017 context_length: 262144,
20018 tokenizer: "qwen2".into(),
20019 instruct_type: Some("chatml".into()),
20020 ..Default::default()
20021 };
20022 let entry = model_entry_openrouter("main", Some(&caps), metadata.get("main"));
20023
20024 assert_eq!(entry["schema_version"], "2.4");
20025 assert_eq!(entry["id"], "main");
20026 assert_eq!(entry["name"], "main");
20027 assert_eq!(entry["hugging_face_id"], "Qwen/Qwen3.6-27B");
20028 assert_eq!(entry["created"], 1786032000u64);
20029 assert_eq!(entry["quantization"], "nvfp4");
20030 assert_eq!(entry["tokenizer"], "qwen2");
20031 assert_eq!(entry["description"], "Qwen3.6 27B served by memra.");
20032 assert!(
20033 entry.get("object").is_none(),
20034 "OpenRouter schema 2.4 rejects unknown OpenAI fields"
20035 );
20036
20037 let input = &entry["input_modalities"][0];
20038 assert_eq!(input["type"], "text");
20039 assert_eq!(
20040 input["supported_inputs"]["max_context_length"]["value"],
20041 262144
20042 );
20043 assert_eq!(
20044 input["supported_inputs"]["max_prompt_length"]["value"],
20045 245760
20046 );
20047 let input_prices = input["pricing"].as_array().unwrap();
20048 let input_price = |kind: &str| {
20049 input_prices
20050 .iter()
20051 .find(|price| price["type"] == kind)
20052 .unwrap()
20053 };
20054 assert_eq!(input_price("prompt")["cost_usd"], "0.000000234");
20055 assert_eq!(input_price("cached_prompt")["cost_usd"], "0.0000000585");
20056 assert_eq!(input_price("cache_write")["cost_usd"], "0.000000234");
20057 assert_eq!(input["capacity"][0]["value"], 1000000);
20058 assert_eq!(input["capacity"][1]["value"], 2000000);
20059
20060 let output = &entry["output_modalities"][0];
20061 assert_eq!(output["type"], "text");
20062 assert_eq!(output["max_length"]["value"], 16384);
20063 assert_eq!(output["streaming"], true);
20064 assert_eq!(output["supported_parameters"]["tools"]["type"], "boolean");
20065 assert_eq!(
20066 output["supported_parameters"]["structured_outputs"]["type"],
20067 "boolean"
20068 );
20069 assert_eq!(
20070 output["supported_parameters"]["reasoning"]["type"],
20071 "boolean"
20072 );
20073 assert_eq!(output["pricing"][0]["type"], "completion");
20074 assert_eq!(output["pricing"][0]["cost_usd"], "0.000001872");
20075 assert_eq!(output["pricing"][1]["type"], "internal_reasoning");
20076 assert_eq!(output["capacity"][0]["value"], 500000);
20077 assert_eq!(output["capacity"][1]["type"], "concurrency");
20078 assert_eq!(output["capacity"][1]["value"], 64);
20079
20080 assert_eq!(entry["pricing"][0]["type"], "request");
20081 assert_eq!(entry["pricing"][0]["cost_usd"], "0.01");
20082 assert_eq!(entry["capacity"][0]["value"], 1000);
20083 assert_eq!(entry["is_ready"], true);
20084 assert_eq!(entry["is_free"], false);
20085 assert_eq!(entry["discount_to_user"], 0.1);
20086 assert_eq!(entry["openrouter"]["slug"], "qwen/qwen3.6-27b");
20087 assert_eq!(entry["datacenters"][0]["country_code"], "US");
20088 assert_eq!(entry["compliance"]["zdr"], true);
20089 assert_eq!(entry["compliance"]["hipaa"], false);
20090 }
20091
20092 const GATEWAY_REGISTRY_FIXTURE: &str = r#"
20097[models."qwen/qwen3.6-35b-a3b"]
20098hugging_face_id = "Qwen/Qwen3.6-35B-A3B"
20099created = 1777260255
20100quantization = "int4"
20101description = "Qwen3.6 35B-A3B fixture entry."
20102max_prompt_length = 262144
20103max_output_length = 262144
20104default_output_length = 8192
20105is_ready = true
20106is_free = false
20107discount_to_user = 0.0
20108openrouter_slug = "qwen/qwen3.6-35b-a3b"
20109zdr = false
20110hipaa = false
20111
20112[[models."qwen/qwen3.6-35b-a3b".datacenters]]
20113country_code = "CA"
20114region = "Ontario"
20115
20116[models."qwen/qwen3.6-35b-a3b".pricing]
20117prompt = "0.0000000931"
20118cached_prompt = "0.0000000652"
20119completion = "0.0000009025"
20120
20121[models."qwen/qwen3.6-35b-a3b".capacity]
20122prompt_tpm = 780000
20123cached_prompt_tpm = 310000
20124completion_tpm = 9600
20125request_rpm = 160
20126concurrency = 16
20127
20128[planned_models."qwen/qwen3.8-27b"]
20129description = "Planned fixture entry; must never be emitted."
20130max_prompt_length = 262144
20131max_output_length = 262144
20132default_output_length = 8192
20133is_ready = false
20134is_free = false
20135discount_to_user = 0.0
20136openrouter_slug = "qwen/qwen3.8-27b"
20137zdr = false
20138hipaa = false
20139
20140[planned_models."qwen/qwen3.8-27b".pricing]
20141prompt = "0.0000002745"
20142cached_prompt = "0.0000001922"
20143completion = "0.0000022800"
20144
20145[planned_models."google/gemma-4-26b-a4b-it"]
20146hugging_face_id = "google/gemma-4-26B-A4B-it"
20147created = 1775227989
20148quantization = "int4"
20149description = "Planned fixture entry; must never be emitted."
20150max_prompt_length = 262144
20151max_output_length = 262144
20152default_output_length = 8192
20153is_ready = false
20154is_free = false
20155discount_to_user = 0.0
20156openrouter_slug = "google/gemma-4-26b-a4b-it"
20157zdr = false
20158hipaa = false
20159
20160[planned_models."google/gemma-4-26b-a4b-it".pricing]
20161prompt = "0.0000000665"
20162cached_prompt = "0.0000000466"
20163completion = "0.0000003230"
20164"#;
20165
20166 #[test]
20167 fn gateway_registry_generates_the_staged_active_shape() {
20168 let metadata = OpenRouterMetadataFile::from_toml(GATEWAY_REGISTRY_FIXTURE).unwrap();
20169 let caps = ModelCaps {
20170 tools_branch: true,
20171 qwen_think: true,
20172 think_switch: true,
20173 chat_ok: true,
20174 context_length: 262144,
20175 tokenizer: "qwen2".into(),
20176 instruct_type: Some("chatml".into()),
20177 ..Default::default()
20178 };
20179 let q35_entry = model_entry_openrouter(
20180 "qwen/qwen3.6-35b-a3b",
20181 Some(&caps),
20182 metadata.get("qwen/qwen3.6-35b-a3b"),
20183 );
20184 assert_eq!(q35_entry["created"], 1777260255u64);
20185 assert_eq!(q35_entry["quantization"], "int4");
20186 assert_eq!(q35_entry["is_ready"], true);
20187 assert_eq!(
20188 q35_entry["input_modalities"][0]["supported_inputs"]["max_context_length"]["value"],
20189 262144
20190 );
20191 assert_eq!(
20192 q35_entry["input_modalities"][0]["supported_inputs"]["max_prompt_length"]["value"],
20193 262144
20194 );
20195 assert_eq!(
20196 q35_entry["output_modalities"][0]["max_length"]["value"],
20197 262144
20198 );
20199 let prices = q35_entry["input_modalities"][0]["pricing"]
20200 .as_array()
20201 .unwrap();
20202 assert_eq!(prices[0]["cost_usd"], "0.0000000931");
20203 assert_eq!(prices[1]["cost_usd"], "0.0000000652");
20204 assert_eq!(
20208 q35_entry["input_modalities"][0]["capacity"][0]["value"],
20209 780000
20210 );
20211 assert_eq!(
20212 q35_entry["input_modalities"][0]["capacity"][1]["value"],
20213 310000
20214 );
20215 assert_eq!(
20216 q35_entry["output_modalities"][0]["supported_parameters"]["max_tokens"]["max"],
20217 262144
20218 );
20219 assert_eq!(
20220 q35_entry["output_modalities"][0]["capacity"][0]["value"],
20221 9600
20222 );
20223 assert_eq!(
20224 q35_entry["output_modalities"][0]["capacity"][1]["value"],
20225 16
20226 );
20227 assert_eq!(
20228 q35_entry["output_modalities"][0]["pricing"][0]["cost_usd"],
20229 "0.0000009025"
20230 );
20231 assert_eq!(q35_entry["capacity"][0]["value"], 160); assert_eq!(q35_entry["datacenters"][0]["country_code"], "CA");
20233
20234 assert_eq!(
20235 metadata.len(),
20236 1,
20237 "planned models must never enter the active map"
20238 );
20239 assert!(!metadata.contains_key("qwen/qwen3.6-27b"));
20240 assert!(!metadata.contains_key("qwen/qwen3.8-27b"));
20241 assert!(!metadata.contains_key("google/gemma-4-26b-a4b-it"));
20242
20243 let openmodels = model_entry_openmodels(
20244 "qwen/qwen3.6-35b-a3b",
20245 Some(&caps),
20246 metadata.get("qwen/qwen3.6-35b-a3b"),
20247 )
20248 .unwrap();
20249 assert_eq!(openmodels["currency"], "USD");
20250 assert_eq!(openmodels["max_output_length"], 262144);
20251 assert_eq!(openmodels["is_ready"], true);
20252 assert_eq!(openmodels["is_free"], false);
20253 assert_eq!(openmodels["discount_to_user"], 0.0);
20254 }
20255
20256 #[test]
20257 fn gateway_registry_limits_are_live_request_limits() {
20258 let metadata_file = OpenRouterMetadataFile::from_toml(GATEWAY_REGISTRY_FIXTURE).unwrap();
20259 let metadata = metadata_file.get("qwen/qwen3.6-35b-a3b").unwrap();
20260 let caps = ModelCaps {
20261 context_length: 262_144,
20262 ..Default::default()
20263 };
20264 let build = |value: serde_json::Value| {
20265 let req: CompletionReq = serde_json::from_value(value).unwrap();
20266 let (tx, _rx) = worker::event_channel();
20267 build_request(&req, tx, lanes::Lane::Interactive, None)
20268 };
20269
20270 let mut omitted = build(json!({
20271 "model": "qwen/qwen3.6-35b-a3b",
20272 "prompt_ids": [1, 2, 3]
20273 }));
20274 apply_model_request_limits(&mut omitted, Some(metadata), Some(&caps)).unwrap();
20275 assert_eq!(omitted.params.max_new, 8_192);
20276 assert_eq!(omitted.max_prompt_tokens, Some(262_144));
20277
20278 let mut field_top = build(json!({
20279 "model": "qwen/qwen3.6-35b-a3b",
20280 "prompt_ids": [1],
20281 "max_tokens": 262144
20282 }));
20283 apply_model_request_limits(&mut field_top, Some(metadata), Some(&caps)).unwrap();
20284 assert_eq!(field_top.params.max_new, 262_144);
20285 assert_eq!(
20286 budget_completion_bound(&field_top, 100, Some(&caps)).unwrap(),
20287 262_044,
20288 "the field-top output request is accepted but bounded by remaining trained context",
20289 );
20290
20291 let mut too_much_output = build(json!({
20292 "model": "qwen/qwen3.6-35b-a3b",
20293 "prompt_ids": [1],
20294 "max_tokens": 262145
20295 }));
20296 let (message, param) =
20297 apply_model_request_limits(&mut too_much_output, Some(metadata), Some(&caps))
20298 .unwrap_err();
20299 assert_eq!(param, "max_tokens");
20300 assert!(message.contains("262145"));
20301
20302 let mut oversized_allocation = build(json!({
20303 "model": "qwen/qwen3.6-35b-a3b",
20304 "prompt_ids": [1],
20305 "max_tokens": 1,
20306 "max_ctx": 262145
20307 }));
20308 let (_, param) =
20309 apply_model_request_limits(&mut oversized_allocation, Some(metadata), Some(&caps))
20310 .unwrap_err();
20311 assert_eq!(param, "max_ctx");
20312 }
20313
20314 #[test]
20315 fn planned_registry_entries_are_validated_but_never_activated() {
20316 let parsed = OpenRouterMetadataFile::from_toml(
20317 r#"
20318[planned_models.future]
20319max_output_length = 262144
20320default_output_length = 8192
20321
20322[planned_models.future.pricing]
20323prompt = "0.0000001"
20324"#,
20325 )
20326 .unwrap();
20327 assert!(parsed.is_empty());
20328
20329 let error = OpenRouterMetadataFile::from_toml(
20330 r#"
20331[planned_models.future]
20332default_output_length = 8192
20333"#,
20334 )
20335 .unwrap_err();
20336 assert!(error.contains("requires max_output_length"));
20337 }
20338
20339 #[test]
20344 fn every_catalog_feed_honours_the_declared_surface() {
20345 let metadata = OpenRouterMetadataFile::from_toml(
20346 r#"
20347[models."qwen/qwen3-embedding-8b"]
20348surface = "embedding"
20349created = 1787961600
20350max_output_length = 1
20351is_ready = true
20352is_free = false
20353discount_to_user = 0.0
20354
20355[models."qwen/qwen3-embedding-8b".pricing]
20356prompt = "0.00000001"
20357cached_prompt = "0.0"
20358completion = "0.0"
20359
20360[models."main"]
20361created = 1787443200
20362max_output_length = 32768
20363is_ready = true
20364is_free = false
20365discount_to_user = 0.0
20366
20367[models."main".pricing]
20368prompt = "0.00000025"
20369cached_prompt = "0.00000009"
20370completion = "0.0000012"
20371"#,
20372 )
20373 .unwrap();
20374 let caps = ModelCaps {
20375 tools_branch: true,
20376 qwen_think: true,
20377 think_switch: true,
20382 chat_ok: true,
20383 context_length: 32768,
20384 ..Default::default()
20385 };
20386 let embed = metadata.get("qwen/qwen3-embedding-8b");
20387 let chat = metadata.get("main");
20388
20389 let or = model_entry_openrouter("qwen/qwen3-embedding-8b", Some(&caps), embed);
20391 let out = &or["output_modalities"][0];
20392 assert_eq!(out["type"], "embeddings", "openrouter feed: {or}");
20393 assert!(
20394 out.get("streaming").is_none(),
20395 "the embeddings branch declares no streaming property (additionalProperties:false): {out}"
20396 );
20397 let params = &out["supported_parameters"];
20402 assert_eq!(
20403 params.as_object().map(|o| o.len()),
20404 Some(0),
20405 "no completion parameter belongs on an embedder row: {params}"
20406 );
20407 for field in [
20408 "tools",
20409 "tool_choice",
20410 "reasoning",
20411 "max_tokens",
20412 "json_mode",
20413 "structured_outputs",
20414 "stop",
20415 "temperature",
20416 "seed",
20417 ] {
20418 assert!(params[field].is_null(), "{field} leaked onto an embedder");
20419 }
20420 assert!(
20421 out["max_length"].is_null(),
20422 "a surface emitting no completion tokens advertises no ceiling: {out}"
20423 );
20424
20425 let om = model_entry_openmodels("qwen/qwen3-embedding-8b", Some(&caps), embed)
20427 .expect("openmodels entry builds");
20428 assert_eq!(om["output_modalities"], json!(["embeddings"]));
20429 let features = om["supported_features"].as_array().unwrap();
20430 assert!(
20431 !features
20432 .iter()
20433 .any(|f| f == "tool_calling" || f == "reasoning"),
20434 "chat-only features leaked onto an embedder: {features:?}"
20435 );
20436
20437 let v1 = model_entry_v1("qwen/qwen3-embedding-8b", Some(&caps), embed);
20439 assert_eq!(v1["type"], "embedding");
20440 assert_eq!(v1["capabilities"]["tools"], false);
20441
20442 let or_chat = model_entry_openrouter("main", Some(&caps), chat);
20444 let out_chat = &or_chat["output_modalities"][0];
20445 assert_eq!(out_chat["type"], "text");
20446 assert_eq!(out_chat["streaming"], true);
20447 assert!(!out_chat["supported_parameters"]["tools"].is_null());
20448 assert!(!out_chat["supported_parameters"]["max_tokens"].is_null());
20449 assert!(!out_chat["supported_parameters"]["structured_outputs"].is_null());
20450 assert_eq!(out_chat["max_length"]["value"], 32768u64);
20451 let om_chat = model_entry_openmodels("main", Some(&caps), chat).expect("chat entry builds");
20452 assert_eq!(om_chat["output_modalities"], json!(["text"]));
20453 assert!(
20454 om_chat["supported_features"]
20455 .as_array()
20456 .unwrap()
20457 .iter()
20458 .any(|f| f == "tool_calling")
20459 );
20460 assert_eq!(model_entry_v1("main", Some(&caps), chat)["type"], "chat");
20461 }
20462
20463 #[test]
20470 fn openrouter_output_modality_matches_the_vendored_2_4_schema() {
20471 let raw = std::fs::read_to_string(concat!(
20472 env!("CARGO_MANIFEST_DIR"),
20473 "/../../research/gateway-20260812/raw/sources/",
20474 "openrouter-provider-schema-v2.4-20260812.json"
20475 ))
20476 .expect("vendored Provider Monitor 2.4 schema is in-tree");
20477 let schema: serde_json::Value = serde_json::from_str(&raw).expect("schema parses");
20478 let branches = schema["components"]["schemas"]["OutputModality"]["oneOf"]
20479 .as_array()
20480 .expect("OutputModality is a oneOf");
20481
20482 let metadata = OpenRouterMetadataFile::from_toml(
20483 r#"
20484[models."embed"]
20485surface = "embedding"
20486created = 1787961600
20487max_output_length = 1
20488is_ready = true
20489is_free = false
20490discount_to_user = 0.0
20491
20492[models."embed".pricing]
20493prompt = "0.00000001"
20494cached_prompt = "0.0"
20495completion = "0.0"
20496
20497[models."rr"]
20498surface = "rerank"
20499created = 1787961600
20500max_output_length = 1
20501is_ready = true
20502is_free = false
20503discount_to_user = 0.0
20504
20505[models."rr".pricing]
20506prompt = "0.00000003"
20507cached_prompt = "0.0"
20508completion = "0.0"
20509
20510[models."chatty"]
20511created = 1787443200
20512max_output_length = 32768
20513is_ready = true
20514is_free = false
20515discount_to_user = 0.0
20516
20517[models."chatty".pricing]
20518prompt = "0.00000025"
20519cached_prompt = "0.00000009"
20520completion = "0.0000012"
20521"#,
20522 )
20523 .unwrap();
20524 let caps = ModelCaps {
20525 tools_branch: true,
20526 qwen_think: true,
20527 chat_ok: true,
20528 context_length: 32768,
20529 ..Default::default()
20530 };
20531
20532 for (alias, want_type) in [
20533 ("embed", "embeddings"),
20534 ("rr", "rerank"),
20535 ("chatty", "text"),
20536 ] {
20537 let row = model_entry_openrouter(alias, Some(&caps), metadata.get(alias));
20538 let modality = &row["output_modalities"][0];
20539 assert_eq!(modality["type"], want_type, "{alias}: {row}");
20540
20541 let branch = branches
20543 .iter()
20544 .find(|b| b["properties"]["type"]["enum"][0] == want_type)
20545 .unwrap_or_else(|| panic!("{want_type:?} is not an OutputModality branch"));
20546 let allowed: std::collections::BTreeSet<&str> = branch["properties"]
20547 .as_object()
20548 .expect("branch properties")
20549 .keys()
20550 .map(String::as_str)
20551 .collect();
20552 for key in modality.as_object().expect("modality object").keys() {
20553 assert!(
20554 allowed.contains(key.as_str()),
20555 "{alias}: {key:?} is not a property of the {want_type:?} branch \
20556 (additionalProperties:false); allowed = {allowed:?}"
20557 );
20558 }
20559 for req in branch["required"].as_array().into_iter().flatten() {
20560 let req = req.as_str().expect("required entry is a string");
20561 assert!(
20562 modality.get(req).is_some(),
20563 "{alias}: required property {req:?} missing from the {want_type:?} branch"
20564 );
20565 }
20566 }
20567 }
20568
20569 #[test]
20570 fn openrouter_models_entry_omits_undeclared_optional_fields() {
20571 let entry = model_entry_openrouter("minimal", None, None);
20572 let object = entry.as_object().unwrap();
20573 for field in [
20574 "hugging_face_id",
20575 "created",
20576 "quantization",
20577 "tokenizer",
20578 "description",
20579 "pricing",
20580 "capacity",
20581 "is_ready",
20582 "is_free",
20583 "discount_to_user",
20584 "openrouter",
20585 "datacenters",
20586 "compliance",
20587 ] {
20588 assert!(
20589 !object.contains_key(field),
20590 "optional field {field} must be absent, not null"
20591 );
20592 }
20593 assert_eq!(entry["schema_version"], "2.4");
20594 assert_eq!(entry["input_modalities"][0]["type"], "text");
20595 assert!(
20596 entry["input_modalities"][0]
20597 .get("supported_inputs")
20598 .is_none()
20599 );
20600 assert!(entry["input_modalities"][0].get("pricing").is_none());
20601 assert!(entry["input_modalities"][0].get("capacity").is_none());
20602 assert_eq!(entry["output_modalities"][0]["type"], "text");
20603 assert_eq!(entry["output_modalities"][0]["streaming"], true);
20604 assert!(entry["output_modalities"][0]["supported_parameters"].is_object());
20605 assert!(entry["output_modalities"][0].get("max_length").is_none());
20606 assert!(entry["output_modalities"][0].get("pricing").is_none());
20607 assert!(entry["output_modalities"][0].get("capacity").is_none());
20608 }
20609
20610 #[test]
20611 fn openmodels_entry_serializes_standard_provider_shape() {
20612 let metadata = OpenRouterMetadataFile::from_toml(
20613 r#"
20614[models."qwen/qwen3.6-27b"]
20615created = 1786032000
20616max_output_length = 16384
20617is_ready = true
20618is_free = false
20619discount_to_user = 0.05
20620
20621[models."qwen/qwen3.6-27b".pricing]
20622prompt = "0.000000291"
20623cached_prompt = "0.000000291"
20624completion = "0.000002763"
20625request = "0"
20626"#,
20627 )
20628 .unwrap();
20629 let caps = ModelCaps {
20630 tools_branch: true,
20631 qwen_think: true,
20632 chat_ok: true,
20633 context_length: 262144,
20634 ..Default::default()
20635 };
20636 let entry = model_entry_openmodels(
20637 "qwen/qwen3.6-27b",
20638 Some(&caps),
20639 metadata.get("qwen/qwen3.6-27b"),
20640 )
20641 .unwrap();
20642
20643 assert_eq!(entry["id"], "qwen/qwen3.6-27b");
20644 assert_eq!(entry["name"], "qwen/qwen3.6-27b");
20645 assert_eq!(entry["created"], 1786032000u64);
20646 assert_eq!(entry["input_modalities"], json!(["text"]));
20647 assert_eq!(entry["output_modalities"], json!(["text"]));
20648 assert_eq!(entry["context_length"], 262144u64);
20649 assert_eq!(entry["max_output_length"], 16384u64);
20650 assert_eq!(entry["currency"], "USD");
20651 assert_eq!(entry["pricing"]["prompt"], "0.000000291");
20652 assert_eq!(entry["pricing"]["completion"], "0.000002763");
20653 assert_eq!(entry["pricing"]["input_cache_read"], "0.000000291");
20654 assert_eq!(entry["pricing"]["request"], "0");
20655 assert_eq!(
20656 entry["supported_features"],
20657 json!(["tool_calling", "reasoning"])
20658 );
20659 assert_eq!(entry["is_ready"], true);
20660 assert_eq!(entry["is_free"], false);
20661 assert_eq!(entry["discount_to_user"], 0.05);
20662 assert!(entry.get("schema_version").is_none());
20663 assert!(entry.get("quantization").is_none());
20664 }
20665
20666 #[test]
20667 fn openmodels_entry_rejects_missing_operator_metadata() {
20668 let caps = ModelCaps {
20669 context_length: 262144,
20670 ..Default::default()
20671 };
20672 let error = model_entry_openmodels("qwen/qwen3.6-27b", Some(&caps), None).unwrap_err();
20673 assert_eq!(
20674 error,
20675 "OpenModels feed requires MEMRA_MODEL_METADATA for model \"qwen/qwen3.6-27b\""
20676 );
20677 }
20678
20679 #[tokio::test]
20680 async fn blocking_response_excludes_stop_text_across_token_events() {
20681 let (tx, rx) = worker::event_channel();
20682 tx.send(Event::Token {
20683 id: 1,
20684 text: "answer\nPro".into(),
20685 })
20686 .unwrap();
20687 tx.send(Event::Token {
20688 id: 2,
20689 text: "blem: leaked prompt".into(),
20690 })
20691 .unwrap();
20692 tx.send(Event::Done {
20693 stop_reason: "Callback".into(),
20694 n_tokens: 2,
20695 n_prompt: 8,
20696 n_cached: 0,
20697 elapsed_s: 0.5,
20698 spec: None,
20699 })
20700 .unwrap();
20701 drop(tx);
20702 let response = blocking_response(
20703 rx,
20704 "plain_quant".into(),
20705 false,
20706 vec!["Problem:".into()],
20707 None,
20708 Envelope::new(false),
20709 )
20710 .await;
20711 assert_eq!(response.status(), StatusCode::OK);
20712 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
20713 .await
20714 .unwrap();
20715 let payload: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
20716 assert_eq!(payload["text"], "answer\n");
20717 assert_eq!(payload["stop_reason"], "Callback");
20718 }
20719
20720 #[test]
20724 fn step_walker_expansion_and_separator_law() {
20725 const PNG64: &str = "iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAY0lEQVR4nO3PQQ3AIADAQEANmlCD9IngcVnSU9DOe/b4s6UDXjWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgNaA1oDWgfeKYAYIDsx/LAAAAAElFTkSuQmCC";
20727 let uri = format!("data:image/png;base64,{PNG64}");
20728 let content = serde_json::json!([
20729 {"type": "text", "text": "look at"},
20730 {"type": "text", "text": "this:"},
20731 {"type": "image_url", "image_url": {"url": uri}},
20732 {"type": "text", "text": "what is it?"},
20733 ]);
20734 let mut pending: Vec<PendingStepImage> = Vec::new();
20735 let out = content_to_text_vision_step(&content, &mut pending).unwrap();
20736 let mut expansion = String::from("<im_start>");
20737 for _ in 0..memra_engine::vision_step::SV_MAIN_ROWS {
20738 expansion.push_str("<im_patch>");
20739 }
20740 expansion.push_str("<im_end>");
20741 assert_eq!(out, format!("look at this:{expansion}what is it?"));
20744 assert_eq!(pending.len(), 1);
20745 assert_eq!(pending[0].plan.n_tiles, 0);
20746 assert_eq!(pending[0].plan.n_prompt_tokens(), 171);
20747
20748 let vid = serde_json::json!([{ "type": "video_url", "video_url": {"url": uri} }]);
20750 assert!(content_to_text_vision_step(&vid, &mut Vec::new()).is_err());
20751 let http = serde_json::json!([
20752 {"type": "image_url", "image_url": {"url": "http://example.com/x.png"}}
20753 ]);
20754 assert!(content_to_text_vision_step(&http, &mut Vec::new()).is_err());
20755 }
20756}
20757
20758#[cfg(test)]
20764mod build_identity_tests {
20765 use super::{BUILD_GIT_SHA, BUILD_ID_NOTE, BUILD_ID_SRC, SYSTEM_FINGERPRINT, build_id};
20766
20767 #[test]
20769 fn baked_fingerprint_is_real_and_well_formed() {
20770 assert!(!SYSTEM_FINGERPRINT.is_empty());
20771 assert_ne!(SYSTEM_FINGERPRINT, "memra-unknown");
20772 assert!(
20773 !SYSTEM_FINGERPRINT.contains("unknown"),
20774 "fingerprint {SYSTEM_FINGERPRINT:?} still carries the degraded literal"
20775 );
20776 assert!(
20777 build_id::fingerprint_is_well_formed(SYSTEM_FINGERPRINT),
20778 "fingerprint {SYSTEM_FINGERPRINT:?} is not memra-<version>-<12 hex>"
20779 );
20780 assert!(
20783 SYSTEM_FINGERPRINT.starts_with(concat!("memra-", env!("CARGO_PKG_VERSION"), "-")),
20784 "fingerprint {SYSTEM_FINGERPRINT:?} does not name this crate version"
20785 );
20786 }
20787
20788 #[test]
20791 fn the_shape_check_rejects_what_shipped_to_prod() {
20792 assert!(!build_id::fingerprint_is_well_formed("memra-unknown"));
20793 assert!(!build_id::fingerprint_is_well_formed(
20794 "memra-0.123.0-unknown"
20795 ));
20796 let old_form = format!("memra-{}", "0".repeat(12));
20801 assert!(!build_id::fingerprint_is_well_formed(&old_form));
20802 assert!(!build_id::fingerprint_is_well_formed(""));
20803 assert!(!build_id::fingerprint_is_well_formed("memra-"));
20804 assert!(!build_id::fingerprint_is_well_formed("memra-0.123.0-"));
20805 assert!(!build_id::fingerprint_is_well_formed("memra-0.123.0-abc"));
20807 assert!(!build_id::fingerprint_is_well_formed(
20808 "memra-0.123.0-ABCDEF012345"
20809 ));
20810 assert!(!build_id::fingerprint_is_well_formed(
20811 "memra-0.123.0-zzzzzzzzzzzz"
20812 ));
20813 assert!(build_id::fingerprint_is_well_formed(
20815 "memra-0.123.0-4b1f9c02d7a3"
20816 ));
20817 }
20818
20819 #[test]
20827 fn build_id_is_rederivable_from_the_source_tree() {
20828 let root = build_id::workspace_root(env!("CARGO_MANIFEST_DIR"));
20829 let scan = root.as_deref().and_then(build_id::content_id);
20830 match scan {
20831 Some(scan) => {
20832 assert_eq!(
20833 BUILD_ID_SRC,
20834 build_id::BUILD_ID_SRC_TREE,
20835 "the source tree is readable, so the baked id must come from it"
20836 );
20837 assert!(BUILD_ID_NOTE.is_empty(), "note set on a non-degraded build");
20838 let expected =
20839 format!(concat!("memra-", env!("CARGO_PKG_VERSION"), "-{}"), scan.id);
20840 assert_eq!(
20841 SYSTEM_FINGERPRINT,
20842 expected,
20843 "baked fingerprint disagrees with a re-derivation over {} files: the id \
20844 is not a pure function of the source tree, or the build script did not \
20845 re-run after an edit",
20846 scan.files.len()
20847 );
20848 assert!(scan.files.len() > 100, "suspiciously small hashed file set");
20849 }
20850 None => {
20851 assert_eq!(BUILD_ID_SRC, build_id::BUILD_ID_SRC_DEGRADED);
20855 assert!(
20856 !BUILD_ID_NOTE.is_empty(),
20857 "a degraded build must state its reason so the boot WARN can print it"
20858 );
20859 }
20860 }
20861 }
20862
20863 #[test]
20866 fn identity_is_independent_of_git_history() {
20867 let id = SYSTEM_FINGERPRINT.rsplit_once('-').unwrap().1;
20868 assert_ne!(
20869 id, BUILD_GIT_SHA,
20870 "the content id equals the git sha; the identity must not be history, it has to \
20871 survive a rewrite that changes every commit"
20872 );
20873 assert!(
20874 !SYSTEM_FINGERPRINT.contains(BUILD_GIT_SHA),
20875 "the git sha leaked into the customer-visible fingerprint {SYSTEM_FINGERPRINT:?}"
20876 );
20877 assert!(!BUILD_GIT_SHA.is_empty());
20880 }
20881
20882 #[test]
20885 fn content_digest_is_deterministic_and_change_sensitive() {
20886 let a = build_id::degraded_build_id("memra-server", "0.123.0");
20887 let b = build_id::degraded_build_id("memra-server", "0.123.0");
20888 assert_eq!(a, b, "the digest is not deterministic");
20889 assert_eq!(a.len(), build_id::BUILD_ID_HEX);
20890 assert!(
20891 a.chars()
20892 .all(|c| c.is_ascii_digit() || ('a'..='f').contains(&c))
20893 );
20894 assert_ne!(a, build_id::degraded_build_id("memra-server", "0.123.1"));
20895 assert_ne!(a, build_id::degraded_build_id("memra-serve", "r0.123.0"));
20896 assert_eq!(build_id::render_build_id(0).len(), build_id::BUILD_ID_HEX);
20898 assert_eq!(
20899 build_id::render_build_id(0),
20900 "0".repeat(build_id::BUILD_ID_HEX)
20901 );
20902 }
20903
20904 #[test]
20907 fn two_scans_of_one_tree_agree() {
20908 let Some(root) = build_id::workspace_root(env!("CARGO_MANIFEST_DIR")) else {
20909 assert_eq!(BUILD_ID_SRC, build_id::BUILD_ID_SRC_DEGRADED);
20910 return;
20911 };
20912 let first = build_id::content_id(&root).expect("first scan");
20913 let second = build_id::content_id(&root).expect("second scan");
20914 assert_eq!(first.id, second.id);
20915 assert_eq!(first.files.len(), second.files.len());
20916 }
20917}
20918
20919#[cfg(test)]
20925mod vision_placement_gate_tests {
20926 use super::vision_media_admissible;
20927
20928 #[test]
20929 fn a_media_part_is_admitted_only_when_the_placement_admits() {
20930 assert_eq!(vision_media_admissible(true, "image"), Ok(()));
20931 assert_eq!(vision_media_admissible(true, "video"), Ok(()));
20932 let err = vision_media_admissible(false, "image").unwrap_err();
20933 assert!(
20934 err.starts_with("image input is not enabled on this deployment"),
20935 "same named refusal the armed-off path gives, so clients see one contract: {err}"
20936 );
20937 assert!(
20938 err.contains("placement"),
20939 "the refusal names its cause: {err}"
20940 );
20941 let err = vision_media_admissible(false, "video").unwrap_err();
20942 assert!(
20943 err.starts_with("video input is not enabled on this deployment"),
20944 "{err}"
20945 );
20946 }
20947
20948 fn live_src() -> String {
20949 let src: String = include_str!("lib.rs")
20950 .lines()
20951 .map(|l| match l.find("//") {
20952 Some(i) => &l[..i],
20953 None => l,
20954 })
20955 .collect::<Vec<_>>()
20956 .join("\n");
20957 let end = src
20958 .find("\nmod vision_placement_gate_tests")
20959 .expect("this test module exists");
20960 src[..end].to_string()
20961 }
20962
20963 fn item_body<'a>(live: &'a str, head: &str) -> &'a str {
20965 let start = live
20966 .find(head)
20967 .unwrap_or_else(|| panic!("{head} not found — did it get renamed?"));
20968 let body = &live[start..];
20969 let end = body.find("\n}\n").expect("item body closes");
20970 &body[..end]
20971 }
20972
20973 fn head_of(s: &str, n: usize) -> &str {
20975 match s.char_indices().nth(n) {
20976 Some((i, _)) => &s[..i],
20977 None => s,
20978 }
20979 }
20980
20981 #[test]
20986 fn no_family_switch_reads_the_placement_decision() {
20987 let live = live_src();
20988 for switch in [
20989 "fn vision_enabled()",
20990 "fn gemma_vision_enabled()",
20991 "fn step_vision_enabled()",
20992 ] {
20993 let body = item_body(&live, switch);
20994 assert!(
20995 !body.contains("vision_placement_serving")
20996 && !body.contains("vision_placement_admits"),
20997 "{switch} routes text rendering; it must stay keyed on the operator knobs alone"
20998 );
20999 }
21000 let walker = item_body(&live, "fn content_to_text_vision(");
21001 assert!(
21002 walker.contains(
21003 "if step_vision_enabled() {\n return content_to_text_vision_step(v, step_images);"
21004 ),
21005 "the step walker dispatch is keyed on the armed switch alone"
21006 );
21007 }
21008
21009 #[test]
21012 fn every_media_accepting_arm_passes_the_placement_gate() {
21013 let live = live_src();
21014 let step = item_body(&live, "fn content_to_text_vision_step(");
21015 let arm = step
21016 .split("Some(\"image_url\") => {")
21017 .nth(1)
21018 .expect("the step walker has an image arm");
21019 assert!(
21020 head_of(arm, 120).contains("vision_placement_admits(\"image\")?;"),
21021 "the step image arm must pass the placement gate first: {}",
21022 head_of(arm, 120)
21023 );
21024 let walker = item_body(&live, "fn content_to_text_vision(");
21025 for (head, kind) in [
21026 (
21027 "Some(\"image_url\") if gemma_vision_enabled() => {",
21028 "image",
21029 ),
21030 ("Some(\"image_url\") => {", "image"),
21031 ("Some(\"video_url\") => {", "video"),
21032 ] {
21033 let arm = walker
21034 .split(head)
21035 .nth(1)
21036 .unwrap_or_else(|| panic!("{head} is not an arm of the walker"));
21037 let window = head_of(arm, 400);
21038 assert!(
21039 window.contains(&format!("vision_placement_admits(\"{kind}\")?;")),
21040 "{head} must pass the placement gate before planning anything: {window}"
21041 );
21042 }
21043 assert!(live.contains("GLM5_VISION_SERVING.load(std::sync::atomic::Ordering::Acquire)"));
21047 let gate = item_body(&live, "fn vision_placement_admits(");
21049 assert!(gate.contains("vision_media_admissible(vision_placement_serving(), kind)"));
21050 }
21051}