[rust-renderer] Standalone preprocessing (#36718)
Signed-off-by: Sage Ahrac <sagiahrak@gmail.com> Co-authored-by: Shangming Cai <csmthu@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com> Co-authored-by: Rain Jiang <96632942+rainj-me@users.noreply.github.com>
This commit is contained in:
co-authored by
Shangming Cai
Liangsheng Yin
Rain Jiang
parent
6880a47955
commit
7b1c2ed0a4
@@ -0,0 +1,37 @@
|
||||
//! HTTP chat completion adapter.
|
||||
|
||||
use super::{
|
||||
ChatCompletionRequest,
|
||||
error::{json_rejection_response, response_error},
|
||||
response::sse_response,
|
||||
};
|
||||
use crate::openai::chat::serialize_chat_stream_response;
|
||||
use crate::openai::{OpenAIService, OperationResponse};
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{State, rejection::JsonRejection},
|
||||
response::{IntoResponse, Response},
|
||||
routing::post,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(super) fn routes() -> Router<Arc<OpenAIService>> {
|
||||
Router::new().route("/v1/chat/completions", post(chat_completions))
|
||||
}
|
||||
|
||||
async fn chat_completions(
|
||||
State(state): State<Arc<OpenAIService>>,
|
||||
body: Result<Json<ChatCompletionRequest>, JsonRejection>,
|
||||
) -> Response {
|
||||
let request = match body {
|
||||
Ok(Json(request)) => request,
|
||||
Err(error) => return json_rejection_response(error),
|
||||
};
|
||||
match state.chat(request).await {
|
||||
Ok(OperationResponse::Unary(response)) => Json(response).into_response(),
|
||||
Ok(OperationResponse::Stream(chunks)) => {
|
||||
sse_response(chunks, serialize_chat_stream_response)
|
||||
}
|
||||
Err(error) => response_error(error),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
//! HTTP completion adapter.
|
||||
|
||||
use super::{
|
||||
CompletionRequest,
|
||||
error::{json_rejection_response, response_error},
|
||||
response::sse_response,
|
||||
};
|
||||
use crate::openai::{OpenAIService, OperationResponse};
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{State, rejection::JsonRejection},
|
||||
response::{IntoResponse, Response},
|
||||
routing::post,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(super) fn routes() -> Router<Arc<OpenAIService>> {
|
||||
Router::new().route("/v1/completions", post(completions))
|
||||
}
|
||||
|
||||
async fn completions(
|
||||
State(state): State<Arc<OpenAIService>>,
|
||||
body: Result<Json<CompletionRequest>, JsonRejection>,
|
||||
) -> Response {
|
||||
let request = match body {
|
||||
Ok(Json(request)) => request,
|
||||
Err(error) => return json_rejection_response(error),
|
||||
};
|
||||
match state.complete(request).await {
|
||||
Ok(OperationResponse::Unary(response)) => Json(response).into_response(),
|
||||
Ok(OperationResponse::Stream(chunks)) => sse_response(chunks, |chunk| {
|
||||
serde_json::to_string(&chunk).expect("OpenAI response must serialize")
|
||||
}),
|
||||
Err(error) => response_error(error),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
use axum::{
|
||||
Json,
|
||||
extract::rejection::JsonRejection,
|
||||
http::StatusCode,
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
|
||||
use crate::ResponseError;
|
||||
|
||||
fn openai_error(code: StatusCode, message: impl Into<String>) -> Response {
|
||||
(code, Json(error_payload(code, message))).into_response()
|
||||
}
|
||||
|
||||
pub(super) fn json_rejection_response(rejection: JsonRejection) -> Response {
|
||||
let status = if rejection.status() == StatusCode::PAYLOAD_TOO_LARGE {
|
||||
StatusCode::PAYLOAD_TOO_LARGE
|
||||
} else {
|
||||
StatusCode::BAD_REQUEST
|
||||
};
|
||||
openai_error(status, rejection.body_text())
|
||||
}
|
||||
|
||||
pub(super) fn response_error(error: ResponseError) -> Response {
|
||||
let status = response_status(&error);
|
||||
openai_error(status, error.message)
|
||||
}
|
||||
|
||||
pub(super) fn response_status(error: &ResponseError) -> StatusCode {
|
||||
use crate::{ResponseErrorKind, UpstreamErrorCode};
|
||||
match error.kind {
|
||||
ResponseErrorKind::InvalidRequest => StatusCode::BAD_REQUEST,
|
||||
ResponseErrorKind::Unavailable => StatusCode::SERVICE_UNAVAILABLE,
|
||||
ResponseErrorKind::Internal => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
ResponseErrorKind::Upstream(UpstreamErrorCode::Http(code)) => {
|
||||
StatusCode::from_u16(code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn error_payload(status: StatusCode, message: impl Into<String>) -> serde_json::Value {
|
||||
let error_type = if status.is_server_error() {
|
||||
"InternalServerError"
|
||||
} else {
|
||||
"BadRequestError"
|
||||
};
|
||||
crate::openai::error_payload(status.as_u16(), message, error_type)
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
//! OpenAI HTTP frontend and render-only routes.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::Router;
|
||||
|
||||
use crate::engine::HttpGenerateClient;
|
||||
use crate::openai::OpenAIService;
|
||||
|
||||
mod chat;
|
||||
mod completions;
|
||||
mod error;
|
||||
mod proxy;
|
||||
mod render;
|
||||
mod response;
|
||||
mod tokenize;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
use crate::openai::protocol::{ChatCompletionRequest, CompletionRequest};
|
||||
|
||||
const DEFAULT_REQUEST_BODY_LIMIT_BYTES: usize = 32 * 1024 * 1024;
|
||||
|
||||
pub(crate) fn inference_routes(frontend: OpenAIService) -> Router<()> {
|
||||
Router::new()
|
||||
.merge(chat::routes())
|
||||
.merge(completions::routes())
|
||||
.with_state(Arc::new(frontend))
|
||||
}
|
||||
|
||||
fn renderer_routes(renderer: Arc<crate::RendererService>) -> Router<()> {
|
||||
render::routes(renderer.clone()).merge(tokenize::routes(renderer))
|
||||
}
|
||||
|
||||
fn with_request_body_limit(routes: Router<()>) -> Router<()> {
|
||||
// Limit JSON extraction without buffering or limiting raw proxy bodies.
|
||||
routes.layer(axum::extract::DefaultBodyLimit::max(
|
||||
DEFAULT_REQUEST_BODY_LIMIT_BYTES,
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn standalone_routes(
|
||||
frontend: OpenAIService,
|
||||
health_client: HttpGenerateClient,
|
||||
) -> Router<()> {
|
||||
let renderer = frontend.renderer.clone();
|
||||
let routes = inference_routes(frontend).merge(renderer_routes(renderer));
|
||||
let routes = routes.merge(render::engine_health_route(health_client));
|
||||
with_request_body_limit(routes)
|
||||
}
|
||||
|
||||
pub(crate) fn render_only_routes(renderer: Arc<crate::RendererService>) -> Router<()> {
|
||||
let routes = renderer_routes(renderer).merge(render::health_route());
|
||||
with_request_body_limit(routes)
|
||||
}
|
||||
|
||||
pub(crate) fn hosted_routes(
|
||||
frontend: OpenAIService,
|
||||
upstream_url: String,
|
||||
) -> Result<Router<()>, String> {
|
||||
let renderer = frontend.renderer.clone();
|
||||
let proxy = proxy::RustServerProxy::new(upstream_url)?;
|
||||
let routes = inference_routes(frontend)
|
||||
.merge(renderer_routes(renderer))
|
||||
.merge(render::readiness_route())
|
||||
.fallback(move |request| {
|
||||
let proxy = proxy.clone();
|
||||
async move { proxy.forward(request).await }
|
||||
});
|
||||
Ok(with_request_body_limit(routes))
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
//! Streaming HTTP fallback to the native Rust server.
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::http::{HeaderMap, HeaderName, Request, Response, StatusCode, header};
|
||||
use axum::response::IntoResponse;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) struct RustServerProxy {
|
||||
client: reqwest::Client,
|
||||
upstream_url: String,
|
||||
}
|
||||
|
||||
impl RustServerProxy {
|
||||
pub(super) fn new(upstream_url: String) -> Result<Self, String> {
|
||||
let upstream_url = upstream_url.trim_end_matches('/').to_owned();
|
||||
reqwest::Url::parse(&upstream_url)
|
||||
.map_err(|error| format!("invalid proxy upstream {upstream_url:?}: {error}"))?;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|error| format!("building Rust-server proxy client failed: {error}"))?;
|
||||
Ok(Self {
|
||||
client,
|
||||
upstream_url,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn forward(&self, request: Request<Body>) -> Response<Body> {
|
||||
let (mut parts, body) = request.into_parts();
|
||||
strip_hop_by_hop_headers(&mut parts.headers);
|
||||
// Let the client set Host for the upstream origin.
|
||||
parts.headers.remove(header::HOST);
|
||||
let path = parts
|
||||
.uri
|
||||
.path_and_query()
|
||||
.map_or("/", axum::http::uri::PathAndQuery::as_str);
|
||||
let upstream = format!("{}{path}", self.upstream_url);
|
||||
let response = self
|
||||
.client
|
||||
.request(parts.method, upstream)
|
||||
.headers(parts.headers)
|
||||
.body(reqwest::Body::wrap_stream(body.into_data_stream()))
|
||||
.send()
|
||||
.await;
|
||||
let response = match response {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
tracing::error!(%error, "Rust-server proxy request failed");
|
||||
return (StatusCode::BAD_GATEWAY, "Rust server unavailable").into_response();
|
||||
}
|
||||
};
|
||||
|
||||
let status = response.status();
|
||||
let mut headers = response.headers().clone();
|
||||
strip_hop_by_hop_headers(&mut headers);
|
||||
let mut builder = Response::builder().status(status);
|
||||
*builder
|
||||
.headers_mut()
|
||||
.expect("response builder must expose headers") = headers;
|
||||
builder
|
||||
.body(Body::from_stream(response.bytes_stream()))
|
||||
.unwrap_or_else(|error| {
|
||||
tracing::error!(%error, "building Rust-server proxy response failed");
|
||||
(StatusCode::BAD_GATEWAY, "Invalid Rust server response").into_response()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_hop_by_hop_headers(headers: &mut HeaderMap) {
|
||||
let connection_headers = headers
|
||||
.get(header::CONNECTION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.into_iter()
|
||||
.flat_map(|value| value.split(','))
|
||||
.filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok())
|
||||
.collect::<Vec<_>>();
|
||||
for name in connection_headers {
|
||||
headers.remove(name);
|
||||
}
|
||||
for name in [
|
||||
header::CONNECTION,
|
||||
header::HeaderName::from_static("keep-alive"),
|
||||
header::PROXY_AUTHENTICATE,
|
||||
header::PROXY_AUTHORIZATION,
|
||||
header::TE,
|
||||
header::TRAILER,
|
||||
header::TRANSFER_ENCODING,
|
||||
header::UPGRADE,
|
||||
] {
|
||||
headers.remove(name);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
//! Render-only HTTP routes and renderer health endpoints.
|
||||
|
||||
use super::{
|
||||
ChatCompletionRequest, CompletionRequest,
|
||||
error::{json_rejection_response, response_error},
|
||||
};
|
||||
use crate::{RendererService, engine::HttpGenerateClient};
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{State, rejection::JsonRejection},
|
||||
http::StatusCode,
|
||||
response::{IntoResponse, Response},
|
||||
routing::{get, post},
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(super) fn routes(renderer: Arc<RendererService>) -> Router<()> {
|
||||
Router::new()
|
||||
.route("/v1/chat/completions/render", post(render_chat))
|
||||
.route("/v1/completions/render", post(render_completions))
|
||||
.with_state(renderer)
|
||||
}
|
||||
|
||||
pub(super) fn health_route() -> Router<()> {
|
||||
Router::new().route("/health", get(health))
|
||||
}
|
||||
|
||||
pub(super) fn engine_health_route(generate_client: HttpGenerateClient) -> Router<()> {
|
||||
Router::new()
|
||||
.route("/health", get(engine_health))
|
||||
.with_state(generate_client)
|
||||
}
|
||||
|
||||
pub(super) fn readiness_route() -> Router<()> {
|
||||
Router::new().route("/_sglang_renderer/ready", get(readiness))
|
||||
}
|
||||
|
||||
async fn health() -> StatusCode {
|
||||
StatusCode::OK
|
||||
}
|
||||
|
||||
async fn engine_health(State(generate_client): State<HttpGenerateClient>) -> StatusCode {
|
||||
match generate_client.health_status().await {
|
||||
Ok(status) => status,
|
||||
Err(error) => {
|
||||
tracing::warn!(message = %error.message, "engine health check failed");
|
||||
StatusCode::SERVICE_UNAVAILABLE
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn readiness() -> impl IntoResponse {
|
||||
(StatusCode::NO_CONTENT, [("x-sglang-renderer", "ready")])
|
||||
}
|
||||
|
||||
async fn render_chat(
|
||||
State(renderer): State<Arc<RendererService>>,
|
||||
body: Result<Json<ChatCompletionRequest>, JsonRejection>,
|
||||
) -> Response {
|
||||
let request = match body {
|
||||
Ok(Json(request)) => request,
|
||||
Err(error) => return json_rejection_response(error),
|
||||
};
|
||||
match crate::openai::render::render_chat(&renderer, request).await {
|
||||
Ok(request) => Json(request).into_response(),
|
||||
Err(error) => response_error(error),
|
||||
}
|
||||
}
|
||||
|
||||
async fn render_completions(
|
||||
State(renderer): State<Arc<RendererService>>,
|
||||
body: Result<Json<CompletionRequest>, JsonRejection>,
|
||||
) -> Response {
|
||||
let request = match body {
|
||||
Ok(Json(request)) => request,
|
||||
Err(error) => return json_rejection_response(error),
|
||||
};
|
||||
match crate::openai::render::render_completions(&renderer, request).await {
|
||||
Ok(requests) => Json(requests).into_response(),
|
||||
Err(error) => response_error(error),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use axum::{
|
||||
body::{Body, to_bytes},
|
||||
http::Request,
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use crate::{RendererConfig, RendererError, RendererLimits, SamplingDefaults, TextTokenizer};
|
||||
|
||||
struct WordTokenizer;
|
||||
|
||||
impl TextTokenizer for WordTokenizer {
|
||||
fn encode(&self, text: &str, _add_special_tokens: bool) -> Result<Vec<i32>, RendererError> {
|
||||
Ok(text.split_whitespace().map(|_| 7).collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn app() -> Router<()> {
|
||||
let config = RendererConfig {
|
||||
served_model_name: "model".into(),
|
||||
tokenizer_path: ".".into(),
|
||||
revision: None,
|
||||
model_path: String::new(),
|
||||
chat_template: Some("chatml".into()),
|
||||
tool_call_parser: None,
|
||||
reasoning_parser: None,
|
||||
default_chat_template_kwargs: Default::default(),
|
||||
stream_response_default_include_usage: false,
|
||||
default_sampling_params: SamplingDefaults::default(),
|
||||
limits: RendererLimits {
|
||||
vocab_size: 100,
|
||||
context_len: 64,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
};
|
||||
routes(Arc::new(RendererService::with_tokenizer(
|
||||
config,
|
||||
Arc::new(WordTokenizer),
|
||||
2,
|
||||
2,
|
||||
)))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completion_render_returns_token_only_generate_requests() {
|
||||
let response = app()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/completions/render")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one two", "three"],
|
||||
"n": 2,
|
||||
"max_tokens": 5,
|
||||
"top_k": 17,
|
||||
"min_p": 0.2,
|
||||
"min_tokens": 3,
|
||||
"stop_regex": "END[0-9]",
|
||||
"rid": "request-id",
|
||||
"cache_salt": "tenant-a",
|
||||
"extra_key": "interactive",
|
||||
"priority": 7,
|
||||
"bootstrap_host": "prefill",
|
||||
"bootstrap_port": 8998,
|
||||
"bootstrap_room": 42,
|
||||
"routed_dp_rank": 2,
|
||||
"disagg_prefill_dp_rank": 1
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body: serde_json::Value =
|
||||
serde_json::from_slice(&to_bytes(response.into_body(), 64 * 1024).await.unwrap())
|
||||
.unwrap();
|
||||
assert_eq!(body[0]["input_ids"], serde_json::json!([7, 7]));
|
||||
assert_eq!(body[1]["input_ids"], serde_json::json!([7, 7]));
|
||||
assert_eq!(body[2]["input_ids"], serde_json::json!([7]));
|
||||
assert_eq!(body[3]["input_ids"], serde_json::json!([7]));
|
||||
assert!(
|
||||
body.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.all(|request| request.get("text").is_none())
|
||||
);
|
||||
assert_eq!(body[0]["sampling_params"]["top_k"], 17);
|
||||
assert_eq!(body[0]["sampling_params"]["min_p"], 0.2);
|
||||
assert_eq!(body[0]["sampling_params"]["min_new_tokens"], 3);
|
||||
assert_eq!(
|
||||
body[0]["sampling_params"]["stop_regex"],
|
||||
serde_json::json!(["END[0-9]"])
|
||||
);
|
||||
assert_eq!(body[0]["rid"], "request-id-0");
|
||||
assert_eq!(body[0]["model"], "model");
|
||||
assert_eq!(body[0]["cache_salt"], "tenant-a");
|
||||
assert_eq!(body[0]["extra_key"], "interactive");
|
||||
assert_eq!(body[0]["priority"], 7);
|
||||
assert_eq!(body[0]["bootstrap_host"], "prefill");
|
||||
assert_eq!(body[0]["bootstrap_port"], 8998);
|
||||
assert_eq!(body[0]["bootstrap_room"], 42);
|
||||
assert_eq!(body[0]["routed_dp_rank"], 2);
|
||||
assert_eq!(body[0]["disagg_prefill_dp_rank"], 1);
|
||||
assert_eq!(body[1]["rid"], "request-id-1");
|
||||
assert_eq!(body[2]["rid"], "request-id-2");
|
||||
assert_eq!(body[3]["rid"], "request-id-3");
|
||||
assert_eq!(body[3]["cache_salt"], "tenant-a");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn chat_render_rejects_multiple_choices() {
|
||||
let response = app()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/chat/completions/render")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"n": 2
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn render_rejects_unimplemented_stateful_fields() {
|
||||
let response = app()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/completions/render")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello",
|
||||
"session_id": "session"
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
//! HTTP SSE framing for typed OpenAI response streams.
|
||||
|
||||
use super::error::error_payload;
|
||||
use crate::ResponseError;
|
||||
use axum::response::{
|
||||
IntoResponse, Response,
|
||||
sse::{Event, Sse},
|
||||
};
|
||||
use futures::{Stream, StreamExt};
|
||||
use std::convert::Infallible;
|
||||
|
||||
pub(super) fn sse_response<T, S, F>(chunks: S, serialize: F) -> Response
|
||||
where
|
||||
T: Send + 'static,
|
||||
S: Stream<Item = Result<T, ResponseError>> + Send + 'static,
|
||||
F: Fn(T) -> String + Send + 'static,
|
||||
{
|
||||
let events = async_stream::stream! {
|
||||
futures::pin_mut!(chunks);
|
||||
while let Some(chunk) = chunks.next().await {
|
||||
let data = match chunk {
|
||||
Ok(chunk) => serialize(chunk),
|
||||
Err(error) => {
|
||||
let status = super::error::response_status(&error);
|
||||
error_payload(status, error.message).to_string()
|
||||
}
|
||||
};
|
||||
yield Ok::<_, Infallible>(Event::default().data(data));
|
||||
}
|
||||
// An error may be followed by the protocol's final usage chunk.
|
||||
// Only this transport owns the SSE terminator.
|
||||
yield Ok::<_, Infallible>(Event::default().data("[DONE]"));
|
||||
};
|
||||
Sse::new(events).into_response()
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,171 @@
|
||||
//! HTTP tokenization adapter.
|
||||
|
||||
use super::error::{json_rejection_response, response_error};
|
||||
use crate::{
|
||||
RendererService,
|
||||
openai::tokenize::{TokenizeRequest, tokenize as tokenize_request},
|
||||
};
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{State, rejection::JsonRejection},
|
||||
response::Response,
|
||||
routing::post,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(super) fn routes(renderer: Arc<RendererService>) -> Router<()> {
|
||||
Router::new()
|
||||
.route("/tokenize", post(tokenize))
|
||||
.route("/v1/tokenize", post(tokenize))
|
||||
.with_state(renderer)
|
||||
}
|
||||
|
||||
async fn tokenize(
|
||||
State(renderer): State<Arc<RendererService>>,
|
||||
body: Result<Json<TokenizeRequest>, JsonRejection>,
|
||||
) -> Result<Json<Value>, Response> {
|
||||
let Json(request) = body.map_err(json_rejection_response)?;
|
||||
tokenize_request(&renderer, request)
|
||||
.await
|
||||
.map(Json)
|
||||
.map_err(response_error)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use axum::{
|
||||
body::{Body, to_bytes},
|
||||
http::{Request, StatusCode},
|
||||
};
|
||||
use serde_json::json;
|
||||
use tower::ServiceExt;
|
||||
|
||||
use crate::{RendererConfig, RendererError, RendererLimits, SamplingDefaults, TextTokenizer};
|
||||
|
||||
struct PrefixTokenizer;
|
||||
|
||||
impl TextTokenizer for PrefixTokenizer {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> Result<Vec<i32>, RendererError> {
|
||||
Ok(add_special_tokens
|
||||
.then_some(1)
|
||||
.into_iter()
|
||||
.chain(text.split_whitespace().map(|_| 7))
|
||||
.chain(add_special_tokens.then_some(2))
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn app() -> Router<()> {
|
||||
let config = RendererConfig {
|
||||
served_model_name: "model".into(),
|
||||
tokenizer_path: ".".into(),
|
||||
revision: None,
|
||||
model_path: String::new(),
|
||||
chat_template: Some("chatml".into()),
|
||||
tool_call_parser: None,
|
||||
reasoning_parser: None,
|
||||
default_chat_template_kwargs: Default::default(),
|
||||
stream_response_default_include_usage: false,
|
||||
default_sampling_params: SamplingDefaults::default(),
|
||||
limits: RendererLimits {
|
||||
vocab_size: 100,
|
||||
context_len: 64,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
};
|
||||
routes(Arc::new(RendererService::with_tokenizer(
|
||||
config,
|
||||
Arc::new(PrefixTokenizer),
|
||||
2,
|
||||
2,
|
||||
)))
|
||||
}
|
||||
|
||||
async fn post(body: Value) -> (StatusCode, Value) {
|
||||
let response = app()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/tokenize")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(body.to_string()))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let status = response.status();
|
||||
let body =
|
||||
serde_json::from_slice(&to_bytes(response.into_body(), 64 * 1024).await.unwrap())
|
||||
.unwrap();
|
||||
(status, body)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn prompt_tokenization_preserves_batch_shape_and_special_token_choice() {
|
||||
let (status, body) = post(json!({
|
||||
"prompt": ["one two", ""],
|
||||
"add_special_tokens": false
|
||||
}))
|
||||
.await;
|
||||
assert_eq!(status, StatusCode::OK);
|
||||
assert_eq!(body["tokens"], json!([[7, 7], []]));
|
||||
assert_eq!(body["count"], json!([2, 0]));
|
||||
|
||||
let (_, body) = post(json!({"prompt": "one"})).await;
|
||||
assert_eq!(body["tokens"], json!([1, 7, 2]));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn chat_tokenization_applies_the_template_without_generation_limits() {
|
||||
let (status, body) = post(json!({
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_completion_tokens": 10_000
|
||||
}))
|
||||
.await;
|
||||
assert_eq!(status, StatusCode::OK);
|
||||
assert!(
|
||||
body["tokens"]
|
||||
.as_array()
|
||||
.is_some_and(|tokens| !tokens.is_empty())
|
||||
);
|
||||
assert_ne!(body["tokens"][0], json!(1));
|
||||
assert_ne!(
|
||||
body["tokens"][body["tokens"].as_array().unwrap().len() - 1],
|
||||
json!(2)
|
||||
);
|
||||
assert_eq!(
|
||||
body["count"],
|
||||
json!(body["tokens"].as_array().unwrap().len())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn chat_tokenization_continues_the_final_assistant_message() {
|
||||
let (_, regular) = post(json!({
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "partial answer"}
|
||||
]
|
||||
}))
|
||||
.await;
|
||||
let (status, continued) = post(json!({
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "partial answer"}
|
||||
],
|
||||
"continue_final_message": true,
|
||||
"chat_template_kwargs": {
|
||||
"continue_final_message": false,
|
||||
"add_generation_prompt": true
|
||||
}
|
||||
}))
|
||||
.await;
|
||||
|
||||
assert_eq!(status, StatusCode::OK);
|
||||
assert!(continued["count"].as_u64().unwrap() < regular["count"].as_u64().unwrap());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
//! Inbound protocol adapters.
|
||||
|
||||
#[cfg(feature = "http")]
|
||||
pub(crate) mod http;
|
||||
Reference in New Issue
Block a user