[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
Generated
+88
@@ -703,6 +703,16 @@ version = "0.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3d52eff69cd5e647efe296129160853a42795992097e8af39800e1060caeea9b"
|
||||
|
||||
[[package]]
|
||||
name = "core-foundation"
|
||||
version = "0.9.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f"
|
||||
dependencies = [
|
||||
"core-foundation-sys",
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "core-foundation-sys"
|
||||
version = "0.8.7"
|
||||
@@ -1076,6 +1086,15 @@ version = "1.0.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "34aa73646ffb006b8f5147f3dc182bd4bcb190227ce861fc4a4844bf8e3cb2c0"
|
||||
|
||||
[[package]]
|
||||
name = "encoding_rs"
|
||||
version = "0.8.35"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "equator"
|
||||
version = "0.4.2"
|
||||
@@ -1510,15 +1529,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97"
|
||||
dependencies = [
|
||||
"dirs",
|
||||
"futures",
|
||||
"http",
|
||||
"indicatif",
|
||||
"libc",
|
||||
"log",
|
||||
"num_cpus",
|
||||
"rand 0.9.5",
|
||||
"reqwest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror",
|
||||
"tokio",
|
||||
"ureq",
|
||||
"windows-sys 0.60.2",
|
||||
]
|
||||
@@ -1637,9 +1659,11 @@ dependencies = [
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"socket2 0.5.10",
|
||||
"system-configuration",
|
||||
"tokio",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
"windows-registry",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -3293,9 +3317,11 @@ checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"encoding_rs",
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
"futures-util",
|
||||
"h2",
|
||||
"http",
|
||||
"http-body",
|
||||
"http-body-util",
|
||||
@@ -3304,6 +3330,7 @@ dependencies = [
|
||||
"hyper-util",
|
||||
"js-sys",
|
||||
"log",
|
||||
"mime",
|
||||
"mime_guess",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
@@ -3696,6 +3723,35 @@ dependencies = [
|
||||
"ureq",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sglang-renderer"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"async-stream",
|
||||
"axum 0.8.9",
|
||||
"clap",
|
||||
"dynamo-parsers",
|
||||
"dynamo-protocols",
|
||||
"dynamo-renderer",
|
||||
"dynamo-tokenizers",
|
||||
"flume",
|
||||
"futures",
|
||||
"hf-hub",
|
||||
"minijinja",
|
||||
"regex-syntax",
|
||||
"reqwest",
|
||||
"rmp-serde",
|
||||
"rmpv",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror",
|
||||
"tokio",
|
||||
"tower 0.5.3",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sglang-server"
|
||||
version = "0.1.0"
|
||||
@@ -3959,6 +4015,27 @@ dependencies = [
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "system-configuration"
|
||||
version = "0.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"core-foundation",
|
||||
"system-configuration-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "system-configuration-sys"
|
||||
version = "0.6.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4"
|
||||
dependencies = [
|
||||
"core-foundation-sys",
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tagptr"
|
||||
version = "0.2.0"
|
||||
@@ -4752,6 +4829,17 @@ version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
|
||||
|
||||
[[package]]
|
||||
name = "windows-registry"
|
||||
version = "0.6.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720"
|
||||
dependencies = [
|
||||
"windows-link",
|
||||
"windows-result",
|
||||
"windows-strings",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-result"
|
||||
version = "0.4.1"
|
||||
|
||||
@@ -3,6 +3,7 @@ resolver = "3"
|
||||
members = [
|
||||
"sglang-grpc",
|
||||
"sglang-mm",
|
||||
"sglang-renderer",
|
||||
"sglang-server"
|
||||
]
|
||||
exclude = ["sglang-radix-tree"]
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
[package]
|
||||
name = "sglang-renderer"
|
||||
description = "Reusable SGLang request preprocessing with an optional OpenAI frontend"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[features]
|
||||
default = []
|
||||
http = [
|
||||
"dep:axum",
|
||||
"dep:clap",
|
||||
"dep:reqwest",
|
||||
"dep:tokio",
|
||||
"dep:tracing-subscriber",
|
||||
"hf-hub/rustls-tls",
|
||||
"hf-hub/tokio",
|
||||
]
|
||||
|
||||
[[bin]]
|
||||
name = "sglang-renderer"
|
||||
path = "src/main.rs"
|
||||
required-features = ["http"]
|
||||
|
||||
[dependencies]
|
||||
async-stream = { workspace = true }
|
||||
flume = "0.12.0"
|
||||
futures = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
tracing = { workspace = true }
|
||||
tokio = { workspace = true, optional = true }
|
||||
uuid = { workspace = true }
|
||||
|
||||
axum = { version = "0.8.9", features = ["json"], optional = true }
|
||||
clap = { version = "4", features = ["derive"], optional = true }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"], optional = true }
|
||||
tracing-subscriber = { workspace = true, optional = true }
|
||||
dynamo-parsers = "7.0.1"
|
||||
dynamo-protocols = "5.1.0"
|
||||
dynamo-renderer = "5.0.0"
|
||||
hf-hub = { version = "0.4", default-features = false }
|
||||
minijinja = { version = "2.24.0", features = ["unstable_machinery"] }
|
||||
regex-syntax = "=0.8.11"
|
||||
|
||||
# Keep this paired with the server until Dynamo exposes a smaller tokenizer API.
|
||||
dynamo-tokenizers = "1.7.0"
|
||||
|
||||
[dev-dependencies]
|
||||
tokio = { workspace = true }
|
||||
rmp-serde = "1"
|
||||
rmpv = { version = "1", features = ["with-serde"] }
|
||||
tower = { version = "0.5", features = ["util"] }
|
||||
@@ -0,0 +1,104 @@
|
||||
# SGLang renderer
|
||||
|
||||
The renderer runs as a separate service. It owns text preprocessing, token decoding,
|
||||
and OpenAI chat/completion responses. It submits token IDs through the native
|
||||
Rust server's existing `/generate` endpoint.
|
||||
|
||||
The renderer targets the existing `/generate` contract on SGLang main and must
|
||||
work with an unmodified Rust server. It accepts both cumulative and incremental
|
||||
streaming responses, using the engine's configured format. Additional generate
|
||||
request fields or server behavior changes are deferred to separate PRs.
|
||||
|
||||
## Build and run
|
||||
|
||||
From the repository root, build the standalone renderer. Rendering and
|
||||
tokenization work without an engine; generation requires a running SGLang engine.
|
||||
|
||||
```sh
|
||||
cargo build --manifest-path rust/Cargo.toml -p sglang-renderer --release --features http --locked
|
||||
```
|
||||
|
||||
Start the engine in one terminal.
|
||||
|
||||
```sh
|
||||
SGLANG_RUST_SERVER=1 python -m sglang.launch_server \
|
||||
--model-path meta-llama/Llama-3.1-8B-Instruct \
|
||||
--host 127.0.0.1 --port 30001 --skip-server-warmup
|
||||
```
|
||||
|
||||
Keep engine tokenization enabled for stop conditions and minimum-token handling.
|
||||
|
||||
Start the renderer in another terminal. Match the engine's model revision,
|
||||
tokenizer, context limit, and sampling defaults. Set tool and reasoning parsers
|
||||
on the renderer when needed.
|
||||
|
||||
```sh
|
||||
rust/target/release/sglang-renderer meta-llama/Llama-3.1-8B-Instruct \
|
||||
--engine-url http://127.0.0.1:30001 \
|
||||
--host 127.0.0.1 --port 30000 \
|
||||
--sampling-defaults openai --proxy-unhandled-routes
|
||||
```
|
||||
|
||||
Send OpenAI requests to port 30000. With `--proxy-unhandled-routes`, routes such as
|
||||
`/v1/models` and engine health checks are forwarded to the engine. The renderer's
|
||||
own `/_sglang_renderer/ready` endpoint returns HTTP 204 with
|
||||
`x-sglang-renderer: ready`; engine readiness is checked separately.
|
||||
|
||||
For preprocessing without an engine, omit `--engine-url`. This mode serves render
|
||||
and tokenization endpoints without inference.
|
||||
|
||||
```sh
|
||||
rust/target/release/sglang-renderer meta-llama/Llama-3.1-8B-Instruct \
|
||||
--host 127.0.0.1 --port 30000 --sampling-defaults openai
|
||||
```
|
||||
|
||||
The CLI defaults to sampling parameters from the model's generation config.
|
||||
`--sampling-defaults openai` matches SGLang's OpenAI API defaults. Use
|
||||
`--help` for template, parser, and limit options. A custom Cargo target directory
|
||||
or compilation target changes the executable path shown above.
|
||||
|
||||
## Tool-call parser support
|
||||
|
||||
`--tool-call-parser` uses Dynamo's parsers. See
|
||||
[Dynamo's supported tool-call parsers](https://docs.nvidia.com/dynamo/dev/parsing/tool-call-parsing#supported-tool-call-parsers)
|
||||
for parser names and model formats. These SGLang names need special attention:
|
||||
|
||||
| SGLang name | Renderer support |
|
||||
| --- | --- |
|
||||
| `llama3` | Accepted alias for `llama3_json` |
|
||||
| `qwen` | Accepted alias for `qwen25` |
|
||||
| `glm`, `glm45` | Accepted aliases for `glm47` |
|
||||
| `deepseekv3` | Use `deepseek_v3` |
|
||||
| `gpt-oss` | Use `harmony` |
|
||||
| `step3` | Unsupported |
|
||||
|
||||
Reasoning parsers are configured separately with `--reasoning-parser`.
|
||||
|
||||
## Docker image
|
||||
|
||||
Build the CPU-only renderer image from the repository root (`linux/amd64` or
|
||||
`linux/arm64`).
|
||||
|
||||
```sh
|
||||
docker buildx build --load -f docker/renderer.Dockerfile \
|
||||
-t local/sglang-renderer:dev .
|
||||
```
|
||||
|
||||
Run preprocessing without an engine.
|
||||
|
||||
```sh
|
||||
docker run --rm -p 30000:30000 \
|
||||
-v renderer-cache:/home/sglang/.cache/huggingface \
|
||||
-e HF_TOKEN \
|
||||
local/sglang-renderer:dev meta-llama/Llama-3.1-8B-Instruct \
|
||||
--host 0.0.0.0 --sampling-defaults openai
|
||||
```
|
||||
|
||||
For inference, add `--engine-url` with a URL reachable from the container.
|
||||
|
||||
## Current scope
|
||||
|
||||
OpenAI serving supports text chat and completions. Multimodal OpenAI inputs,
|
||||
`/responses`, and `/messages` are deferred. Automatic engine launch and packaged
|
||||
renderer installation are also deferred; manage both processes explicitly.
|
||||
The renderer does not implement API-key authentication or TLS.
|
||||
@@ -0,0 +1,37 @@
|
||||
//! Immutable configuration required during request rendering.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct SamplingDefaults {
|
||||
pub temperature: Option<f64>,
|
||||
pub top_p: Option<f64>,
|
||||
pub top_k: Option<i64>,
|
||||
pub min_p: Option<f64>,
|
||||
pub repetition_penalty: Option<f64>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct RendererLimits {
|
||||
pub vocab_size: u64,
|
||||
pub context_len: u64,
|
||||
pub num_reserved_tokens: u64,
|
||||
pub allow_auto_truncate: bool,
|
||||
pub enable_return_hidden_states: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct RendererConfig {
|
||||
pub served_model_name: String,
|
||||
pub tokenizer_path: String,
|
||||
pub revision: Option<String>,
|
||||
pub model_path: String,
|
||||
pub chat_template: Option<String>,
|
||||
pub tool_call_parser: Option<String>,
|
||||
pub reasoning_parser: Option<String>,
|
||||
#[serde(default)]
|
||||
pub default_chat_template_kwargs: std::collections::HashMap<String, serde_json::Value>,
|
||||
pub stream_response_default_include_usage: bool,
|
||||
pub default_sampling_params: SamplingDefaults,
|
||||
pub limits: RendererLimits,
|
||||
}
|
||||
@@ -0,0 +1,427 @@
|
||||
//! Prompt and generated-token decoding, including local text stops.
|
||||
|
||||
use super::{internal, invalid};
|
||||
use crate::{
|
||||
GenerateRequest, GenerationOutput, GenerationOutputExtras, ResponseError, TokenIds,
|
||||
TokenLogprob,
|
||||
};
|
||||
|
||||
use super::{GenerationFinishReason, GenerationStream, MatchedStop, TokenStream};
|
||||
use futures::StreamExt;
|
||||
|
||||
/// Shared tokenizer handle for prompt and generated-output decoding.
|
||||
pub(crate) struct TokenDecoder {
|
||||
tokenizer: dynamo_tokenizers::Tokenizer,
|
||||
}
|
||||
|
||||
pub(super) struct DecodeState {
|
||||
decoder: dynamo_tokenizers::DecodeStream,
|
||||
stops: Option<StopStringMatcher>,
|
||||
logprob_text: bool,
|
||||
}
|
||||
|
||||
impl TokenDecoder {
|
||||
pub(crate) fn new(tokenizer: dynamo_tokenizers::Tokenizer) -> Self {
|
||||
Self { tokenizer }
|
||||
}
|
||||
|
||||
pub(crate) fn detokenize_prompt(&self, token_ids: TokenIds) -> Result<String, ResponseError> {
|
||||
let ids = token_ids
|
||||
.into_iter()
|
||||
.map(u32::try_from)
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|_| invalid("token IDs must be non-negative"))?;
|
||||
self.tokenizer
|
||||
.decode(&ids, true)
|
||||
.map(String::from)
|
||||
.map_err(|error| invalid(format!("detokenizing prompt failed: {error}")))
|
||||
}
|
||||
|
||||
pub(super) fn prepare(
|
||||
&self,
|
||||
request: &mut GenerateRequest,
|
||||
) -> Result<DecodeState, ResponseError> {
|
||||
let stops = text_stop_matcher(request);
|
||||
let prompt_ids = request
|
||||
.input_ids
|
||||
.iter()
|
||||
.map(|&id| u32::try_from(id))
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|_| invalid("input_ids must be non-negative"))?;
|
||||
let logprob_text = request.return_text_in_logprobs.unwrap_or(false);
|
||||
request.return_text_in_logprobs = Some(false);
|
||||
Ok(DecodeState {
|
||||
decoder: self
|
||||
.tokenizer
|
||||
.decode_stream(&prompt_ids, request.sampling_params.skip_special_tokens),
|
||||
stops,
|
||||
logprob_text,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn decode(
|
||||
&self,
|
||||
mut tokens: TokenStream,
|
||||
mut state: DecodeState,
|
||||
) -> GenerationStream {
|
||||
let tokenizer = self.tokenizer.clone();
|
||||
async_stream::try_stream! {
|
||||
while let Some(delta) = tokens.next().await {
|
||||
let mut output = GenerationOutput::from(delta?);
|
||||
let matched = decode_output(&mut state.decoder, &mut output, state.stops.as_mut())?;
|
||||
if state.logprob_text {
|
||||
fill_logprob_text(&tokenizer, output.extras.as_deref_mut());
|
||||
}
|
||||
let stopped = matched.is_some();
|
||||
if let Some(stop) = matched {
|
||||
output.finish_reason = Some(GenerationFinishReason::Stop(Some(MatchedStop::Text(stop))));
|
||||
}
|
||||
if stopped {
|
||||
drop(tokens);
|
||||
yield output;
|
||||
return;
|
||||
}
|
||||
yield output;
|
||||
}
|
||||
}.boxed()
|
||||
}
|
||||
}
|
||||
|
||||
/// Match text stops locally without removing them from the engine request.
|
||||
///
|
||||
/// The engine uses the same stops to end decoding promptly. The renderer still
|
||||
/// needs its own matcher because it owns text decoding, stop trimming, and the
|
||||
/// OpenAI-facing finish reason.
|
||||
pub(super) fn text_stop_matcher(request: &GenerateRequest) -> Option<StopStringMatcher> {
|
||||
let params = &request.sampling_params;
|
||||
StopStringMatcher::new(params.stop.clone(), params.no_stop_trim)
|
||||
}
|
||||
|
||||
pub(super) struct StopStringMatcher {
|
||||
stops: Vec<String>,
|
||||
pending: String,
|
||||
include_stop: bool,
|
||||
}
|
||||
|
||||
struct StopMatch {
|
||||
text: String,
|
||||
matched: Option<String>,
|
||||
}
|
||||
|
||||
impl StopStringMatcher {
|
||||
fn new(stops: Vec<String>, include_stop: bool) -> Option<Self> {
|
||||
(!stops.is_empty()).then_some(Self {
|
||||
stops,
|
||||
pending: String::new(),
|
||||
include_stop,
|
||||
})
|
||||
}
|
||||
|
||||
fn push(&mut self, text: &str) -> StopMatch {
|
||||
self.pending.push_str(text);
|
||||
if let Some((position, stop)) = self
|
||||
.stops
|
||||
.iter()
|
||||
.filter_map(|stop| {
|
||||
self.pending
|
||||
.find(stop)
|
||||
.map(|position| (position, stop.clone()))
|
||||
})
|
||||
.min_by_key(|(position, _)| *position)
|
||||
{
|
||||
if stop.is_empty() {
|
||||
return StopMatch {
|
||||
text: std::mem::take(&mut self.pending),
|
||||
matched: Some(stop),
|
||||
};
|
||||
}
|
||||
let end = if self.include_stop {
|
||||
position + stop.len()
|
||||
} else {
|
||||
position
|
||||
};
|
||||
let text = self.pending[..end].to_owned();
|
||||
self.pending.clear();
|
||||
return StopMatch {
|
||||
text,
|
||||
matched: Some(stop),
|
||||
};
|
||||
}
|
||||
|
||||
let held_start = self
|
||||
.pending
|
||||
.char_indices()
|
||||
.map(|(start, _)| start)
|
||||
.chain(std::iter::once(self.pending.len()))
|
||||
.find(|&start| {
|
||||
self.stops
|
||||
.iter()
|
||||
.any(|stop| stop.starts_with(&self.pending[start..]))
|
||||
})
|
||||
.unwrap_or(self.pending.len());
|
||||
let held = self.pending.split_off(held_start);
|
||||
let text = std::mem::replace(&mut self.pending, held);
|
||||
StopMatch {
|
||||
text,
|
||||
matched: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> String {
|
||||
std::mem::take(&mut self.pending)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn decode_output(
|
||||
decoder: &mut dynamo_tokenizers::DecodeStream,
|
||||
output: &mut GenerationOutput,
|
||||
mut stop_matcher: Option<&mut StopStringMatcher>,
|
||||
) -> Result<Option<String>, ResponseError> {
|
||||
let mut text = String::new();
|
||||
for index in 0..output.token_ids.len() {
|
||||
let id = output.token_ids[index];
|
||||
let id = u32::try_from(id).map_err(|_| internal("engine returned a negative token ID"))?;
|
||||
let delta = decoder
|
||||
.step(id)
|
||||
.map_err(|error| internal(format!("detokenizing engine output failed: {error}")))?;
|
||||
if let Some(matcher) = stop_matcher.as_deref_mut() {
|
||||
let matched = matcher.push(delta.as_deref().unwrap_or_default());
|
||||
text.push_str(&matched.text);
|
||||
if let Some(stop) = matched.matched {
|
||||
truncate_output(output, index + 1)?;
|
||||
output.text = text;
|
||||
return Ok(Some(stop));
|
||||
}
|
||||
} else if let Some(delta) = delta {
|
||||
text.push_str(&delta);
|
||||
}
|
||||
}
|
||||
if output.finish_reason.is_some()
|
||||
&& let Some(matcher) = stop_matcher
|
||||
{
|
||||
text.push_str(&matcher.flush());
|
||||
}
|
||||
output.text = text;
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn truncate_output(output: &mut GenerationOutput, kept_tokens: usize) -> Result<(), ResponseError> {
|
||||
output.token_ids.truncate(kept_tokens);
|
||||
output.completion_tokens = u64::try_from(kept_tokens).unwrap_or(u64::MAX);
|
||||
let Some(extras) = output.extras.as_deref_mut() else {
|
||||
return Ok(());
|
||||
};
|
||||
truncate_optional(
|
||||
&mut extras.output_logprobs,
|
||||
kept_tokens,
|
||||
"output logprob positions",
|
||||
)
|
||||
}
|
||||
|
||||
fn truncate_optional<T>(
|
||||
values: &mut Vec<T>,
|
||||
length: usize,
|
||||
description: &str,
|
||||
) -> Result<(), ResponseError> {
|
||||
if values.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
if values.len() < length {
|
||||
return Err(internal(format!(
|
||||
"engine returned {} {description} values for {length} retained tokens",
|
||||
values.len()
|
||||
)));
|
||||
}
|
||||
values.truncate(length);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn fill_logprob_text(
|
||||
tokenizer: &dynamo_tokenizers::Tokenizer,
|
||||
extras: Option<&mut GenerationOutputExtras>,
|
||||
) {
|
||||
let Some(extras) = extras else { return };
|
||||
for position in extras
|
||||
.output_logprobs
|
||||
.iter_mut()
|
||||
.chain(&mut extras.input_logprobs)
|
||||
{
|
||||
fill_text(tokenizer, &mut position.token);
|
||||
for token in &mut position.top {
|
||||
fill_text(tokenizer, token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn fill_text(tokenizer: &dynamo_tokenizers::Tokenizer, token: &mut TokenLogprob) {
|
||||
if token.text.is_some() {
|
||||
return;
|
||||
}
|
||||
token.text = Some(
|
||||
u32::try_from(token.token_id)
|
||||
.ok()
|
||||
.and_then(|id| tokenizer.decode(&[id], false).ok())
|
||||
.map(String::from)
|
||||
.unwrap_or_default(),
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::test_utils::{position, tiny_tokenizer};
|
||||
use super::*;
|
||||
use crate::{GenerationOptions, SamplingParams, TokenIdsRequest};
|
||||
|
||||
fn request(stop: Vec<&str>) -> GenerateRequest {
|
||||
TokenIdsRequest {
|
||||
rid: "r".into(),
|
||||
input_ids: vec![1],
|
||||
options: GenerationOptions {
|
||||
sampling_params: SamplingParams {
|
||||
stop_strs: stop.into_iter().map(str::to_owned).collect(),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
metadata: Default::default(),
|
||||
}
|
||||
.into()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_stops_reach_the_frontend_and_engine() {
|
||||
let mut request = request(vec!["<eos>"]);
|
||||
request.sampling_params.stop_token_ids = Some(vec![9]);
|
||||
let matcher = text_stop_matcher(&request);
|
||||
|
||||
assert!(matcher.is_some());
|
||||
assert_eq!(request.sampling_params.stop_token_ids, Some(vec![9]));
|
||||
assert_eq!(request.sampling_params.stop, ["<eos>"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn regex_stops_and_min_tokens_reach_the_engine() {
|
||||
let mut request = request(vec!["END"]);
|
||||
request.sampling_params.stop_regex = vec!["[0-9]{3}".into()];
|
||||
request.sampling_params.min_new_tokens = 4;
|
||||
|
||||
text_stop_matcher(&request);
|
||||
|
||||
assert_eq!(request.sampling_params.stop, ["END"]);
|
||||
assert_eq!(request.sampling_params.stop_regex, ["[0-9]{3}"]);
|
||||
assert_eq!(request.sampling_params.min_new_tokens, 4);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decoded_stop_matcher_handles_cross_frame_matches_and_order() {
|
||||
let mut matcher = StopStringMatcher::new(vec!["END".into(), "ND".into()], false).unwrap();
|
||||
|
||||
let first = matcher.push("value E");
|
||||
assert_eq!(first.text, "value ");
|
||||
assert!(first.matched.is_none());
|
||||
|
||||
let second = matcher.push("ND trailing");
|
||||
assert_eq!(second.text, "");
|
||||
assert_eq!(second.matched.as_deref(), Some("END"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decoded_stop_matcher_uses_the_earliest_match() {
|
||||
let mut matcher =
|
||||
StopStringMatcher::new(vec!["later".into(), "first".into()], false).unwrap();
|
||||
|
||||
let matched = matcher.push("first then later");
|
||||
|
||||
assert_eq!(matched.text, "");
|
||||
assert_eq!(matched.matched.as_deref(), Some("first"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_stop_trim_includes_the_matched_text() {
|
||||
let mut matcher = StopStringMatcher::new(vec!["END".into()], true).unwrap();
|
||||
let matched = matcher.push("value END trailing");
|
||||
|
||||
assert_eq!(matched.text, "value END");
|
||||
assert_eq!(matched.matched.as_deref(), Some("END"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_stop_truncates_token_aligned_logprobs() {
|
||||
let mut output = GenerationOutput {
|
||||
token_ids: vec![7, 8, 9],
|
||||
completion_tokens: 3,
|
||||
extras: Some(Box::new(GenerationOutputExtras {
|
||||
output_logprobs: vec![
|
||||
position(7, -0.1, &[(7, -0.1), (6, -1.0)]),
|
||||
position(8, -0.2, &[(8, -0.2)]),
|
||||
position(9, -0.3, &[(9, -0.3)]),
|
||||
],
|
||||
..Default::default()
|
||||
})),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
truncate_output(&mut output, 2).unwrap();
|
||||
|
||||
assert_eq!(output.token_ids, [7, 8]);
|
||||
assert_eq!(output.completion_tokens, 2);
|
||||
let extras = output.extras.unwrap();
|
||||
assert_eq!(extras.output_logprobs.len(), 2);
|
||||
assert_eq!(extras.output_logprobs[0].top.len(), 2);
|
||||
assert_eq!(extras.output_logprobs[1].top.len(), 1);
|
||||
assert_eq!(extras.output_logprobs[1].token.token_id, 8);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_stops_are_matched_on_contextual_decoder_output() {
|
||||
let tokenizer = tiny_tokenizer();
|
||||
let token_ids = tokenizer
|
||||
.encode("hello")
|
||||
.unwrap()
|
||||
.token_ids()
|
||||
.iter()
|
||||
.map(|&id| id as i32)
|
||||
.collect::<Vec<_>>();
|
||||
let mut expected_decoder = tokenizer.decode_stream(&[65], true);
|
||||
let mut decoded = String::new();
|
||||
for &id in &token_ids {
|
||||
if let Some(delta) = expected_decoder.step(id as u32).unwrap() {
|
||||
decoded.push_str(&delta);
|
||||
}
|
||||
}
|
||||
assert!(!decoded.is_empty());
|
||||
|
||||
let mut decoder = tokenizer.decode_stream(&[65], true);
|
||||
let mut output = GenerationOutput {
|
||||
token_ids,
|
||||
completion_tokens: 1,
|
||||
..Default::default()
|
||||
};
|
||||
let mut matcher = StopStringMatcher::new(vec![decoded.clone()], false).unwrap();
|
||||
|
||||
let matched = decode_output(&mut decoder, &mut output, Some(&mut matcher)).unwrap();
|
||||
|
||||
assert_eq!(matched.as_deref(), Some(decoded.as_str()));
|
||||
assert!(output.text.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_stop_matches_after_the_first_generated_token() {
|
||||
let tokenizer = tiny_tokenizer();
|
||||
let mut decoder = tokenizer.decode_stream(&[65], true);
|
||||
let mut output = GenerationOutput {
|
||||
token_ids: vec![104, 101],
|
||||
completion_tokens: 2,
|
||||
..Default::default()
|
||||
};
|
||||
let mut matcher = StopStringMatcher::new(vec!["never".into(), String::new()], false)
|
||||
.expect("the empty stop must remain active");
|
||||
|
||||
let matched = decode_output(&mut decoder, &mut output, Some(&mut matcher)).unwrap();
|
||||
|
||||
assert_eq!(matched.as_deref(), Some(""));
|
||||
assert_eq!(output.token_ids, [104]);
|
||||
assert_eq!(output.completion_tokens, 1);
|
||||
assert_eq!(output.text, "h");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,550 @@
|
||||
//! HTTP client from renderer-owned generation requests to SGLang `/generate`.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use async_stream::stream;
|
||||
use futures::{StreamExt, future::BoxFuture};
|
||||
|
||||
use super::{GenerateTransport, TokenStream, internal};
|
||||
use crate::{GenerateRequest, ResponseError};
|
||||
use protocol::{engine_error_message, normalize_engine_output, parse_engine_frame};
|
||||
|
||||
mod protocol;
|
||||
|
||||
// SGLang's deep health probe defaults to 20 seconds. Leave it time to return
|
||||
// its own status while still bounding a peer that never sends response headers.
|
||||
const ENGINE_HEALTH_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
fn unavailable(message: impl Into<String>) -> ResponseError {
|
||||
ResponseError {
|
||||
kind: crate::ResponseErrorKind::Unavailable,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct HttpGenerateClient {
|
||||
client: reqwest::Client,
|
||||
generate_url: reqwest::Url,
|
||||
health_url: reqwest::Url,
|
||||
health_timeout: Duration,
|
||||
}
|
||||
|
||||
impl HttpGenerateClient {
|
||||
pub fn new(engine_url: impl AsRef<str>) -> Result<Self, String> {
|
||||
let engine_url = engine_url.as_ref();
|
||||
let base_url = reqwest::Url::parse(engine_url)
|
||||
.map_err(|error| format!("invalid engine URL {engine_url:?}: {error}"))?;
|
||||
let is_http_origin = matches!(base_url.scheme(), "http" | "https")
|
||||
&& base_url.host_str().is_some()
|
||||
&& base_url.username().is_empty()
|
||||
&& base_url.password().is_none()
|
||||
&& base_url.path() == "/"
|
||||
&& base_url.query().is_none()
|
||||
&& base_url.fragment().is_none();
|
||||
if !is_http_origin {
|
||||
return Err(format!(
|
||||
"invalid engine URL {engine_url:?}: expected an HTTP(S) origin without credentials, a path, query, or fragment"
|
||||
));
|
||||
}
|
||||
let generate_url = base_url
|
||||
.join("/generate")
|
||||
.map_err(|error| format!("joining /generate to engine URL failed: {error}"))?;
|
||||
let health_url = base_url
|
||||
.join("/health")
|
||||
.map_err(|error| format!("joining /health to engine URL failed: {error}"))?;
|
||||
let client = reqwest::Client::builder()
|
||||
.connect_timeout(Duration::from_secs(10))
|
||||
.build()
|
||||
.map_err(|error| format!("building engine HTTP client failed: {error}"))?;
|
||||
Ok(Self {
|
||||
client,
|
||||
generate_url,
|
||||
health_url,
|
||||
health_timeout: ENGINE_HEALTH_REQUEST_TIMEOUT,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn with_health_timeout(mut self, timeout: Duration) -> Self {
|
||||
self.health_timeout = timeout;
|
||||
self
|
||||
}
|
||||
|
||||
pub(crate) async fn health_status(&self) -> Result<reqwest::StatusCode, ResponseError> {
|
||||
let request = self
|
||||
.client
|
||||
.get(self.health_url.clone())
|
||||
.timeout(self.health_timeout);
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| unavailable(format!("engine health check failed: {error}")))?;
|
||||
Ok(response.status())
|
||||
}
|
||||
}
|
||||
|
||||
impl GenerateTransport for HttpGenerateClient {
|
||||
fn generate(
|
||||
&self,
|
||||
mut request: GenerateRequest,
|
||||
) -> BoxFuture<'_, Result<TokenStream, ResponseError>> {
|
||||
Box::pin(async move {
|
||||
// Always consume token deltas, including for unary frontend requests.
|
||||
request.stream = true;
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(self.generate_url.clone())
|
||||
.json(&request)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| unavailable(format!("engine request failed: {error}")))?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(
|
||||
status.as_u16(),
|
||||
)),
|
||||
message: engine_error_message(&body)
|
||||
.unwrap_or_else(|| format!("engine returned HTTP {status}")),
|
||||
});
|
||||
}
|
||||
|
||||
let mut chunks = response.bytes_stream();
|
||||
let events = stream! {
|
||||
let mut parser = SseParser::default();
|
||||
let mut terminal = false;
|
||||
let mut emitted_tokens = 0;
|
||||
while let Some(chunk) = chunks.next().await {
|
||||
let chunk = match chunk {
|
||||
Ok(chunk) => chunk,
|
||||
Err(error) => {
|
||||
yield Err(unavailable(format!("engine stream failed: {error}")));
|
||||
return;
|
||||
}
|
||||
};
|
||||
for payload in parser.push(&chunk) {
|
||||
if payload == "[DONE]" {
|
||||
if !terminal {
|
||||
yield Err(internal("engine stream ended before a terminal frame"));
|
||||
}
|
||||
return;
|
||||
}
|
||||
let mut output = match parse_engine_frame(&payload) {
|
||||
Ok(output) => output,
|
||||
Err(error) => {
|
||||
yield Err(error);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if let Err(error) = normalize_engine_output(&mut output, &mut emitted_tokens) {
|
||||
yield Err(error);
|
||||
return;
|
||||
}
|
||||
terminal = output.finish_reason.is_some();
|
||||
yield Ok(output);
|
||||
}
|
||||
}
|
||||
if !terminal {
|
||||
yield Err(internal("engine response closed before [DONE]"));
|
||||
}
|
||||
}
|
||||
.boxed();
|
||||
|
||||
Ok(events)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct SseParser {
|
||||
bytes: Vec<u8>,
|
||||
}
|
||||
|
||||
impl SseParser {
|
||||
fn push(&mut self, chunk: &[u8]) -> Vec<String> {
|
||||
self.bytes.extend_from_slice(chunk);
|
||||
let mut payloads = Vec::new();
|
||||
while let Some((end, separator_len)) = event_end(&self.bytes) {
|
||||
let event = self.bytes.drain(..end).collect::<Vec<_>>();
|
||||
self.bytes.drain(..separator_len);
|
||||
let event = String::from_utf8_lossy(&event);
|
||||
let data = event
|
||||
.lines()
|
||||
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
if !data.is_empty() {
|
||||
payloads.push(data);
|
||||
}
|
||||
}
|
||||
payloads
|
||||
}
|
||||
}
|
||||
|
||||
fn event_end(bytes: &[u8]) -> Option<(usize, usize)> {
|
||||
let crlf = bytes
|
||||
.windows(4)
|
||||
.position(|window| window == b"\r\n\r\n")
|
||||
.map(|position| (position, 4));
|
||||
let lf = bytes
|
||||
.windows(2)
|
||||
.position(|window| window == b"\n\n")
|
||||
.map(|position| (position, 2));
|
||||
match (crlf, lf) {
|
||||
(Some(crlf), Some(lf)) => Some(crlf.min(lf)),
|
||||
(Some(crlf), None) => Some(crlf),
|
||||
(None, Some(lf)) => Some(lf),
|
||||
(None, None) => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::engine::test_utils::tiny_tokenizer;
|
||||
use crate::{GenerationOptions, TokenIds, TokenIdsRequest};
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::State,
|
||||
response::sse::{Event, Sse},
|
||||
routing::post,
|
||||
};
|
||||
use std::convert::Infallible;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
#[test]
|
||||
fn sse_parser_handles_split_crlf_and_lf_frames() {
|
||||
let mut parser = SseParser::default();
|
||||
assert!(parser.push(b"data: {\"a\":1}\r\n").is_empty());
|
||||
assert_eq!(
|
||||
parser.push(b"\r\ndata: [DONE]\n\n"),
|
||||
["{\"a\":1}", "[DONE]"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_parser_uses_the_earliest_mixed_delimiter() {
|
||||
let mut parser = SseParser::default();
|
||||
|
||||
let payloads = parser.push(b"data: {\"a\":1}\n\ndata: {\"b\":2}\r\n\r\n");
|
||||
|
||||
assert_eq!(payloads, ["{\"a\":1}", "{\"b\":2}"]);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct EngineState {
|
||||
requests: Arc<Mutex<Vec<serde_json::Value>>>,
|
||||
output_ids: TokenIds,
|
||||
}
|
||||
|
||||
async fn generate(
|
||||
State(state): State<EngineState>,
|
||||
Json(body): Json<serde_json::Value>,
|
||||
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
|
||||
state.requests.lock().unwrap().push(body);
|
||||
let frame = serde_json::json!({
|
||||
"output_ids": state.output_ids,
|
||||
"meta_info": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": state.output_ids.len(),
|
||||
"finish_reason": {"type": "stop", "matched": null}
|
||||
}
|
||||
})
|
||||
.to_string();
|
||||
Sse::new(futures::stream::iter([
|
||||
Ok(Event::default().data(frame)),
|
||||
Ok(Event::default().data("[DONE]")),
|
||||
]))
|
||||
}
|
||||
|
||||
async fn streaming_generate(
|
||||
State(cumulative): State<bool>,
|
||||
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
|
||||
let frame = |completion_tokens, finish_reason: serde_json::Value| {
|
||||
Event::default().data(
|
||||
serde_json::json!({
|
||||
"output_ids": if cumulative { vec![104; completion_tokens] } else { vec![104] },
|
||||
"meta_info": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": completion_tokens,
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
)
|
||||
};
|
||||
Sse::new(futures::stream::iter([
|
||||
Ok(frame(1, serde_json::Value::Null)),
|
||||
Ok(frame(2, serde_json::json!({"type": "length", "length": 2}))),
|
||||
Ok(Event::default().data("[DONE]")),
|
||||
]))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn engine_origins_are_validated_and_joined_during_client_construction() {
|
||||
for invalid_url in [
|
||||
"127.0.0.1:30001",
|
||||
"ftp://engine.example",
|
||||
"http://user@engine.example",
|
||||
"http://engine.example/base",
|
||||
"http://engine.example?query",
|
||||
"http://engine.example#fragment",
|
||||
] {
|
||||
let error = match HttpGenerateClient::new(invalid_url) {
|
||||
Ok(_) => panic!("{invalid_url:?} must be rejected"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.contains("invalid engine URL"));
|
||||
}
|
||||
|
||||
let client = HttpGenerateClient::new("http://engine.example:30001/").unwrap();
|
||||
assert_eq!(
|
||||
client.generate_url.as_str(),
|
||||
"http://engine.example:30001/generate"
|
||||
);
|
||||
assert_eq!(
|
||||
client.health_url.as_str(),
|
||||
"http://engine.example:30001/health"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn backend_posts_token_ids_and_decodes_the_engine_stream() {
|
||||
let tokenizer = tiny_tokenizer();
|
||||
let output_ids = tokenizer
|
||||
.encode("hello")
|
||||
.unwrap()
|
||||
.token_ids()
|
||||
.iter()
|
||||
.map(|&id| id as i32)
|
||||
.collect::<Vec<_>>();
|
||||
let requests = Arc::new(Mutex::new(Vec::new()));
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(
|
||||
axum::serve(
|
||||
listener,
|
||||
Router::new()
|
||||
.route("/generate", post(generate))
|
||||
.with_state(EngineState {
|
||||
requests: requests.clone(),
|
||||
output_ids: output_ids.clone(),
|
||||
}),
|
||||
)
|
||||
.into_future(),
|
||||
);
|
||||
|
||||
let client = HttpGenerateClient::new(format!("http://{address}")).unwrap();
|
||||
let request = TokenIdsRequest {
|
||||
rid: "client-request".into(),
|
||||
input_ids: vec![65],
|
||||
options: GenerationOptions {
|
||||
return_text_in_logprobs: Some(true),
|
||||
..Default::default()
|
||||
},
|
||||
metadata: Default::default(),
|
||||
};
|
||||
let service = crate::engine::GenerationService::new(
|
||||
Arc::new(client),
|
||||
crate::engine::TokenDecoder::new(tokenizer.clone()),
|
||||
);
|
||||
let mut events = service.generate(request.into()).await.unwrap();
|
||||
let output = events.next().await.unwrap().unwrap();
|
||||
assert!(output.finish_reason.is_some());
|
||||
|
||||
let mut expected_decoder = tokenizer.decode_stream(&[65], true);
|
||||
let mut expected = String::new();
|
||||
for id in output_ids {
|
||||
if let Some(delta) = expected_decoder.step(id as u32).unwrap() {
|
||||
expected.push_str(&delta);
|
||||
}
|
||||
}
|
||||
assert_eq!(output.text, expected);
|
||||
assert_eq!(output.prompt_tokens, 1);
|
||||
let request = requests.lock().unwrap().pop().unwrap();
|
||||
assert_eq!(request["rid"], "client-request");
|
||||
assert_eq!(request["input_ids"], serde_json::json!([65]));
|
||||
assert_eq!(request["stream"], true);
|
||||
assert!(request.get("incremental_streaming_output").is_none());
|
||||
assert_eq!(request["return_text_in_logprobs"], false);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn transport_requires_a_terminal_frame_and_rejects_malformed_output() {
|
||||
async fn scripted(
|
||||
Json(request): Json<serde_json::Value>,
|
||||
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
|
||||
let case = request["rid"].as_str().unwrap();
|
||||
let terminal = case == "terminal-eof";
|
||||
let frame = serde_json::json!({
|
||||
"output_ids": [],
|
||||
"meta_info": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 0,
|
||||
"finish_reason": if terminal { serde_json::json!({"type": "length"}) } else { serde_json::Value::Null },
|
||||
}
|
||||
}).to_string();
|
||||
let frames = match case {
|
||||
"malformed" => vec!["{".to_owned()],
|
||||
"early-done" => vec![frame, "[DONE]".to_owned()],
|
||||
_ => vec![frame],
|
||||
};
|
||||
Sse::new(futures::stream::iter(
|
||||
frames
|
||||
.into_iter()
|
||||
.map(|frame| Ok(Event::default().data(frame))),
|
||||
))
|
||||
}
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(
|
||||
axum::serve(listener, Router::new().route("/generate", post(scripted))).into_future(),
|
||||
);
|
||||
let client = HttpGenerateClient::new(format!("http://{address}")).unwrap();
|
||||
for (case, error_message) in [
|
||||
("malformed", Some("invalid engine frame")),
|
||||
(
|
||||
"early-done",
|
||||
Some("engine stream ended before a terminal frame"),
|
||||
),
|
||||
(
|
||||
"unfinished-eof",
|
||||
Some("engine response closed before [DONE]"),
|
||||
),
|
||||
("terminal-eof", None),
|
||||
] {
|
||||
let request = TokenIdsRequest {
|
||||
rid: case.into(),
|
||||
input_ids: vec![65],
|
||||
options: GenerationOptions::default(),
|
||||
metadata: Default::default(),
|
||||
};
|
||||
let events = client
|
||||
.generate(request.into())
|
||||
.await
|
||||
.unwrap()
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
if let Some(message) = error_message {
|
||||
let error = events.last().unwrap().as_ref().unwrap_err();
|
||||
assert_eq!(error.kind, crate::ResponseErrorKind::Internal);
|
||||
assert!(
|
||||
error.message.starts_with(message),
|
||||
"{case}: {}",
|
||||
error.message
|
||||
);
|
||||
assert_eq!(events.iter().filter(|event| event.is_err()).count(), 1);
|
||||
} else {
|
||||
assert_eq!(events.len(), 1);
|
||||
assert!(events[0].as_ref().unwrap().finish_reason.is_some());
|
||||
}
|
||||
}
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn engine_frames_are_forwarded_once() {
|
||||
for cumulative in [false, true] {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(
|
||||
axum::serve(
|
||||
listener,
|
||||
Router::new()
|
||||
.route("/generate", post(streaming_generate))
|
||||
.with_state(cumulative),
|
||||
)
|
||||
.into_future(),
|
||||
);
|
||||
|
||||
let client = HttpGenerateClient::new(format!("http://{address}")).unwrap();
|
||||
let mut events = client
|
||||
.generate(
|
||||
TokenIdsRequest {
|
||||
rid: "incremental".into(),
|
||||
input_ids: vec![65],
|
||||
options: GenerationOptions::default(),
|
||||
metadata: Default::default(),
|
||||
}
|
||||
.into(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let first = events.next().await.unwrap().unwrap();
|
||||
assert!(first.finish_reason.is_none());
|
||||
assert_eq!(first.token_ids, [104]);
|
||||
assert_eq!(first.completion_tokens, 1);
|
||||
|
||||
let second = events.next().await.unwrap().unwrap();
|
||||
assert!(second.finish_reason.is_some());
|
||||
assert_eq!(second.token_ids, [104]);
|
||||
assert_eq!(second.completion_tokens, 1);
|
||||
assert!(events.next().await.is_none());
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
|
||||
struct DropNotice(Option<tokio::sync::oneshot::Sender<()>>);
|
||||
|
||||
impl Drop for DropNotice {
|
||||
fn drop(&mut self) {
|
||||
if let Some(sender) = self.0.take() {
|
||||
let _ = sender.send(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn slow_generate(
|
||||
State(notice): State<Arc<Mutex<Option<tokio::sync::oneshot::Sender<()>>>>>,
|
||||
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
|
||||
let guard = DropNotice(notice.lock().unwrap().take());
|
||||
Sse::new(stream! {
|
||||
let _guard = guard;
|
||||
yield Ok(Event::default().data(serde_json::json!({
|
||||
"output_ids": [104],
|
||||
"meta_info": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 1,
|
||||
"finish_reason": null
|
||||
}
|
||||
}).to_string()));
|
||||
futures::future::pending::<()>().await;
|
||||
})
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_renderer_events_closes_the_engine_stream() {
|
||||
let (notice_tx, notice_rx) = tokio::sync::oneshot::channel();
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(
|
||||
axum::serve(
|
||||
listener,
|
||||
Router::new()
|
||||
.route("/generate", post(slow_generate))
|
||||
.with_state(Arc::new(Mutex::new(Some(notice_tx)))),
|
||||
)
|
||||
.into_future(),
|
||||
);
|
||||
let client = HttpGenerateClient::new(format!("http://{address}")).unwrap();
|
||||
let request = TokenIdsRequest {
|
||||
rid: "cancel-me".into(),
|
||||
input_ids: vec![65],
|
||||
options: GenerationOptions::default(),
|
||||
metadata: Default::default(),
|
||||
};
|
||||
let mut events = client.generate(request.into()).await.unwrap();
|
||||
assert!(events.next().await.is_some());
|
||||
drop(events);
|
||||
tokio::time::timeout(Duration::from_secs(2), notice_rx)
|
||||
.await
|
||||
.expect("engine response stream was not dropped")
|
||||
.unwrap();
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,455 @@
|
||||
//! SGLang engine frame parsing and normalization into generation deltas.
|
||||
|
||||
use super::internal;
|
||||
use crate::engine::TokenDelta;
|
||||
use crate::{
|
||||
GenerationFinishReason, GenerationOutputExtras, MatchedStop, PositionLogprobs, ResponseError,
|
||||
TokenIds, TokenLogprob,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
|
||||
type WireLogprob = (Option<f32>, i32, Option<String>);
|
||||
type WireTopLogprobs = Vec<Option<Vec<WireLogprob>>>;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct EngineFrame {
|
||||
#[serde(default)]
|
||||
output_ids: TokenIds,
|
||||
meta_info: EngineMeta,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct EngineMeta {
|
||||
#[serde(default)]
|
||||
prompt_tokens: u32,
|
||||
#[serde(default)]
|
||||
completion_tokens: u64,
|
||||
#[serde(default)]
|
||||
finish_reason: Option<EngineFinishReason>,
|
||||
#[serde(default)]
|
||||
output_token_logprobs: Vec<WireLogprob>,
|
||||
#[serde(default)]
|
||||
input_token_logprobs: Vec<WireLogprob>,
|
||||
#[serde(default)]
|
||||
output_top_logprobs: WireTopLogprobs,
|
||||
#[serde(default)]
|
||||
input_top_logprobs: WireTopLogprobs,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct EngineFinishReason {
|
||||
#[serde(rename = "type")]
|
||||
kind: String,
|
||||
#[serde(default)]
|
||||
matched: Option<EngineMatchedStop>,
|
||||
#[serde(default)]
|
||||
status_code: Option<u16>,
|
||||
#[serde(default)]
|
||||
message: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum EngineMatchedStop {
|
||||
Token(i64),
|
||||
Text(String),
|
||||
Tokens(Vec<i64>),
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct EngineErrorEnvelope {
|
||||
error: EngineError,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct EngineError {
|
||||
#[serde(default = "default_error_code")]
|
||||
code: u16,
|
||||
message: String,
|
||||
}
|
||||
|
||||
fn default_error_code() -> u16 {
|
||||
500
|
||||
}
|
||||
|
||||
pub(super) fn parse_engine_frame(payload: &str) -> Result<TokenDelta, ResponseError> {
|
||||
if let Ok(error) = serde_json::from_str::<EngineErrorEnvelope>(payload) {
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(
|
||||
error.error.code,
|
||||
)),
|
||||
message: error.error.message,
|
||||
});
|
||||
}
|
||||
let frame: EngineFrame = serde_json::from_str(payload)
|
||||
.map_err(|error| internal(format!("invalid engine frame: {error}")))?;
|
||||
if let Some(reason) = frame.meta_info.finish_reason.as_ref()
|
||||
&& reason.kind == "abort"
|
||||
&& let Some(status_code) = reason.status_code
|
||||
{
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(status_code)),
|
||||
message: reason
|
||||
.message
|
||||
.clone()
|
||||
.unwrap_or_else(|| "request aborted".to_owned()),
|
||||
});
|
||||
}
|
||||
let finish_reason = frame
|
||||
.meta_info
|
||||
.finish_reason
|
||||
.map(|reason| match reason.kind.as_str() {
|
||||
"stop" => GenerationFinishReason::Stop(reason.matched.map(|matched| match matched {
|
||||
EngineMatchedStop::Token(id) => MatchedStop::Token(id),
|
||||
EngineMatchedStop::Text(text) => MatchedStop::Text(text),
|
||||
EngineMatchedStop::Tokens(ids) => MatchedStop::Tokens(ids),
|
||||
})),
|
||||
"length" => GenerationFinishReason::Length,
|
||||
"abort" => GenerationFinishReason::Abort,
|
||||
"content_filter" => GenerationFinishReason::ContentFilter,
|
||||
other => GenerationFinishReason::Other(other.to_owned()),
|
||||
});
|
||||
let has_extras = !frame.meta_info.output_token_logprobs.is_empty()
|
||||
|| !frame.meta_info.input_token_logprobs.is_empty()
|
||||
|| !frame.meta_info.output_top_logprobs.is_empty()
|
||||
|| !frame.meta_info.input_top_logprobs.is_empty();
|
||||
let output_logprobs = group_logprobs(
|
||||
frame.meta_info.output_token_logprobs,
|
||||
frame.meta_info.output_top_logprobs,
|
||||
"output",
|
||||
)?;
|
||||
let input_logprobs = group_logprobs(
|
||||
frame.meta_info.input_token_logprobs,
|
||||
frame.meta_info.input_top_logprobs,
|
||||
"input",
|
||||
)?;
|
||||
let extras = has_extras.then_some(Box::new(GenerationOutputExtras {
|
||||
output_logprobs,
|
||||
input_logprobs,
|
||||
}));
|
||||
Ok(TokenDelta {
|
||||
token_ids: frame.output_ids,
|
||||
finish_reason,
|
||||
prompt_tokens: frame.meta_info.prompt_tokens,
|
||||
completion_tokens: frame.meta_info.completion_tokens,
|
||||
extras,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn normalize_engine_output(
|
||||
output: &mut TokenDelta,
|
||||
emitted_tokens: &mut u64,
|
||||
) -> Result<(), ResponseError> {
|
||||
let total = output.completion_tokens;
|
||||
let delta = total.checked_sub(*emitted_tokens).ok_or_else(|| {
|
||||
internal(format!(
|
||||
"engine completion token count decreased from {} to {total}",
|
||||
*emitted_tokens
|
||||
))
|
||||
})?;
|
||||
let output_len = u64::try_from(output.token_ids.len()).unwrap_or(u64::MAX);
|
||||
let trimmed_stop_tokens = match output.finish_reason.as_ref() {
|
||||
Some(GenerationFinishReason::Stop(Some(MatchedStop::Token(_)))) => 1,
|
||||
Some(GenerationFinishReason::Stop(Some(MatchedStop::Tokens(ids)))) => {
|
||||
u64::try_from(ids.len()).unwrap_or(u64::MAX)
|
||||
}
|
||||
_ => 0,
|
||||
};
|
||||
let cumulative =
|
||||
output_len == total || output_len.checked_add(trimmed_stop_tokens) == Some(total);
|
||||
let incremental =
|
||||
output_len == delta || output_len.checked_add(trimmed_stop_tokens) == Some(delta);
|
||||
|
||||
if cumulative {
|
||||
let prefix = usize::try_from(*emitted_tokens)
|
||||
.map_err(|_| internal("engine completion token count exceeds addressable memory"))?;
|
||||
if prefix > output.token_ids.len() {
|
||||
return Err(internal(format!(
|
||||
"engine returned {output_len} cumulative output token IDs after {prefix} were already emitted"
|
||||
)));
|
||||
}
|
||||
output.token_ids.drain(..prefix);
|
||||
if let Some(extras) = output.extras.as_deref_mut() {
|
||||
trim_cumulative_output_extras(extras, prefix)?;
|
||||
}
|
||||
} else if !incremental {
|
||||
return Err(internal(format!(
|
||||
"engine returned {output_len} output token IDs after reporting {delta} new completion tokens"
|
||||
)));
|
||||
}
|
||||
|
||||
output.completion_tokens = delta;
|
||||
*emitted_tokens = total;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn trim_cumulative_output_extras(
|
||||
extras: &mut GenerationOutputExtras,
|
||||
prefix: usize,
|
||||
) -> Result<(), ResponseError> {
|
||||
drain_optional_prefix(
|
||||
&mut extras.output_logprobs,
|
||||
prefix,
|
||||
"output logprob positions",
|
||||
)
|
||||
}
|
||||
|
||||
fn drain_prefix<T>(
|
||||
values: &mut Vec<T>,
|
||||
prefix: usize,
|
||||
description: &str,
|
||||
) -> Result<(), ResponseError> {
|
||||
if values.len() < prefix {
|
||||
return Err(internal(format!(
|
||||
"engine returned {} {description} values for a {prefix}-token cumulative prefix",
|
||||
values.len()
|
||||
)));
|
||||
}
|
||||
values.drain(..prefix);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn drain_optional_prefix<T>(
|
||||
values: &mut Vec<T>,
|
||||
prefix: usize,
|
||||
description: &str,
|
||||
) -> Result<(), ResponseError> {
|
||||
if values.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
drain_prefix(values, prefix, description)
|
||||
}
|
||||
|
||||
fn wire_logprob((logprob, token_id, text): WireLogprob) -> TokenLogprob {
|
||||
TokenLogprob {
|
||||
logprob,
|
||||
token_id,
|
||||
text,
|
||||
}
|
||||
}
|
||||
|
||||
fn group_logprobs(
|
||||
values: Vec<WireLogprob>,
|
||||
top_values: WireTopLogprobs,
|
||||
kind: &str,
|
||||
) -> Result<Vec<PositionLogprobs>, ResponseError> {
|
||||
// P/D can send a single null position when top logprobs are disabled.
|
||||
if top_values.iter().all(Option::is_none) {
|
||||
return Ok(values
|
||||
.into_iter()
|
||||
.map(|token| PositionLogprobs {
|
||||
token: wire_logprob(token),
|
||||
top: Vec::new(),
|
||||
})
|
||||
.collect());
|
||||
}
|
||||
|
||||
if top_values.len() != values.len() {
|
||||
return Err(internal(format!(
|
||||
"engine returned {} {kind} top-logprob positions for {} selected-token positions",
|
||||
top_values.len(),
|
||||
values.len()
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(values
|
||||
.into_iter()
|
||||
.zip(top_values)
|
||||
.map(|(token, top)| PositionLogprobs {
|
||||
token: wire_logprob(token),
|
||||
top: top
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(wire_logprob)
|
||||
.collect(),
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
pub(super) fn engine_error_message(body: &str) -> Option<String> {
|
||||
serde_json::from_str::<EngineErrorEnvelope>(body)
|
||||
.ok()
|
||||
.map(|error| error.error.message)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::engine::test_utils::position;
|
||||
|
||||
#[test]
|
||||
fn engine_frame_maps_tokens_usage_finish_and_logprobs() {
|
||||
let output = parse_engine_frame(
|
||||
r#"{
|
||||
"output_ids":[7],
|
||||
"meta_info":{
|
||||
"prompt_tokens":3,
|
||||
"completion_tokens":1,
|
||||
"finish_reason":{"type":"stop","matched":9},
|
||||
"output_token_logprobs":[[-0.25,7,null]],
|
||||
"output_top_logprobs":[[[-0.25,7,null],[-1.0,8,null]]]
|
||||
}
|
||||
}"#,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(output.token_ids, [7]);
|
||||
assert_eq!(output.prompt_tokens, 3);
|
||||
assert_eq!(output.completion_tokens, 1);
|
||||
assert_eq!(
|
||||
output.finish_reason,
|
||||
Some(GenerationFinishReason::Stop(Some(MatchedStop::Token(9))))
|
||||
);
|
||||
let extras = output.extras.unwrap();
|
||||
assert_eq!(extras.output_logprobs.len(), 1);
|
||||
assert_eq!(extras.output_logprobs[0].token.token_id, 7);
|
||||
assert_eq!(extras.output_logprobs[0].top.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn engine_frame_preserves_selected_logprobs_with_absent_top_positions() {
|
||||
let output = parse_engine_frame(
|
||||
r#"{
|
||||
"output_ids":[12095,13],
|
||||
"meta_info":{
|
||||
"prompt_tokens":5,
|
||||
"completion_tokens":2,
|
||||
"output_token_logprobs":[
|
||||
[-0.42652416229248047,12095,null],
|
||||
[-0.7053262591362,13,null]
|
||||
],
|
||||
"output_top_logprobs":[null]
|
||||
}
|
||||
}"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(output.token_ids, [12095, 13]);
|
||||
assert_eq!(
|
||||
output.extras.unwrap().output_logprobs,
|
||||
[
|
||||
position(12095, -0.42652416, &[]),
|
||||
position(13, -0.70532626, &[])
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn engine_frame_rejects_misaligned_logprob_positions() {
|
||||
let error = parse_engine_frame(
|
||||
r#"{
|
||||
"output_ids":[7,8],
|
||||
"meta_info":{
|
||||
"completion_tokens":2,
|
||||
"output_token_logprobs":[[-0.25,7,null],[-0.5,8,null]],
|
||||
"output_top_logprobs":[[[-0.25,7,null]]]
|
||||
}
|
||||
}"#,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(error.kind, crate::ResponseErrorKind::Internal);
|
||||
assert_eq!(
|
||||
error.message,
|
||||
"engine returned 1 output top-logprob positions for 2 selected-token positions"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn engine_error_frame_preserves_status_and_message() {
|
||||
let error = parse_engine_frame(
|
||||
r#"{"error":{"message":"too long","type":"BadRequestError","code":400}}"#,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert_eq!(
|
||||
error.kind,
|
||||
crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(400))
|
||||
);
|
||||
assert_eq!(error.message, "too long");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn coded_abort_frame_preserves_status_and_message() {
|
||||
let error = parse_engine_frame(
|
||||
r#"{"output_ids":[],"meta_info":{"finish_reason":{"type":"abort","status_code":503,"message":"out of memory"}}}"#,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(
|
||||
error.kind,
|
||||
crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(503))
|
||||
);
|
||||
assert_eq!(error.message, "out of memory");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uncoded_abort_frame_remains_a_finish_reason() {
|
||||
let output = parse_engine_frame(
|
||||
r#"{"output_ids":[],"meta_info":{"finish_reason":{"type":"abort","status_code":null,"message":"cancelled"}}}"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(output.finish_reason, Some(GenerationFinishReason::Abort));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cumulative_engine_frames_become_deltas() {
|
||||
let mut emitted_tokens = 1;
|
||||
let mut output = TokenDelta {
|
||||
token_ids: vec![7, 8],
|
||||
completion_tokens: 2,
|
||||
extras: Some(Box::new(GenerationOutputExtras {
|
||||
output_logprobs: vec![
|
||||
position(7, -0.5, &[(7, -0.5), (9, -1.0)]),
|
||||
position(8, -0.25, &[(8, -0.25)]),
|
||||
],
|
||||
..Default::default()
|
||||
})),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
normalize_engine_output(&mut output, &mut emitted_tokens).unwrap();
|
||||
|
||||
assert_eq!(output.token_ids, [8]);
|
||||
assert_eq!(output.completion_tokens, 1);
|
||||
assert_eq!(emitted_tokens, 2);
|
||||
let extras = output.extras.unwrap();
|
||||
assert_eq!(extras.output_logprobs.len(), 1);
|
||||
assert_eq!(extras.output_logprobs[0].token.token_id, 8);
|
||||
assert_eq!(extras.output_logprobs[0].top.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_stops_may_be_trimmed_from_incremental_or_cumulative_frames() {
|
||||
for token_ids in [vec![], vec![7]] {
|
||||
let mut emitted_tokens = 1;
|
||||
let mut output = TokenDelta {
|
||||
token_ids,
|
||||
completion_tokens: 2,
|
||||
finish_reason: Some(GenerationFinishReason::Stop(Some(MatchedStop::Token(9)))),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
normalize_engine_output(&mut output, &mut emitted_tokens).unwrap();
|
||||
|
||||
assert!(output.token_ids.is_empty());
|
||||
assert_eq!(output.completion_tokens, 1);
|
||||
assert_eq!(emitted_tokens, 2);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inconsistent_engine_token_counts_are_rejected() {
|
||||
let mut emitted_tokens = 2;
|
||||
let mut output = TokenDelta {
|
||||
token_ids: vec![7, 8],
|
||||
completion_tokens: 3,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let error = normalize_engine_output(&mut output, &mut emitted_tokens).unwrap_err();
|
||||
|
||||
assert_eq!(error.kind, crate::ResponseErrorKind::Internal);
|
||||
assert!(error.message.contains("2 output token IDs"));
|
||||
assert_eq!(emitted_tokens, 2);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
//! Token-only generation transport and decoded engine output.
|
||||
|
||||
use futures::{StreamExt, TryStreamExt, future::BoxFuture};
|
||||
|
||||
use crate::{GenerateRequest, ResponseError};
|
||||
|
||||
mod decode;
|
||||
#[cfg(feature = "http")]
|
||||
mod http;
|
||||
pub(crate) mod response;
|
||||
mod types;
|
||||
|
||||
pub(crate) use decode::TokenDecoder;
|
||||
#[cfg(feature = "http")]
|
||||
pub(crate) use http::HttpGenerateClient;
|
||||
pub(crate) use types::{
|
||||
GenerationFinishReason, GenerationOutput, GenerationOutputExtras, GenerationStream,
|
||||
MatchedStop, PositionLogprobs, TokenDelta, TokenLogprob,
|
||||
};
|
||||
|
||||
pub(crate) type TokenStream =
|
||||
futures::stream::BoxStream<'static, Result<TokenDelta, ResponseError>>;
|
||||
|
||||
/// Backend generation from prepared token requests to normalized token deltas.
|
||||
///
|
||||
/// Successful streams carry a finish reason on their terminal output. The caller
|
||||
/// owns the submission future and response stream; dropping either must release the
|
||||
/// corresponding transport work. HTTP health checks and proxying are separate.
|
||||
pub(crate) trait GenerateTransport: Send + Sync {
|
||||
fn generate(
|
||||
&self,
|
||||
request: GenerateRequest,
|
||||
) -> BoxFuture<'_, Result<TokenStream, ResponseError>>;
|
||||
}
|
||||
|
||||
// Bound pending submissions per request without duplicating scheduler admission.
|
||||
const CONCURRENT_ENGINE_SUBMISSIONS: usize = 32;
|
||||
|
||||
/// Shared generation policy and decoding, independent of the engine transport.
|
||||
pub(crate) struct GenerationService {
|
||||
transport: std::sync::Arc<dyn GenerateTransport>,
|
||||
pub(crate) decoder: TokenDecoder,
|
||||
}
|
||||
|
||||
impl GenerationService {
|
||||
pub(crate) fn new(
|
||||
transport: std::sync::Arc<dyn GenerateTransport>,
|
||||
decoder: TokenDecoder,
|
||||
) -> Self {
|
||||
Self { transport, decoder }
|
||||
}
|
||||
|
||||
pub(crate) async fn generate(
|
||||
&self,
|
||||
mut request: GenerateRequest,
|
||||
) -> Result<GenerationStream, ResponseError> {
|
||||
let decode = self.decoder.prepare(&mut request)?;
|
||||
let tokens = self.transport.generate(request).await?;
|
||||
Ok(self.decoder.decode(tokens, decode))
|
||||
}
|
||||
|
||||
/// Establish all choice streams before consumption, retaining input order.
|
||||
pub(crate) async fn generate_many(
|
||||
&self,
|
||||
inputs: Vec<GenerateRequest>,
|
||||
) -> Result<Vec<GenerationStream>, ResponseError> {
|
||||
futures::stream::iter(inputs.into_iter().map(|input| self.generate(input)))
|
||||
.buffered(CONCURRENT_ENGINE_SUBMISSIONS)
|
||||
.try_collect()
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
fn invalid(message: impl Into<String>) -> ResponseError {
|
||||
ResponseError {
|
||||
kind: crate::ResponseErrorKind::InvalidRequest,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn internal(message: impl Into<String>) -> ResponseError {
|
||||
ResponseError {
|
||||
kind: crate::ResponseErrorKind::Internal,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) mod test_utils;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
@@ -0,0 +1,116 @@
|
||||
//! Generation stream merging and aggregation.
|
||||
|
||||
use crate::{GenerationOutput, GenerationStream, ResponseError};
|
||||
use futures::{StreamExt, stream::BoxStream};
|
||||
|
||||
pub(crate) fn merge_indexed(
|
||||
streams: Vec<GenerationStream>,
|
||||
) -> BoxStream<'static, (usize, Result<GenerationOutput, ResponseError>)> {
|
||||
let streams = streams
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(index, events)| events.map(move |event| (index, event)).boxed());
|
||||
futures::stream::select_all(streams).boxed()
|
||||
}
|
||||
|
||||
pub(crate) async fn collect_output(
|
||||
mut events: GenerationStream,
|
||||
) -> Result<GenerationOutput, ResponseError> {
|
||||
let mut collected = GenerationOutput::default();
|
||||
while let Some(item) = events.next().await {
|
||||
let output = item?;
|
||||
let finished = output.finish_reason.is_some();
|
||||
fold_output(&mut collected, output);
|
||||
if finished {
|
||||
return Ok(collected);
|
||||
}
|
||||
}
|
||||
Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Internal,
|
||||
message: "response truncated before completion".into(),
|
||||
})
|
||||
}
|
||||
|
||||
fn fold_output(collected: &mut GenerationOutput, output: GenerationOutput) {
|
||||
collected.text.push_str(&output.text);
|
||||
collected.token_ids.extend(output.token_ids);
|
||||
collected.prompt_tokens = output.prompt_tokens;
|
||||
collected.completion_tokens = collected
|
||||
.completion_tokens
|
||||
.saturating_add(output.completion_tokens);
|
||||
if output.finish_reason.is_some() {
|
||||
collected.finish_reason = output.finish_reason;
|
||||
}
|
||||
if let Some(output) = output.extras {
|
||||
let collected = collected
|
||||
.extras
|
||||
.get_or_insert_with(|| Box::new(crate::GenerationOutputExtras::default()));
|
||||
collected.output_logprobs.extend(output.output_logprobs);
|
||||
if !output.input_logprobs.is_empty() {
|
||||
collected.input_logprobs = output.input_logprobs;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use futures::{StreamExt, stream};
|
||||
|
||||
use super::super::test_utils::position;
|
||||
use super::{fold_output, merge_indexed};
|
||||
use crate::{GenerationOutput, GenerationOutputExtras};
|
||||
|
||||
#[test]
|
||||
fn unary_output_appends_generated_logprobs_and_replaces_prompt_logprobs() {
|
||||
let mut collected = GenerationOutput::default();
|
||||
for (output_token, input_token) in [(1, 10), (2, 20)] {
|
||||
fold_output(
|
||||
&mut collected,
|
||||
GenerationOutput {
|
||||
extras: Some(Box::new(GenerationOutputExtras {
|
||||
output_logprobs: vec![position(output_token, -0.1, &[])],
|
||||
input_logprobs: vec![position(input_token, -0.2, &[])],
|
||||
})),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
}
|
||||
let extras = collected.extras.unwrap();
|
||||
assert_eq!(extras.output_logprobs[0].token.token_id, 1);
|
||||
assert_eq!(extras.output_logprobs[1].token.token_id, 2);
|
||||
assert_eq!(extras.input_logprobs[0].token.token_id, 20);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn merged_stream_preserves_choice_indexes() {
|
||||
let choice0 = stream::iter([
|
||||
Ok(GenerationOutput {
|
||||
text: "a".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
Ok(GenerationOutput {
|
||||
text: "b".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
])
|
||||
.boxed();
|
||||
let choice1 = stream::iter([Ok(GenerationOutput {
|
||||
text: "x".into(),
|
||||
..Default::default()
|
||||
})])
|
||||
.boxed();
|
||||
|
||||
let events = merge_indexed(vec![choice0, choice1])
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
let mut observed = events
|
||||
.into_iter()
|
||||
.map(|(index, event)| (index, event.unwrap().text))
|
||||
.collect::<Vec<_>>();
|
||||
observed.sort();
|
||||
assert_eq!(
|
||||
observed,
|
||||
[(0, "a".into()), (0, "b".into()), (1, "x".into())]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
use crate::{PositionLogprobs, TokenLogprob};
|
||||
|
||||
fn logprob(token_id: i32, logprob: f32) -> TokenLogprob {
|
||||
TokenLogprob {
|
||||
logprob: Some(logprob),
|
||||
token_id,
|
||||
text: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn position(token_id: i32, value: f32, top: &[(i32, f32)]) -> PositionLogprobs {
|
||||
PositionLogprobs {
|
||||
token: logprob(token_id, value),
|
||||
top: top
|
||||
.iter()
|
||||
.map(|&(token_id, logprob)| self::logprob(token_id, logprob))
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn tiny_tokenizer() -> dynamo_tokenizers::Tokenizer {
|
||||
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../../experimental/sgl-router/tests/fixtures/tiny_tokenizer.json");
|
||||
dynamo_tokenizers::Tokenizer::from_file_with_options(
|
||||
path.to_str().unwrap(),
|
||||
dynamo_tokenizers::TokenizerOptions {
|
||||
add_special_tokens: false,
|
||||
},
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
use std::sync::{
|
||||
Arc, Mutex,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
};
|
||||
|
||||
use futures::{FutureExt, StreamExt, future::BoxFuture};
|
||||
|
||||
use super::{
|
||||
GenerateTransport, GenerationService, TokenDecoder, TokenDelta, TokenStream,
|
||||
test_utils::{position, tiny_tokenizer},
|
||||
};
|
||||
use crate::{
|
||||
GenerateRequest, GenerationFinishReason, GenerationOptions, GenerationOutputExtras,
|
||||
MatchedStop, ResponseError, TokenIdsRequest,
|
||||
};
|
||||
|
||||
struct DropNotice(Arc<AtomicUsize>);
|
||||
|
||||
impl Drop for DropNotice {
|
||||
fn drop(&mut self) {
|
||||
self.0.fetch_add(1, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
struct MemoryTransport {
|
||||
pending_submission: bool,
|
||||
dropped: Arc<AtomicUsize>,
|
||||
requests: Mutex<Vec<GenerateRequest>>,
|
||||
}
|
||||
|
||||
impl GenerateTransport for MemoryTransport {
|
||||
fn generate(
|
||||
&self,
|
||||
request: GenerateRequest,
|
||||
) -> BoxFuture<'_, Result<TokenStream, ResponseError>> {
|
||||
Box::pin(async move {
|
||||
let guard = DropNotice(self.dropped.clone());
|
||||
self.requests.lock().unwrap().push(request);
|
||||
if self.pending_submission {
|
||||
futures::future::pending::<()>().await;
|
||||
}
|
||||
Ok(async_stream::stream! {
|
||||
let _guard = guard;
|
||||
for ids in [vec![104], vec![101, 108]] {
|
||||
yield Ok(TokenDelta {
|
||||
completion_tokens: ids.len() as u64,
|
||||
extras: Some(Box::new(GenerationOutputExtras {
|
||||
output_logprobs: ids.iter().map(|&id| position(id, -0.1, &[(id, -0.1)])).collect(),
|
||||
..Default::default()
|
||||
})),
|
||||
token_ids: ids,
|
||||
prompt_tokens: 1,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
futures::future::pending::<()>().await;
|
||||
}.boxed())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn transport(pending_submission: bool) -> Arc<MemoryTransport> {
|
||||
Arc::new(MemoryTransport {
|
||||
pending_submission,
|
||||
dropped: Arc::new(AtomicUsize::new(0)),
|
||||
requests: Mutex::new(Vec::new()),
|
||||
})
|
||||
}
|
||||
|
||||
fn request() -> GenerateRequest {
|
||||
TokenIdsRequest {
|
||||
rid: "generate".into(),
|
||||
input_ids: vec![65],
|
||||
options: GenerationOptions::default(),
|
||||
metadata: Default::default(),
|
||||
}
|
||||
.into()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shared_decoder_stops_across_chunks_and_releases_transport() {
|
||||
for no_stop_trim in [false, true] {
|
||||
let transport = transport(false);
|
||||
let service =
|
||||
GenerationService::new(transport.clone(), TokenDecoder::new(tiny_tokenizer()));
|
||||
let mut request = request();
|
||||
request.sampling_params.stop = vec!["he".into()];
|
||||
request.sampling_params.stop_token_ids = Some(vec![9]);
|
||||
request.sampling_params.no_stop_trim = no_stop_trim;
|
||||
request.return_text_in_logprobs = Some(true);
|
||||
let mut events = service.generate(request).await.unwrap();
|
||||
|
||||
let first = events.next().await.unwrap().unwrap();
|
||||
assert!(first.text.is_empty());
|
||||
let last = events.next().await.unwrap().unwrap();
|
||||
assert_eq!(last.text, if no_stop_trim { "he" } else { "" });
|
||||
assert_eq!(last.token_ids, [101]);
|
||||
assert_eq!(last.completion_tokens, 1);
|
||||
assert_eq!(
|
||||
last.finish_reason,
|
||||
Some(GenerationFinishReason::Stop(Some(MatchedStop::Text(
|
||||
"he".into()
|
||||
))))
|
||||
);
|
||||
let positions = &last.extras.unwrap().output_logprobs;
|
||||
assert_eq!(positions.len(), 1);
|
||||
assert_eq!(positions[0].token.text.as_deref(), Some("e"));
|
||||
assert_eq!(positions[0].top[0].text.as_deref(), Some("e"));
|
||||
// Release upstream as soon as a local stop is emitted, even if the caller
|
||||
// keeps the completed response stream alive without polling it again.
|
||||
assert_eq!(transport.dropped.load(Ordering::SeqCst), 1);
|
||||
assert!(events.next().await.is_none());
|
||||
|
||||
let sent = transport.requests.lock().unwrap();
|
||||
assert_eq!(sent[0].sampling_params.stop, ["he"]);
|
||||
assert_eq!(sent[0].sampling_params.stop_token_ids, Some(vec![9]));
|
||||
assert_eq!(sent[0].return_text_in_logprobs, Some(false));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_releases_pending_submissions_and_unpolled_streams() {
|
||||
for pending_submission in [true, false] {
|
||||
let transport = transport(pending_submission);
|
||||
let service =
|
||||
GenerationService::new(transport.clone(), TokenDecoder::new(tiny_tokenizer()));
|
||||
let submission = service.generate_many(vec![request(), request(), request()]);
|
||||
if pending_submission {
|
||||
// Poll every submission once, then cancel the aggregate future.
|
||||
assert!(submission.now_or_never().is_none());
|
||||
} else {
|
||||
let streams = submission.await.unwrap();
|
||||
assert_eq!(transport.dropped.load(Ordering::SeqCst), 0);
|
||||
drop(streams);
|
||||
}
|
||||
assert_eq!(transport.requests.lock().unwrap().len(), 3);
|
||||
assert_eq!(transport.dropped.load(Ordering::SeqCst), 3);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
//! Generated output shared by the OpenAI response paths.
|
||||
|
||||
use futures::stream::BoxStream;
|
||||
|
||||
use crate::{ResponseError, TokenIds};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum MatchedStop {
|
||||
Token(i64),
|
||||
Text(String),
|
||||
Tokens(Vec<i64>),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum GenerationFinishReason {
|
||||
Stop(Option<MatchedStop>),
|
||||
Length,
|
||||
Abort,
|
||||
ContentFilter,
|
||||
Other(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct TokenLogprob {
|
||||
pub logprob: Option<f32>,
|
||||
pub token_id: i32,
|
||||
pub text: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct PositionLogprobs {
|
||||
pub token: TokenLogprob,
|
||||
pub top: Vec<TokenLogprob>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct GenerationOutputExtras {
|
||||
pub output_logprobs: Vec<PositionLogprobs>,
|
||||
pub input_logprobs: Vec<PositionLogprobs>,
|
||||
}
|
||||
|
||||
/// One decoded engine delta. All owned buffers are moved across the boundary.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct GenerationOutput {
|
||||
pub text: String,
|
||||
pub token_ids: TokenIds,
|
||||
pub finish_reason: Option<GenerationFinishReason>,
|
||||
pub prompt_tokens: u32,
|
||||
pub completion_tokens: u64,
|
||||
pub extras: Option<Box<GenerationOutputExtras>>,
|
||||
}
|
||||
|
||||
pub type GenerationStream = BoxStream<'static, Result<GenerationOutput, ResponseError>>;
|
||||
|
||||
/// Normalized engine token delta, before renderer-owned text decoding.
|
||||
/// Completion counts are deltas; prompt counts describe the complete prompt.
|
||||
/// A successful stream includes a terminal finish reason.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(crate) struct TokenDelta {
|
||||
pub token_ids: TokenIds,
|
||||
pub finish_reason: Option<GenerationFinishReason>,
|
||||
pub prompt_tokens: u32,
|
||||
pub completion_tokens: u64,
|
||||
pub extras: Option<Box<GenerationOutputExtras>>,
|
||||
}
|
||||
|
||||
impl From<TokenDelta> for GenerationOutput {
|
||||
fn from(delta: TokenDelta) -> Self {
|
||||
Self {
|
||||
text: String::new(),
|
||||
token_ids: delta.token_ids,
|
||||
finish_reason: delta.finish_reason,
|
||||
prompt_tokens: delta.prompt_tokens,
|
||||
completion_tokens: delta.completion_tokens,
|
||||
extras: delta.extras,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
//! Transport-neutral renderer failures.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum RendererErrorKind {
|
||||
InvalidRequest,
|
||||
Tokenize,
|
||||
Unavailable,
|
||||
Internal,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Error)]
|
||||
pub enum RendererError {
|
||||
#[error("{0}")]
|
||||
Request(String),
|
||||
#[error("validation failed: {0}")]
|
||||
Validation(String),
|
||||
#[error("tokenize failed: {0}")]
|
||||
Tokenize(String),
|
||||
#[error("renderer is shutting down")]
|
||||
Unavailable,
|
||||
#[error("render preprocessing worker failed")]
|
||||
WorkerDropped,
|
||||
#[error("internal renderer error: {0}")]
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
impl From<String> for RendererError {
|
||||
fn from(message: String) -> Self {
|
||||
Self::Request(message)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for RendererError {
|
||||
fn from(message: &str) -> Self {
|
||||
Self::Request(message.to_owned())
|
||||
}
|
||||
}
|
||||
|
||||
impl RendererError {
|
||||
pub fn kind(&self) -> RendererErrorKind {
|
||||
match self {
|
||||
Self::Request(_) | Self::Validation(_) => RendererErrorKind::InvalidRequest,
|
||||
Self::Tokenize(_) => RendererErrorKind::Tokenize,
|
||||
Self::Unavailable => RendererErrorKind::Unavailable,
|
||||
Self::WorkerDropped | Self::Internal(_) => RendererErrorKind::Internal,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A host error carried through semantic processing without interpreting it.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ResponseError {
|
||||
pub kind: ResponseErrorKind,
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
impl From<RendererError> for ResponseError {
|
||||
fn from(error: RendererError) -> Self {
|
||||
let kind = match error.kind() {
|
||||
RendererErrorKind::InvalidRequest => ResponseErrorKind::InvalidRequest,
|
||||
RendererErrorKind::Unavailable => ResponseErrorKind::Unavailable,
|
||||
RendererErrorKind::Tokenize | RendererErrorKind::Internal => {
|
||||
ResponseErrorKind::Internal
|
||||
}
|
||||
};
|
||||
ResponseError {
|
||||
kind,
|
||||
message: error.to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Failure category interpreted by the receiving transport adapter.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum ResponseErrorKind {
|
||||
InvalidRequest,
|
||||
Unavailable,
|
||||
Internal,
|
||||
Upstream(UpstreamErrorCode),
|
||||
}
|
||||
|
||||
/// Original upstream code, preserved without imposing response transport policy.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum UpstreamErrorCode {
|
||||
Http(u16),
|
||||
}
|
||||
@@ -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;
|
||||
@@ -0,0 +1,720 @@
|
||||
//! Process launch configuration for the standalone renderer.
|
||||
|
||||
use std::collections::{BTreeSet, HashMap};
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use clap::{Parser, ValueEnum};
|
||||
use hf_hub::api::tokio::{ApiBuilder, ApiRepo};
|
||||
use hf_hub::{Cache, Repo, RepoType};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::preprocessing::{resolve_model_file, resolve_tokenizer_file};
|
||||
use crate::{RendererConfig, RendererLimits, RendererRuntimeConfig, SamplingDefaults, serve};
|
||||
|
||||
const DEFAULT_CONTEXT_LEN: u64 = 2048;
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "sglang-renderer",
|
||||
about = "Run the SGLang Rust renderer with an optional SGLang engine"
|
||||
)]
|
||||
struct Cli {
|
||||
/// Model directory, config file, or Hugging Face repository id.
|
||||
#[arg(value_name = "MODEL")]
|
||||
model: String,
|
||||
|
||||
/// Optional SGLang engine origin exposing /generate.
|
||||
///
|
||||
/// When omitted, only rendering and tokenization routes are served.
|
||||
#[arg(long, value_name = "URL")]
|
||||
engine_url: Option<String>,
|
||||
|
||||
/// Proxy routes not owned by the renderer to the SGLang engine origin.
|
||||
#[arg(long, requires = "engine_url")]
|
||||
proxy_unhandled_routes: bool,
|
||||
|
||||
#[arg(long)]
|
||||
tokenizer_path: Option<String>,
|
||||
#[arg(long)]
|
||||
revision: Option<String>,
|
||||
#[arg(long)]
|
||||
served_model_name: Option<String>,
|
||||
#[arg(long, default_value_t = IpAddr::V4(Ipv4Addr::LOCALHOST))]
|
||||
host: IpAddr,
|
||||
#[arg(long, default_value_t = 30000)]
|
||||
port: u16,
|
||||
#[arg(long, default_value_t = 2)]
|
||||
http_workers: usize,
|
||||
#[arg(long, default_value_t = 1)]
|
||||
tokenizer_workers: usize,
|
||||
#[arg(long, default_value_t = 128)]
|
||||
queue_capacity: usize,
|
||||
#[arg(long)]
|
||||
chat_template: Option<String>,
|
||||
#[arg(long)]
|
||||
tool_call_parser: Option<String>,
|
||||
#[arg(long)]
|
||||
reasoning_parser: Option<String>,
|
||||
#[arg(long, value_parser = parse_json_object)]
|
||||
default_chat_template_kwargs: Option<HashMap<String, Value>>,
|
||||
#[arg(long, value_enum, default_value_t)]
|
||||
sampling_defaults: SamplingDefaultsSource,
|
||||
/// Already-resolved sampling defaults. When set with context length and
|
||||
/// vocabulary size, model metadata is not reopened by this process.
|
||||
#[arg(long, value_parser = parse_sampling_defaults)]
|
||||
resolved_sampling_params: Option<SamplingDefaults>,
|
||||
#[arg(long)]
|
||||
context_length: Option<u64>,
|
||||
#[arg(long)]
|
||||
vocab_size: Option<u64>,
|
||||
#[arg(long, default_value_t = 0)]
|
||||
num_reserved_tokens: u64,
|
||||
#[arg(long)]
|
||||
allow_auto_truncate: bool,
|
||||
#[arg(long)]
|
||||
enable_return_hidden_states: bool,
|
||||
#[arg(long)]
|
||||
stream_response_default_include_usage: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, ValueEnum)]
|
||||
enum SamplingDefaultsSource {
|
||||
#[default]
|
||||
Model,
|
||||
Openai,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct DirectArgs {
|
||||
model: String,
|
||||
engine_url: Option<String>,
|
||||
proxy_unhandled_routes: bool,
|
||||
tokenizer_path: String,
|
||||
revision: Option<String>,
|
||||
served_model_name: String,
|
||||
http_addr: SocketAddr,
|
||||
http_workers: usize,
|
||||
tokenizer_workers: usize,
|
||||
queue_capacity: usize,
|
||||
chat_template: Option<String>,
|
||||
tool_call_parser: Option<String>,
|
||||
reasoning_parser: Option<String>,
|
||||
default_chat_template_kwargs: HashMap<String, Value>,
|
||||
sampling_defaults: SamplingDefaultsSource,
|
||||
resolved_sampling_params: Option<SamplingDefaults>,
|
||||
context_length: Option<u64>,
|
||||
vocab_size: Option<u64>,
|
||||
num_reserved_tokens: u64,
|
||||
allow_auto_truncate: bool,
|
||||
enable_return_hidden_states: bool,
|
||||
stream_response_default_include_usage: bool,
|
||||
}
|
||||
|
||||
pub fn run_cli() -> Result<(), String> {
|
||||
let args = Cli::parse().into_direct_args();
|
||||
let http_workers = args.http_workers;
|
||||
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env().unwrap_or_else(|_| "info".into()),
|
||||
)
|
||||
.init();
|
||||
|
||||
let runtime = tokio::runtime::Builder::new_multi_thread()
|
||||
.worker_threads(http_workers.max(1))
|
||||
.enable_all()
|
||||
.build()
|
||||
.map_err(|error| format!("building renderer runtime failed: {error}"))?;
|
||||
runtime.block_on(async { serve(args.resolve().await?).await })
|
||||
}
|
||||
|
||||
impl Cli {
|
||||
fn into_direct_args(self) -> DirectArgs {
|
||||
let model = self.model;
|
||||
let tokenizer_path = self.tokenizer_path.unwrap_or_else(|| model.clone());
|
||||
let served_model_name = self.served_model_name.unwrap_or_else(|| model.clone());
|
||||
let http_addr = SocketAddr::new(self.host, self.port);
|
||||
DirectArgs {
|
||||
model,
|
||||
engine_url: self.engine_url,
|
||||
proxy_unhandled_routes: self.proxy_unhandled_routes,
|
||||
tokenizer_path,
|
||||
revision: self.revision,
|
||||
served_model_name,
|
||||
http_addr,
|
||||
http_workers: self.http_workers,
|
||||
tokenizer_workers: self.tokenizer_workers,
|
||||
queue_capacity: self.queue_capacity,
|
||||
chat_template: self.chat_template,
|
||||
tool_call_parser: self.tool_call_parser,
|
||||
reasoning_parser: self.reasoning_parser,
|
||||
default_chat_template_kwargs: self.default_chat_template_kwargs.unwrap_or_default(),
|
||||
sampling_defaults: self.sampling_defaults,
|
||||
resolved_sampling_params: self.resolved_sampling_params,
|
||||
context_length: self.context_length,
|
||||
vocab_size: self.vocab_size,
|
||||
num_reserved_tokens: self.num_reserved_tokens,
|
||||
allow_auto_truncate: self.allow_auto_truncate,
|
||||
enable_return_hidden_states: self.enable_return_hidden_states,
|
||||
stream_response_default_include_usage: self.stream_response_default_include_usage,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl DirectArgs {
|
||||
async fn resolve(self) -> Result<RendererRuntimeConfig, String> {
|
||||
let (context_len, vocab_size, default_sampling_params) = match self.resolved_sampling_params
|
||||
{
|
||||
Some(default_sampling_params) => {
|
||||
let context_len = self.context_length.ok_or_else(|| {
|
||||
"--resolved-sampling-params requires --context-length".to_string()
|
||||
})?;
|
||||
let vocab_size = self.vocab_size.ok_or_else(|| {
|
||||
"--resolved-sampling-params requires --vocab-size".to_string()
|
||||
})?;
|
||||
(context_len, vocab_size, default_sampling_params)
|
||||
}
|
||||
None => {
|
||||
let files = resolve_required_files(
|
||||
&self.model,
|
||||
&self.tokenizer_path,
|
||||
self.revision.as_deref(),
|
||||
)
|
||||
.await?;
|
||||
let model_config = read_json(&files.config_path)?;
|
||||
let derived_context_len = derive_context_len(&model_config)?;
|
||||
let context_len = match self.context_length {
|
||||
Some(context_len)
|
||||
if context_len > derived_context_len && !allow_longer_context() =>
|
||||
{
|
||||
return Err(format!(
|
||||
"user-specified context length {context_len} exceeds the model-derived context length {derived_context_len}; set SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN=1 to allow it"
|
||||
));
|
||||
}
|
||||
Some(context_len) => context_len,
|
||||
None => derived_context_len,
|
||||
};
|
||||
let vocab_size = self
|
||||
.vocab_size
|
||||
.or_else(|| derive_vocab_size(&model_config))
|
||||
.ok_or_else(|| {
|
||||
"model config does not define vocab_size; pass --vocab-size explicitly"
|
||||
.to_string()
|
||||
})?;
|
||||
let default_sampling_params = match self.sampling_defaults {
|
||||
SamplingDefaultsSource::Openai => SamplingDefaults::default(),
|
||||
SamplingDefaultsSource::Model => files
|
||||
.generation_config_path
|
||||
.as_deref()
|
||||
.map(read_sampling_defaults)
|
||||
.transpose()?
|
||||
.unwrap_or_default(),
|
||||
};
|
||||
(context_len, vocab_size, default_sampling_params)
|
||||
}
|
||||
};
|
||||
|
||||
Ok(RendererRuntimeConfig {
|
||||
http_addr: self.http_addr,
|
||||
http_workers: self.http_workers,
|
||||
tokenizer_workers: self.tokenizer_workers,
|
||||
queue_capacity: self.queue_capacity,
|
||||
engine_url: self.engine_url,
|
||||
proxy_unhandled_routes: self.proxy_unhandled_routes,
|
||||
renderer: RendererConfig {
|
||||
served_model_name: self.served_model_name,
|
||||
tokenizer_path: self.tokenizer_path,
|
||||
revision: self.revision,
|
||||
model_path: self.model,
|
||||
chat_template: self.chat_template,
|
||||
tool_call_parser: self.tool_call_parser,
|
||||
reasoning_parser: self.reasoning_parser,
|
||||
default_chat_template_kwargs: self.default_chat_template_kwargs,
|
||||
stream_response_default_include_usage: self.stream_response_default_include_usage,
|
||||
default_sampling_params,
|
||||
limits: RendererLimits {
|
||||
vocab_size,
|
||||
context_len,
|
||||
num_reserved_tokens: self.num_reserved_tokens,
|
||||
allow_auto_truncate: self.allow_auto_truncate,
|
||||
enable_return_hidden_states: self.enable_return_hidden_states,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ResolvedFiles {
|
||||
config_path: PathBuf,
|
||||
generation_config_path: Option<PathBuf>,
|
||||
}
|
||||
|
||||
async fn resolve_required_files(
|
||||
model: &str,
|
||||
tokenizer: &str,
|
||||
revision: Option<&str>,
|
||||
) -> Result<ResolvedFiles, String> {
|
||||
let model_is_local = Path::new(model).exists();
|
||||
let tokenizer_is_local = Path::new(tokenizer).exists();
|
||||
let mut config_path = resolve_model_file(model, revision, "config.json").map(PathBuf::from);
|
||||
let mut tokenizer_ready = resolve_tokenizer_file(tokenizer, revision).is_some();
|
||||
|
||||
if model_is_local && config_path.is_none() {
|
||||
return Err(format!(
|
||||
"local model source {model:?} does not contain config.json"
|
||||
));
|
||||
}
|
||||
if tokenizer_is_local && !tokenizer_ready {
|
||||
return Err(format!(
|
||||
"local tokenizer source {tokenizer:?} does not contain tokenizer.json, tiktoken.model, or *.tiktoken"
|
||||
));
|
||||
}
|
||||
|
||||
let need_model = config_path.is_none();
|
||||
let need_tokenizer = !tokenizer_ready;
|
||||
if need_model || need_tokenizer {
|
||||
if offline_mode() {
|
||||
return Err(format!(
|
||||
"required renderer metadata is not cached for model {model:?} and tokenizer {tokenizer:?}, and HF_HUB_OFFLINE is enabled"
|
||||
));
|
||||
}
|
||||
if model == tokenizer {
|
||||
download_repository(model, revision, need_model, need_tokenizer).await?;
|
||||
} else {
|
||||
if need_model {
|
||||
download_repository(model, revision, true, false).await?;
|
||||
}
|
||||
if need_tokenizer {
|
||||
download_repository(tokenizer, revision, false, true).await?;
|
||||
}
|
||||
}
|
||||
config_path = resolve_model_file(model, revision, "config.json").map(PathBuf::from);
|
||||
tokenizer_ready = resolve_tokenizer_file(tokenizer, revision).is_some();
|
||||
}
|
||||
|
||||
let config_path = config_path.ok_or_else(|| {
|
||||
format!(
|
||||
"model {model:?} does not expose config.json at revision {:?}",
|
||||
revision.unwrap_or("main")
|
||||
)
|
||||
})?;
|
||||
if !tokenizer_ready {
|
||||
return Err(format!(
|
||||
"tokenizer {tokenizer:?} does not expose tokenizer.json, tiktoken.model, or *.tiktoken at revision {:?}",
|
||||
revision.unwrap_or("main")
|
||||
));
|
||||
}
|
||||
let generation_config_path =
|
||||
resolve_model_file(model, revision, "generation_config.json").map(PathBuf::from);
|
||||
Ok(ResolvedFiles {
|
||||
config_path,
|
||||
generation_config_path,
|
||||
})
|
||||
}
|
||||
|
||||
async fn download_repository(
|
||||
repo_id: &str,
|
||||
revision: Option<&str>,
|
||||
include_model_metadata: bool,
|
||||
include_tokenizer: bool,
|
||||
) -> Result<(), String> {
|
||||
let mut builder = ApiBuilder::from_env()
|
||||
.with_cache_dir(hf_cache().path().clone())
|
||||
.with_progress(false);
|
||||
if let Ok(token) = std::env::var("HF_TOKEN")
|
||||
&& !token.is_empty()
|
||||
{
|
||||
builder = builder.with_token(Some(token));
|
||||
}
|
||||
let api = builder
|
||||
.build()
|
||||
.map_err(|error| format!("building Hugging Face client failed: {error}"))?;
|
||||
let repo = api.repo(Repo::with_revision(
|
||||
repo_id.to_string(),
|
||||
RepoType::Model,
|
||||
revision.unwrap_or("main").to_string(),
|
||||
));
|
||||
let info = repo.info().await.map_err(|error| {
|
||||
format!(
|
||||
"fetching Hugging Face metadata for {repo_id:?} at revision {:?} failed: {error}",
|
||||
revision.unwrap_or("main")
|
||||
)
|
||||
})?;
|
||||
let siblings = info
|
||||
.siblings
|
||||
.into_iter()
|
||||
.map(|sibling| sibling.rfilename)
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
if include_model_metadata {
|
||||
if !siblings.contains("config.json") {
|
||||
return Err(format!(
|
||||
"Hugging Face model {repo_id:?} does not contain config.json"
|
||||
));
|
||||
}
|
||||
download_file(&repo, repo_id, "config.json").await?;
|
||||
if siblings.contains("generation_config.json") {
|
||||
download_file(&repo, repo_id, "generation_config.json").await?;
|
||||
}
|
||||
}
|
||||
if include_tokenizer {
|
||||
for filename in ["tokenizer_config.json", "config.json"] {
|
||||
if siblings.contains(filename) {
|
||||
download_file(&repo, repo_id, filename).await?;
|
||||
}
|
||||
}
|
||||
let mut tokenizer_names = Vec::new();
|
||||
if siblings.contains("tokenizer.json") {
|
||||
tokenizer_names.push("tokenizer.json");
|
||||
}
|
||||
if siblings.contains("tiktoken.model") {
|
||||
tokenizer_names.push("tiktoken.model");
|
||||
} else if let Some(name) = siblings.iter().find(|name| name.ends_with(".tiktoken")) {
|
||||
tokenizer_names.push(name);
|
||||
}
|
||||
if tokenizer_names.is_empty() {
|
||||
return Err(format!(
|
||||
"Hugging Face model {repo_id:?} does not contain tokenizer.json, tiktoken.model, or *.tiktoken"
|
||||
));
|
||||
}
|
||||
for tokenizer_name in tokenizer_names {
|
||||
download_file(&repo, repo_id, tokenizer_name).await?;
|
||||
}
|
||||
let template_name = ["chat_template.json", "chat_template.jinja"]
|
||||
.into_iter()
|
||||
.find(|name| siblings.contains(*name))
|
||||
.or_else(|| {
|
||||
siblings
|
||||
.iter()
|
||||
.find(|name| name.ends_with(".jinja"))
|
||||
.map(String::as_str)
|
||||
});
|
||||
if let Some(template_name) = template_name {
|
||||
download_file(&repo, repo_id, template_name).await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn download_file(repo: &ApiRepo, repo_id: &str, filename: &str) -> Result<PathBuf, String> {
|
||||
repo.get(filename).await.map_err(|error| {
|
||||
format!("downloading {filename:?} for Hugging Face model {repo_id:?} failed: {error}")
|
||||
})
|
||||
}
|
||||
|
||||
fn hf_cache() -> Cache {
|
||||
["HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE"]
|
||||
.iter()
|
||||
.find_map(|name| std::env::var(name).ok())
|
||||
.map(PathBuf::from)
|
||||
.map(Cache::new)
|
||||
.unwrap_or_else(Cache::from_env)
|
||||
}
|
||||
|
||||
fn read_json(path: &Path) -> Result<Value, String> {
|
||||
let contents = std::fs::read_to_string(path)
|
||||
.map_err(|error| format!("reading {} failed: {error}", path.display()))?;
|
||||
serde_json::from_str(&contents)
|
||||
.map_err(|error| format!("parsing {} failed: {error}", path.display()))
|
||||
}
|
||||
|
||||
fn read_sampling_defaults(path: &Path) -> Result<SamplingDefaults, String> {
|
||||
let value = read_json(path)?;
|
||||
serde_json::from_value(value).map_err(|error| {
|
||||
format!(
|
||||
"parsing sampling defaults from {} failed: {error}",
|
||||
path.display()
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn derive_context_len(config: &Value) -> Result<u64, String> {
|
||||
let text = effective_text_config(config);
|
||||
let factor = inherited_value(text, config, "rope_scaling")
|
||||
.and_then(Value::as_object)
|
||||
.map(|rope| {
|
||||
if rope.contains_key("original_max_position_embeddings")
|
||||
|| rope.get("rope_type").and_then(Value::as_str) == Some("llama3")
|
||||
{
|
||||
1.0
|
||||
} else {
|
||||
rope.get("factor").and_then(Value::as_f64).unwrap_or(1.0)
|
||||
}
|
||||
})
|
||||
.unwrap_or(1.0);
|
||||
for key in [
|
||||
"max_sequence_length",
|
||||
"seq_length",
|
||||
"max_seq_len",
|
||||
"model_max_length",
|
||||
"max_position_embeddings",
|
||||
] {
|
||||
if let Some(value) = inherited_value(text, config, key).and_then(Value::as_u64) {
|
||||
let scaled = factor * value as f64;
|
||||
if !scaled.is_finite() || scaled <= 0.0 || scaled > u64::MAX as f64 {
|
||||
return Err(format!(
|
||||
"invalid context length {value} with rope scaling factor {factor}"
|
||||
));
|
||||
}
|
||||
return Ok(scaled as u64);
|
||||
}
|
||||
}
|
||||
Ok(DEFAULT_CONTEXT_LEN)
|
||||
}
|
||||
|
||||
fn derive_vocab_size(config: &Value) -> Option<u64> {
|
||||
let text = effective_text_config(config);
|
||||
let architecture = config
|
||||
.get("architectures")
|
||||
.and_then(Value::as_array)
|
||||
.and_then(|architectures| architectures.first())
|
||||
.and_then(Value::as_str);
|
||||
let key = if architecture == Some("GlmImageForConditionalGeneration") {
|
||||
"vision_vocab_size"
|
||||
} else {
|
||||
"vocab_size"
|
||||
};
|
||||
inherited_value(text, config, key).and_then(Value::as_u64)
|
||||
}
|
||||
|
||||
fn effective_text_config(config: &Value) -> &Value {
|
||||
let is_non_hf_llava = config
|
||||
.get("architectures")
|
||||
.and_then(Value::as_array)
|
||||
.and_then(|architectures| architectures.first())
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|architecture| {
|
||||
architecture.starts_with("Llava") && architecture.ends_with("ForCausalLM")
|
||||
});
|
||||
if is_non_hf_llava {
|
||||
return config;
|
||||
}
|
||||
if let Some(thinker) = config.get("thinker_config") {
|
||||
return thinker.get("text_config").unwrap_or(thinker);
|
||||
}
|
||||
for key in ["llm_config", "language_config", "text_config"] {
|
||||
if let Some(text) = config.get(key) {
|
||||
return text;
|
||||
}
|
||||
}
|
||||
config
|
||||
}
|
||||
|
||||
fn inherited_value<'a>(text: &'a Value, root: &'a Value, key: &str) -> Option<&'a Value> {
|
||||
text.get(key).or_else(|| root.get(key))
|
||||
}
|
||||
|
||||
fn parse_json_object(value: &str) -> Result<HashMap<String, Value>, String> {
|
||||
serde_json::from_str(value).map_err(|error| format!("expected a JSON object: {error}"))
|
||||
}
|
||||
|
||||
fn parse_sampling_defaults(value: &str) -> Result<SamplingDefaults, String> {
|
||||
serde_json::from_str(value)
|
||||
.map_err(|error| format!("expected resolved sampling parameters as JSON: {error}"))
|
||||
}
|
||||
|
||||
fn offline_mode() -> bool {
|
||||
std::env::var("HF_HUB_OFFLINE").ok().is_some_and(|value| {
|
||||
matches!(
|
||||
value.to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "on"
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn allow_longer_context() -> bool {
|
||||
std::env::var("SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN")
|
||||
.ok()
|
||||
.is_some_and(|value| {
|
||||
matches!(
|
||||
value.to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "on"
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::fs;
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn direct_cli(model: &Path) -> Cli {
|
||||
Cli::try_parse_from(["sglang-renderer", model.to_str().unwrap()]).unwrap()
|
||||
}
|
||||
|
||||
fn fixture_model(config: Value, generation_config: Option<Value>) -> PathBuf {
|
||||
let directory =
|
||||
std::env::temp_dir().join(format!("sglang-renderer-{}", uuid::Uuid::new_v4()));
|
||||
fs::create_dir(&directory).unwrap();
|
||||
fs::write(directory.join("config.json"), config.to_string()).unwrap();
|
||||
fs::write(directory.join("tokenizer.json"), "{}").unwrap();
|
||||
if let Some(generation_config) = generation_config {
|
||||
fs::write(
|
||||
directory.join("generation_config.json"),
|
||||
generation_config.to_string(),
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
directory
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cli_uses_sglang_renderer_defaults() {
|
||||
let directory = fixture_model(
|
||||
json!({"vocab_size": 128, "max_position_embeddings": 4096}),
|
||||
None,
|
||||
);
|
||||
let args = direct_cli(&directory).into_direct_args();
|
||||
|
||||
assert_eq!(args.served_model_name, directory.to_string_lossy());
|
||||
assert_eq!(args.tokenizer_path, directory.to_string_lossy());
|
||||
assert_eq!(args.http_addr, "127.0.0.1:30000".parse().unwrap());
|
||||
assert_eq!(args.http_workers, 2);
|
||||
assert_eq!(args.tokenizer_workers, 1);
|
||||
assert_eq!(args.queue_capacity, 128);
|
||||
assert_eq!(args.engine_url, None);
|
||||
assert_eq!(args.sampling_defaults, SamplingDefaultsSource::Model);
|
||||
assert_eq!(args.resolved_sampling_params, None);
|
||||
fs::remove_dir_all(directory).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxying_unhandled_routes_requires_an_engine_url() {
|
||||
let error = Cli::try_parse_from(["sglang-renderer", "model", "--proxy-unhandled-routes"])
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(
|
||||
error.kind(),
|
||||
clap::error::ErrorKind::MissingRequiredArgument
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_resolution_matches_model_metadata_and_cli_overrides() {
|
||||
let directory = fixture_model(
|
||||
json!({
|
||||
"vocab_size": 10,
|
||||
"max_position_embeddings": 8192,
|
||||
"thinker_config": {
|
||||
"text_config": {
|
||||
"vocab_size": 128,
|
||||
"max_position_embeddings": 4096,
|
||||
"rope_scaling": {"factor": 2.0}
|
||||
}
|
||||
}
|
||||
}),
|
||||
Some(json!({
|
||||
"temperature": 0.7,
|
||||
"top_p": 0.9,
|
||||
"top_k": 20,
|
||||
"min_p": 0.1,
|
||||
"repetition_penalty": 1.05,
|
||||
"max_new_tokens": 32
|
||||
})),
|
||||
);
|
||||
let cli = Cli::try_parse_from([
|
||||
"sglang-renderer",
|
||||
directory.to_str().unwrap(),
|
||||
"--engine-url",
|
||||
"http://127.0.0.1:30001",
|
||||
"--proxy-unhandled-routes",
|
||||
"--served-model-name",
|
||||
"fixture",
|
||||
"--context-length",
|
||||
"2048",
|
||||
"--vocab-size",
|
||||
"256",
|
||||
"--num-reserved-tokens",
|
||||
"8",
|
||||
"--default-chat-template-kwargs",
|
||||
r#"{"enable_thinking":false}"#,
|
||||
])
|
||||
.unwrap();
|
||||
let config = cli.into_direct_args().resolve().await.unwrap();
|
||||
|
||||
assert_eq!(config.engine_url.as_deref(), Some("http://127.0.0.1:30001"));
|
||||
assert!(config.proxy_unhandled_routes);
|
||||
assert_eq!(config.renderer.served_model_name, "fixture");
|
||||
assert_eq!(config.renderer.limits.context_len, 2048);
|
||||
assert_eq!(config.renderer.limits.vocab_size, 256);
|
||||
assert_eq!(config.renderer.limits.num_reserved_tokens, 8);
|
||||
assert_eq!(config.renderer.default_sampling_params.top_k, Some(20));
|
||||
assert_eq!(config.renderer.default_sampling_params.min_p, Some(0.1));
|
||||
assert_eq!(
|
||||
config.renderer.default_chat_template_kwargs,
|
||||
HashMap::from([("enable_thinking".to_string(), json!(false))])
|
||||
);
|
||||
fs::remove_dir_all(directory).unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolved_metadata_does_not_reopen_a_gguf_model_source() {
|
||||
let directory =
|
||||
std::env::temp_dir().join(format!("sglang-renderer-{}", uuid::Uuid::new_v4()));
|
||||
let tokenizer = directory.join("tokenizer");
|
||||
let model = directory.join("model.gguf");
|
||||
fs::create_dir_all(&tokenizer).unwrap();
|
||||
fs::write(tokenizer.join("tokenizer.json"), "{}").unwrap();
|
||||
fs::write(&model, "not needed by the renderer").unwrap();
|
||||
|
||||
let cli = Cli::try_parse_from([
|
||||
"sglang-renderer",
|
||||
model.to_str().unwrap(),
|
||||
"--engine-url",
|
||||
"http://127.0.0.1:30001",
|
||||
"--tokenizer-path",
|
||||
tokenizer.to_str().unwrap(),
|
||||
"--context-length",
|
||||
"4096",
|
||||
"--vocab-size",
|
||||
"128",
|
||||
"--resolved-sampling-params",
|
||||
r#"{"temperature":0.7,"top_k":20}"#,
|
||||
])
|
||||
.unwrap();
|
||||
let config = cli.into_direct_args().resolve().await.unwrap();
|
||||
|
||||
assert_eq!(config.renderer.model_path, model.to_string_lossy());
|
||||
assert_eq!(config.renderer.limits.context_len, 4096);
|
||||
assert_eq!(config.renderer.limits.vocab_size, 128);
|
||||
assert_eq!(
|
||||
config.renderer.default_sampling_params,
|
||||
SamplingDefaults {
|
||||
temperature: Some(0.7),
|
||||
top_k: Some(20),
|
||||
..SamplingDefaults::default()
|
||||
}
|
||||
);
|
||||
fs::remove_dir_all(directory).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_derivation_matches_python_key_and_rope_precedence() {
|
||||
assert_eq!(
|
||||
derive_context_len(&json!({
|
||||
"seq_length": 1000,
|
||||
"max_position_embeddings": 2000,
|
||||
"rope_scaling": {"factor": 4.0}
|
||||
}))
|
||||
.unwrap(),
|
||||
4000
|
||||
);
|
||||
assert_eq!(
|
||||
derive_context_len(&json!({
|
||||
"max_position_embeddings": 2000,
|
||||
"rope_scaling": {
|
||||
"factor": 4.0,
|
||||
"original_max_position_embeddings": 2000
|
||||
}
|
||||
}))
|
||||
.unwrap(),
|
||||
2000
|
||||
);
|
||||
assert_eq!(derive_context_len(&json!({})).unwrap(), 2048);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
//! Reusable request preprocessing for SGLang.
|
||||
//!
|
||||
//! The core renders normalized chat requests, lowers textual completions,
|
||||
//! tokenizes prompts, and produces the token-in contract consumed by SGLang.
|
||||
//! OpenAI operations and generation decoding are independent of transport.
|
||||
//! The optional `http` feature adds HTTP adapters, the SGLang HTTP engine client,
|
||||
//! and the process runtime. Protocol adapters own middleware and framing;
|
||||
//! shared services own request preparation, submission policy, and decoding.
|
||||
|
||||
mod config;
|
||||
// Shared serving code is compiled without HTTP; production adapters are optional.
|
||||
#[cfg_attr(not(feature = "http"), allow(dead_code))]
|
||||
mod engine;
|
||||
mod error;
|
||||
mod frontend;
|
||||
#[cfg(feature = "http")]
|
||||
mod launcher;
|
||||
#[cfg_attr(not(feature = "http"), allow(dead_code))]
|
||||
mod openai;
|
||||
mod postprocessing;
|
||||
mod preprocessing;
|
||||
#[cfg(feature = "http")]
|
||||
mod runtime;
|
||||
mod types;
|
||||
|
||||
pub use config::{RendererConfig, RendererLimits, SamplingDefaults};
|
||||
pub(crate) use engine::{
|
||||
GenerationFinishReason, GenerationOutput, GenerationOutputExtras, GenerationStream,
|
||||
MatchedStop, PositionLogprobs, TokenLogprob,
|
||||
};
|
||||
pub use error::{
|
||||
RendererError, RendererErrorKind, ResponseError, ResponseErrorKind, UpstreamErrorCode,
|
||||
};
|
||||
#[cfg(feature = "http")]
|
||||
pub use launcher::run_cli;
|
||||
pub use postprocessing::{
|
||||
ChatEvent, ChatFinishReason, ChatResponseProcessor, ChatToolCallDelta, DecodedChatEvent,
|
||||
};
|
||||
pub(crate) use preprocessing::ChatFormatter;
|
||||
pub(crate) use preprocessing::SamplingParamsOverrides;
|
||||
pub(crate) use preprocessing::{ChatPreprocessor, LoweredChat};
|
||||
pub use preprocessing::{
|
||||
ChatRequest, DynamoTokenizer, PreparedChat, ReasoningEffort, RendererService, SamplingParams,
|
||||
TextTokenizer, load_tokenizer,
|
||||
};
|
||||
pub use preprocessing::{
|
||||
GenerateRequest, GenerateRequestMetadata, GenerateSamplingParams, GenerationOptions,
|
||||
TextRequest, TokenIdsRequest,
|
||||
};
|
||||
#[cfg(feature = "http")]
|
||||
pub use runtime::{RendererRuntimeConfig, serve};
|
||||
pub use types::{OneOrMany, TokenIds};
|
||||
@@ -0,0 +1,8 @@
|
||||
fn main() {
|
||||
sglang_renderer::run_cli().unwrap_or_else(|error| exit(error));
|
||||
}
|
||||
|
||||
fn exit(message: impl std::fmt::Display) -> ! {
|
||||
eprintln!("sglang-renderer: {message}");
|
||||
std::process::exit(2)
|
||||
}
|
||||
@@ -0,0 +1,917 @@
|
||||
//! OpenAI chat preparation, response aggregation, and typed chunks.
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::{
|
||||
ChatEvent, ChatFinishReason, ChatResponseProcessor, ChatToolCallDelta, DecodedChatEvent,
|
||||
GenerationFinishReason, GenerationOutput, GenerationOutputExtras, GenerationStream,
|
||||
ResponseError,
|
||||
};
|
||||
use dynamo_protocols::types::{
|
||||
ChatChoice, ChatChoiceLogprobs, ChatChoiceStream, ChatCompletionMessageContent,
|
||||
ChatCompletionMessageToolCall, ChatCompletionMessageToolCallChunk,
|
||||
ChatCompletionResponseMessage, ChatCompletionStreamResponseDelta,
|
||||
ChatCompletionStreamResponseDeltaFunctionCall, ChatCompletionTokenLogprob, CompletionUsage,
|
||||
CreateChatCompletionResponse, CreateChatCompletionStreamResponse,
|
||||
FinishReason as OpenAIFinishReason, FunctionCall, FunctionCallStream, FunctionType, Role,
|
||||
ServiceTier as ChatServiceTier, TopLogprobs,
|
||||
};
|
||||
use futures::StreamExt;
|
||||
use serde::Serialize;
|
||||
|
||||
use super::protocol::{ChatCompletionRequest, lower_chat_request};
|
||||
use super::{completion_usage, unix_seconds_u32};
|
||||
use crate::engine::response::merge_indexed;
|
||||
|
||||
pub(crate) struct ChatResponseContext {
|
||||
pub(crate) response_id: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) created: u32,
|
||||
pub(crate) want_logprobs: bool,
|
||||
pub(crate) include_usage: bool,
|
||||
pub(crate) service_tier: Option<ChatServiceTier>,
|
||||
}
|
||||
|
||||
pub(crate) async fn prepare_request(
|
||||
renderer: &crate::RendererService,
|
||||
request: ChatCompletionRequest,
|
||||
) -> Result<(String, crate::PreparedChat), ResponseError> {
|
||||
let (response_id, request) = lower_chat_request(renderer.config(), request)?;
|
||||
let chat = renderer.prepare_chat(request).await?;
|
||||
Ok((response_id, chat))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn unary_chat(
|
||||
submitted: Vec<(usize, GenerationStream)>,
|
||||
response_processor: ChatResponseProcessor,
|
||||
response_id: String,
|
||||
model: String,
|
||||
created: u32,
|
||||
want_logprobs: bool,
|
||||
service_tier: Option<ChatServiceTier>,
|
||||
) -> Result<CreateChatCompletionResponse, ResponseError> {
|
||||
let choice_count = submitted.len();
|
||||
let mut accumulated = (0..choice_count)
|
||||
.map(|_| UnaryChatChoice::default())
|
||||
.collect::<Vec<_>>();
|
||||
let mut prompt_tokens = 0u32;
|
||||
let mut completion_tokens = 0u64;
|
||||
let parsed = semantic_chat_stream(submitted, response_processor, want_logprobs);
|
||||
futures::pin_mut!(parsed);
|
||||
while let Some(item) = parsed.next().await {
|
||||
match item {
|
||||
Ok(ChatEvent::Role { .. }) => {}
|
||||
Ok(ChatEvent::Delta {
|
||||
choice,
|
||||
content,
|
||||
reasoning_content,
|
||||
tool_calls,
|
||||
finish_reason,
|
||||
logprobs,
|
||||
}) => {
|
||||
let Some(choice) = accumulated.get_mut(choice) else {
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Internal,
|
||||
message: "chat response choice is out of range".into(),
|
||||
});
|
||||
};
|
||||
if let Some(content) = content {
|
||||
choice.content.push_str(&content);
|
||||
}
|
||||
if let Some(reasoning) = reasoning_content {
|
||||
choice.reasoning_content.push_str(&reasoning);
|
||||
}
|
||||
if let Some(tool_calls) = tool_calls {
|
||||
choice.extend_tool_calls(tool_calls);
|
||||
}
|
||||
if finish_reason.is_some() {
|
||||
choice.finish_reason = finish_reason;
|
||||
}
|
||||
merge_chat_logprobs(&mut choice.logprobs, logprobs);
|
||||
}
|
||||
Ok(ChatEvent::Usage {
|
||||
prompt_tokens: prompt,
|
||||
completion_tokens: completion,
|
||||
}) => {
|
||||
prompt_tokens = prompt;
|
||||
completion_tokens = completion;
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let choices = accumulated
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(index, parsed)| {
|
||||
#[allow(deprecated)]
|
||||
let message = ChatCompletionResponseMessage {
|
||||
content: (!parsed.content.is_empty())
|
||||
.then_some(ChatCompletionMessageContent::Text(parsed.content)),
|
||||
refusal: None,
|
||||
tool_calls: (!parsed.tool_calls.is_empty()).then(|| {
|
||||
parsed
|
||||
.tool_calls
|
||||
.into_values()
|
||||
.map(|call| ChatCompletionMessageToolCall {
|
||||
id: call.id,
|
||||
r#type: FunctionType::Function,
|
||||
function: FunctionCall {
|
||||
name: call.name,
|
||||
arguments: call.arguments,
|
||||
},
|
||||
})
|
||||
.collect()
|
||||
}),
|
||||
role: Role::Assistant,
|
||||
function_call: None,
|
||||
audio: None,
|
||||
// Python: `reasoning_text if reasoning_text else None`.
|
||||
reasoning_content: (!parsed.reasoning_content.is_empty())
|
||||
.then_some(parsed.reasoning_content),
|
||||
};
|
||||
ChatChoice {
|
||||
index: u32::try_from(index).unwrap_or(u32::MAX),
|
||||
message,
|
||||
finish_reason: parsed.finish_reason.map(openai_finish_reason),
|
||||
logprobs: parsed.logprobs,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(CreateChatCompletionResponse {
|
||||
id: response_id,
|
||||
choices,
|
||||
created,
|
||||
model,
|
||||
service_tier,
|
||||
system_fingerprint: None,
|
||||
object: "chat.completion".into(),
|
||||
usage: Some(completion_usage(
|
||||
prompt_tokens,
|
||||
u32::try_from(completion_tokens).unwrap_or(u32::MAX),
|
||||
)),
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct UnaryChatChoice {
|
||||
content: String,
|
||||
reasoning_content: String,
|
||||
tool_calls: BTreeMap<u32, UnaryToolCall>,
|
||||
finish_reason: Option<ChatFinishReason>,
|
||||
logprobs: Option<ChatChoiceLogprobs>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct UnaryToolCall {
|
||||
id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
}
|
||||
|
||||
impl UnaryChatChoice {
|
||||
fn extend_tool_calls(&mut self, deltas: Vec<ChatToolCallDelta>) {
|
||||
for delta in deltas {
|
||||
let call = self.tool_calls.entry(delta.index).or_default();
|
||||
if let Some(id) = delta.id {
|
||||
call.id = id;
|
||||
}
|
||||
if let Some(name) = delta.name {
|
||||
call.name = name;
|
||||
}
|
||||
if let Some(arguments) = delta.arguments {
|
||||
call.arguments.push_str(&arguments);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_chat_logprobs(
|
||||
collected: &mut Option<ChatChoiceLogprobs>,
|
||||
delta: Option<ChatChoiceLogprobs>,
|
||||
) {
|
||||
let Some(mut delta) = delta else {
|
||||
return;
|
||||
};
|
||||
let collected = collected.get_or_insert_with(|| ChatChoiceLogprobs {
|
||||
content: Some(Vec::new()),
|
||||
refusal: None,
|
||||
});
|
||||
if let Some(content) = delta.content.take() {
|
||||
collected
|
||||
.content
|
||||
.get_or_insert_with(Vec::new)
|
||||
.extend(content);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn chat_event_stream(
|
||||
submitted: Vec<(usize, GenerationStream)>,
|
||||
response_processor: ChatResponseProcessor,
|
||||
context: ChatResponseContext,
|
||||
) -> impl futures::Stream<Item = Result<CreateChatCompletionStreamResponse, ResponseError>> {
|
||||
let parsed = semantic_chat_stream(submitted, response_processor, context.want_logprobs);
|
||||
|
||||
async_stream::stream! {
|
||||
futures::pin_mut!(parsed);
|
||||
while let Some(item) = parsed.next().await {
|
||||
match item {
|
||||
Ok(ChatEvent::Role { choice }) => {
|
||||
yield Ok(chat_stream_response(
|
||||
&context.response_id,
|
||||
&context.model,
|
||||
context.created,
|
||||
context.service_tier.clone(),
|
||||
vec![ChatChoiceStream {
|
||||
index: choice as u32,
|
||||
delta: chat_delta(None, Some(Role::Assistant), None, None),
|
||||
finish_reason: None,
|
||||
logprobs: None,
|
||||
}],
|
||||
None,
|
||||
));
|
||||
}
|
||||
Ok(ChatEvent::Delta {
|
||||
choice,
|
||||
content,
|
||||
reasoning_content,
|
||||
tool_calls,
|
||||
finish_reason,
|
||||
logprobs,
|
||||
}) => {
|
||||
yield Ok(chat_stream_response(
|
||||
&context.response_id,
|
||||
&context.model,
|
||||
context.created,
|
||||
context.service_tier.clone(),
|
||||
vec![ChatChoiceStream {
|
||||
index: choice as u32,
|
||||
delta: chat_delta(
|
||||
content,
|
||||
None,
|
||||
tool_calls.map(|calls| {
|
||||
calls.into_iter().map(openai_tool_call_delta).collect()
|
||||
}),
|
||||
reasoning_content,
|
||||
),
|
||||
finish_reason: finish_reason.map(openai_finish_reason),
|
||||
logprobs,
|
||||
}],
|
||||
None,
|
||||
));
|
||||
}
|
||||
Ok(ChatEvent::Usage {
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
}) if context.include_usage => {
|
||||
yield Ok(chat_stream_response(
|
||||
&context.response_id,
|
||||
&context.model,
|
||||
context.created,
|
||||
context.service_tier.clone(),
|
||||
Vec::new(),
|
||||
Some((prompt_tokens, completion_tokens)),
|
||||
));
|
||||
}
|
||||
Ok(ChatEvent::Usage { .. }) => {}
|
||||
Err(error) => {
|
||||
yield Err(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn semantic_chat_stream(
|
||||
submitted: Vec<(usize, GenerationStream)>,
|
||||
response_processor: ChatResponseProcessor,
|
||||
want_logprobs: bool,
|
||||
) -> impl futures::Stream<Item = Result<ChatEvent, ResponseError>> {
|
||||
let raw = async_stream::stream! {
|
||||
let streams = submitted.into_iter().map(|(_, events)| events).collect();
|
||||
let mut events = merge_indexed(streams);
|
||||
while let Some((index, item)) = events.next().await {
|
||||
let output = match item {
|
||||
Ok(output) => output,
|
||||
Err(error) => {
|
||||
yield Err(error);
|
||||
break;
|
||||
}
|
||||
};
|
||||
let finish_reason = chat_finish_reason(&output);
|
||||
let logprobs = want_logprobs.then(|| chat_logprobs(output.extras.as_deref()));
|
||||
yield Ok(DecodedChatEvent {
|
||||
choice: index,
|
||||
text: output.text,
|
||||
token_ids: output.token_ids,
|
||||
finish_reason,
|
||||
logprobs,
|
||||
prompt_tokens: output.prompt_tokens,
|
||||
completion_tokens: output.completion_tokens,
|
||||
});
|
||||
}
|
||||
};
|
||||
response_processor.process_stream(raw)
|
||||
}
|
||||
|
||||
fn chat_finish_reason(output: &GenerationOutput) -> Option<ChatFinishReason> {
|
||||
output.finish_reason.as_ref().map(|reason| match reason {
|
||||
GenerationFinishReason::Length => ChatFinishReason::Length,
|
||||
GenerationFinishReason::ContentFilter => ChatFinishReason::ContentFilter,
|
||||
GenerationFinishReason::Stop(_)
|
||||
| GenerationFinishReason::Abort
|
||||
| GenerationFinishReason::Other(_) => ChatFinishReason::Stop,
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
fn chat_logprobs(extras: Option<&GenerationOutputExtras>) -> ChatChoiceLogprobs {
|
||||
let mut content = Vec::new();
|
||||
let Some(extras) = extras else {
|
||||
return ChatChoiceLogprobs {
|
||||
content: Some(content),
|
||||
refusal: None,
|
||||
};
|
||||
};
|
||||
for position in &extras.output_logprobs {
|
||||
let selected = &position.token;
|
||||
let token = selected
|
||||
.text
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("token_id:{}", selected.token_id));
|
||||
let top_logprobs = position
|
||||
.top
|
||||
.iter()
|
||||
.map(|candidate| {
|
||||
let text = candidate
|
||||
.text
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("token_id:{}", candidate.token_id));
|
||||
TopLogprobs {
|
||||
bytes: Some(text.as_bytes().to_vec()),
|
||||
token: text,
|
||||
logprob: candidate.logprob.unwrap_or(f32::NAN),
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
content.push(ChatCompletionTokenLogprob {
|
||||
bytes: Some(token.as_bytes().to_vec()),
|
||||
token,
|
||||
logprob: selected.logprob.unwrap_or(f32::NAN),
|
||||
token_id: u32::try_from(selected.token_id).ok(),
|
||||
top_logprobs,
|
||||
});
|
||||
}
|
||||
ChatChoiceLogprobs {
|
||||
content: Some(content),
|
||||
refusal: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
fn chat_delta(
|
||||
content: Option<String>,
|
||||
role: Option<Role>,
|
||||
tool_calls: Option<Vec<ChatCompletionMessageToolCallChunk>>,
|
||||
reasoning_content: Option<String>,
|
||||
) -> ChatCompletionStreamResponseDelta {
|
||||
ChatCompletionStreamResponseDelta {
|
||||
content: content.map(ChatCompletionMessageContent::Text),
|
||||
function_call: None,
|
||||
tool_calls,
|
||||
role,
|
||||
refusal: None,
|
||||
reasoning_content,
|
||||
}
|
||||
}
|
||||
|
||||
fn chat_stream_response(
|
||||
response_id: &str,
|
||||
model: &str,
|
||||
created: u32,
|
||||
service_tier: Option<ChatServiceTier>,
|
||||
choices: Vec<ChatChoiceStream>,
|
||||
usage: Option<(u32, u64)>,
|
||||
) -> CreateChatCompletionStreamResponse {
|
||||
CreateChatCompletionStreamResponse {
|
||||
id: response_id.to_owned(),
|
||||
choices,
|
||||
created,
|
||||
model: model.to_owned(),
|
||||
service_tier,
|
||||
system_fingerprint: None,
|
||||
object: "chat.completion.chunk".into(),
|
||||
usage: usage.map(|(prompt, completion)| {
|
||||
completion_usage(prompt, u32::try_from(completion).unwrap_or(u32::MAX))
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn openai_finish_reason(reason: ChatFinishReason) -> OpenAIFinishReason {
|
||||
match reason {
|
||||
ChatFinishReason::Stop => OpenAIFinishReason::Stop,
|
||||
ChatFinishReason::Length => OpenAIFinishReason::Length,
|
||||
ChatFinishReason::ContentFilter => OpenAIFinishReason::ContentFilter,
|
||||
ChatFinishReason::ToolCalls => OpenAIFinishReason::ToolCalls,
|
||||
}
|
||||
}
|
||||
|
||||
fn openai_tool_call_delta(call: ChatToolCallDelta) -> ChatCompletionMessageToolCallChunk {
|
||||
ChatCompletionMessageToolCallChunk {
|
||||
index: call.index,
|
||||
id: call.id,
|
||||
r#type: Some(FunctionType::Function),
|
||||
function: Some(FunctionCallStream {
|
||||
name: call.name,
|
||||
arguments: call.arguments,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn serialize_chat_stream_response(
|
||||
response: CreateChatCompletionStreamResponse,
|
||||
) -> String {
|
||||
serde_json::to_string(&ChatStreamResponseWire::from(&response))
|
||||
.expect("OpenAI response must serialize")
|
||||
}
|
||||
|
||||
/// The Dynamo response type omits an absent `reasoning_content`. SGLang's
|
||||
/// streaming contract emits it explicitly as `null`, so use a borrowed wire
|
||||
/// view instead of building and patching a `serde_json::Value` tree.
|
||||
#[derive(Serialize)]
|
||||
struct ChatStreamResponseWire<'a> {
|
||||
id: &'a str,
|
||||
choices: Vec<ChatChoiceStreamWire<'a>>,
|
||||
created: u32,
|
||||
model: &'a str,
|
||||
service_tier: &'a Option<ChatServiceTier>,
|
||||
system_fingerprint: &'a Option<String>,
|
||||
object: &'a str,
|
||||
usage: &'a Option<CompletionUsage>,
|
||||
}
|
||||
|
||||
impl<'a> From<&'a CreateChatCompletionStreamResponse> for ChatStreamResponseWire<'a> {
|
||||
fn from(response: &'a CreateChatCompletionStreamResponse) -> Self {
|
||||
Self {
|
||||
id: &response.id,
|
||||
choices: response
|
||||
.choices
|
||||
.iter()
|
||||
.map(ChatChoiceStreamWire::from)
|
||||
.collect(),
|
||||
created: response.created,
|
||||
model: &response.model,
|
||||
service_tier: &response.service_tier,
|
||||
system_fingerprint: &response.system_fingerprint,
|
||||
object: &response.object,
|
||||
usage: &response.usage,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ChatChoiceStreamWire<'a> {
|
||||
index: u32,
|
||||
delta: ChatDeltaWire<'a>,
|
||||
finish_reason: &'a Option<OpenAIFinishReason>,
|
||||
logprobs: &'a Option<ChatChoiceLogprobs>,
|
||||
}
|
||||
|
||||
impl<'a> From<&'a ChatChoiceStream> for ChatChoiceStreamWire<'a> {
|
||||
fn from(choice: &'a ChatChoiceStream) -> Self {
|
||||
Self {
|
||||
index: choice.index,
|
||||
delta: ChatDeltaWire::from(&choice.delta),
|
||||
finish_reason: &choice.finish_reason,
|
||||
logprobs: &choice.logprobs,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ChatDeltaWire<'a> {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
content: Option<&'a ChatCompletionMessageContent>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
function_call: Option<&'a ChatCompletionStreamResponseDeltaFunctionCall>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_calls: Option<&'a Vec<ChatCompletionMessageToolCallChunk>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
role: Option<&'a Role>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
refusal: Option<&'a String>,
|
||||
reasoning_content: Option<&'a str>,
|
||||
}
|
||||
|
||||
impl<'a> From<&'a ChatCompletionStreamResponseDelta> for ChatDeltaWire<'a> {
|
||||
fn from(delta: &'a ChatCompletionStreamResponseDelta) -> Self {
|
||||
Self {
|
||||
content: delta.content.as_ref(),
|
||||
function_call: delta.function_call.as_ref(),
|
||||
tool_calls: delta.tool_calls.as_ref(),
|
||||
role: delta.role.as_ref(),
|
||||
refusal: delta.refusal.as_ref(),
|
||||
reasoning_content: delta.reasoning_content.as_deref(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl super::OpenAIService {
|
||||
pub(crate) async fn chat(
|
||||
&self,
|
||||
request: ChatCompletionRequest,
|
||||
) -> Result<
|
||||
super::OperationResponse<CreateChatCompletionResponse, CreateChatCompletionStreamResponse>,
|
||||
ResponseError,
|
||||
> {
|
||||
use super::OperationResponse;
|
||||
let stream = request.stream.unwrap_or(false);
|
||||
let model = request.model.clone();
|
||||
let want_logprobs = request.logprobs.unwrap_or(false);
|
||||
let include_usage = request
|
||||
.stream_options
|
||||
.as_ref()
|
||||
.is_some_and(|options| options.include_usage)
|
||||
|| self.renderer.config().stream_response_default_include_usage;
|
||||
let service_tier = request.service_tier.clone();
|
||||
let (response_id, chat) = prepare_request(&self.renderer, request).await?;
|
||||
let context = ChatResponseContext {
|
||||
response_id,
|
||||
model,
|
||||
created: unix_seconds_u32(),
|
||||
want_logprobs,
|
||||
include_usage,
|
||||
service_tier,
|
||||
};
|
||||
let streams = match self.generation.generate_many(chat.requests).await {
|
||||
Ok(streams) => streams,
|
||||
Err(error) if stream => {
|
||||
return Ok(OperationResponse::Stream(
|
||||
futures::stream::once(async { Err(error) }).boxed(),
|
||||
));
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let submitted = streams.into_iter().enumerate().collect();
|
||||
if stream {
|
||||
Ok(OperationResponse::Stream(
|
||||
chat_event_stream(submitted, chat.response_processor, context).boxed(),
|
||||
))
|
||||
} else {
|
||||
unary_chat(
|
||||
submitted,
|
||||
chat.response_processor,
|
||||
context.response_id,
|
||||
context.model,
|
||||
context.created,
|
||||
context.want_logprobs,
|
||||
context.service_tier,
|
||||
)
|
||||
.await
|
||||
.map(OperationResponse::Unary)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{ChatResponseContext, chat_event_stream, chat_logprobs, unary_chat};
|
||||
use crate::openai::protocol::ChatCompletionRequest;
|
||||
use crate::openai::protocol::{chat_sampling_params, lower_chat_request};
|
||||
use crate::openai::test_utils::{chat_submitted, chunk};
|
||||
use crate::{
|
||||
ChatPreprocessor, GenerationOutputExtras, PositionLogprobs, RendererConfig, RendererLimits,
|
||||
ResponseError, SamplingDefaults, TokenLogprob,
|
||||
};
|
||||
use futures::{FutureExt, StreamExt};
|
||||
|
||||
fn request() -> ChatCompletionRequest {
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"model": "test",
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}))
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn response_processor(
|
||||
reasoning_parser: Option<&str>,
|
||||
choices: usize,
|
||||
) -> crate::ChatResponseProcessor {
|
||||
let config = RendererConfig {
|
||||
model_path: String::new(),
|
||||
served_model_name: "model".into(),
|
||||
tokenizer_path: ".".into(),
|
||||
chat_template: Some("chatml".into()),
|
||||
tool_call_parser: None,
|
||||
reasoning_parser: reasoning_parser.map(str::to_owned),
|
||||
default_chat_template_kwargs: Default::default(),
|
||||
revision: None,
|
||||
stream_response_default_include_usage: false,
|
||||
default_sampling_params: SamplingDefaults::default(),
|
||||
limits: RendererLimits {
|
||||
vocab_size: 128,
|
||||
context_len: 128,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
};
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"n": choices
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, chat) = lower_chat_request(&config, request).unwrap();
|
||||
ChatPreprocessor::new(
|
||||
&config,
|
||||
Some(crate::preprocessing::load_test_chat_formatter("chatml")),
|
||||
)
|
||||
.preprocess(chat)
|
||||
.unwrap()
|
||||
.response_processor
|
||||
}
|
||||
|
||||
fn wire_context(include_usage: bool) -> ChatResponseContext {
|
||||
ChatResponseContext {
|
||||
response_id: "chatcmpl-test".into(),
|
||||
model: "model".into(),
|
||||
created: 1,
|
||||
want_logprobs: false,
|
||||
include_usage,
|
||||
service_tier: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Python `to_sampling_params` priority: user value > model generation
|
||||
/// config (`--sampling-defaults model`) > OpenAI terminal default.
|
||||
#[test]
|
||||
fn sampling_defaults_follow_python_priority_chain() {
|
||||
let model = SamplingDefaults {
|
||||
temperature: Some(0.6),
|
||||
top_p: Some(0.9),
|
||||
top_k: Some(32),
|
||||
min_p: Some(0.1),
|
||||
repetition_penalty: Some(1.1),
|
||||
};
|
||||
// Omitted → model defaults, not the 1.0 OpenAI terminals.
|
||||
let sampling = chat_sampling_params(&request(), &model).unwrap();
|
||||
assert_eq!(sampling.temperature, 0.6);
|
||||
assert_eq!(sampling.top_p, 0.9);
|
||||
assert_eq!(sampling.top_k, 32);
|
||||
assert_eq!(sampling.min_p, 0.1);
|
||||
assert_eq!(sampling.repetition_penalty, 1.1);
|
||||
// Explicit request values win. `Option<f32>` loses precision in f64 —
|
||||
// compare with tolerance.
|
||||
let mut request = request();
|
||||
request.temperature = Some(0.2);
|
||||
request.top_p = Some(0.5);
|
||||
let sampling = chat_sampling_params(&request, &model).unwrap();
|
||||
assert!((sampling.temperature - 0.2).abs() < 1e-6);
|
||||
assert!((sampling.top_p - 0.5).abs() < 1e-6);
|
||||
}
|
||||
|
||||
/// `--sampling-defaults openai` resolves an empty model-config slice, so the
|
||||
/// conversion falls back to the OpenAI terminal defaults.
|
||||
#[test]
|
||||
fn sampling_defaults_fall_back_to_openai_terminals_in_openai_mode() {
|
||||
let openai_mode = SamplingDefaults::default();
|
||||
let sampling = chat_sampling_params(&request(), &openai_mode).unwrap();
|
||||
assert_eq!(sampling.temperature, 1.0);
|
||||
assert_eq!(sampling.top_p, 1.0);
|
||||
assert_eq!(sampling.top_k, 1 << 30);
|
||||
assert_eq!(sampling.min_p, 0.0);
|
||||
assert_eq!(sampling.repetition_penalty, 1.0);
|
||||
}
|
||||
|
||||
/// A request with no `max_tokens`/`max_completion_tokens` stays unbounded —
|
||||
/// no terminal default is imposed.
|
||||
#[test]
|
||||
fn chat_without_a_token_limit_stays_unbounded() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "test",
|
||||
"messages": [{"role": "user", "content": "hello"}]
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
chat_sampling_params(&request, &SamplingDefaults::default())
|
||||
.unwrap()
|
||||
.max_new_tokens,
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_logprobs_use_dynamo_wire_types() {
|
||||
let extras = GenerationOutputExtras {
|
||||
output_logprobs: vec![PositionLogprobs {
|
||||
token: TokenLogprob {
|
||||
logprob: Some(-0.25),
|
||||
token_id: 7,
|
||||
text: Some("x".into()),
|
||||
},
|
||||
top: vec![
|
||||
TokenLogprob {
|
||||
logprob: Some(-0.25),
|
||||
token_id: 7,
|
||||
text: Some("x".into()),
|
||||
},
|
||||
TokenLogprob {
|
||||
logprob: Some(-1.0),
|
||||
token_id: 8,
|
||||
text: Some("y".into()),
|
||||
},
|
||||
],
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
let logprobs = chat_logprobs(Some(&extras));
|
||||
let token = &logprobs.content.unwrap()[0];
|
||||
assert_eq!(token.token, "x");
|
||||
assert_eq!(token.token_id, Some(7));
|
||||
assert_eq!(token.top_logprobs.len(), 2);
|
||||
assert_eq!(token.top_logprobs[1].token, "y");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unary_chat_fans_in_choices_and_usage() {
|
||||
let (choice0, tx0) = chat_submitted(0);
|
||||
let (choice1, tx1) = chat_submitted(1);
|
||||
tx0.send(chunk("Paris", true)).await.unwrap();
|
||||
tx1.send(chunk("Paris", true)).await.unwrap();
|
||||
|
||||
let response = unary_chat(
|
||||
vec![choice0, choice1],
|
||||
response_processor(None, 2),
|
||||
"chatcmpl-test".into(),
|
||||
"model".into(),
|
||||
1,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
let value = serde_json::to_value(response.unwrap()).unwrap();
|
||||
assert_eq!(value["choices"][0]["message"]["role"], "assistant");
|
||||
assert_eq!(value["choices"][0]["message"]["content"], "Paris");
|
||||
assert_eq!(value["choices"][1]["index"], 1);
|
||||
assert_eq!(value["usage"]["prompt_tokens"], 5);
|
||||
assert_eq!(value["usage"]["completion_tokens"], 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unary_chat_separates_reasoning_content_with_parser_configured() {
|
||||
let (choice, tx) = chat_submitted(0);
|
||||
tx.send(chunk("<think>because Paris is famous</think>Paris", true))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let response = unary_chat(
|
||||
vec![choice],
|
||||
response_processor(Some("deepseek-r1"), 1),
|
||||
"chatcmpl-test".into(),
|
||||
"model".into(),
|
||||
1,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
let value = serde_json::to_value(response.unwrap()).unwrap();
|
||||
assert_eq!(
|
||||
value["choices"][0]["message"]["reasoning_content"],
|
||||
"because Paris is famous"
|
||||
);
|
||||
assert_eq!(value["choices"][0]["message"]["content"], "Paris");
|
||||
assert!(value["choices"][0]["message"]["reasoning_content"].is_string());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_chat_separates_reasoning_into_own_deltas() {
|
||||
let (choice, tx) = chat_submitted(0);
|
||||
// Force mode starts in reasoning, so the opener is stripped and the first
|
||||
// reasoning fragment streams immediately.
|
||||
tx.send(chunk("<think>be", false)).await.unwrap();
|
||||
tx.send(chunk("cause</think>Par", false)).await.unwrap();
|
||||
tx.send(chunk("is", true)).await.unwrap();
|
||||
|
||||
let stream = chat_event_stream(
|
||||
vec![choice],
|
||||
response_processor(Some("deepseek-r1"), 1),
|
||||
wire_context(true),
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
let frames: Vec<_> = stream
|
||||
.map(|chunk| serde_json::to_value(chunk.unwrap()).unwrap())
|
||||
.collect()
|
||||
.await;
|
||||
let role = &frames[0];
|
||||
let first_reasoning = &frames[1];
|
||||
let second_reasoning = &frames[2];
|
||||
let content = &frames[3];
|
||||
let terminal = &frames[4];
|
||||
assert_eq!(role["choices"][0]["delta"]["role"], "assistant");
|
||||
assert_eq!(
|
||||
first_reasoning["choices"][0]["delta"]["reasoning_content"],
|
||||
"be"
|
||||
);
|
||||
assert!(first_reasoning["choices"][0]["delta"]["content"].is_null());
|
||||
assert_eq!(
|
||||
second_reasoning["choices"][0]["delta"]["reasoning_content"],
|
||||
"cause"
|
||||
);
|
||||
assert_eq!(content["choices"][0]["delta"]["content"], "Par");
|
||||
assert!(content["choices"][0]["delta"]["reasoning_content"].is_null());
|
||||
assert_eq!(terminal["choices"][0]["delta"]["content"], "is");
|
||||
assert_eq!(terminal["choices"][0]["finish_reason"], "stop");
|
||||
assert_eq!(frames.len(), 6);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_chat_emits_role_deltas_and_usage() {
|
||||
let (choice, tx) = chat_submitted(0);
|
||||
tx.send(chunk("Par", false)).await.unwrap();
|
||||
tx.send(chunk("is", true)).await.unwrap();
|
||||
|
||||
let stream = chat_event_stream(
|
||||
vec![choice],
|
||||
response_processor(None, 1),
|
||||
wire_context(true),
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
let frames: Vec<_> = stream
|
||||
.map(|chunk| serde_json::to_value(chunk.unwrap()).unwrap())
|
||||
.collect()
|
||||
.await;
|
||||
assert_eq!(frames.len(), 4);
|
||||
let role = &frames[0];
|
||||
let delta = &frames[1];
|
||||
let terminal = &frames[2];
|
||||
let usage = &frames[3];
|
||||
assert_eq!(role["choices"][0]["delta"]["role"], "assistant");
|
||||
assert!(role["choices"][0]["delta"]["reasoning_content"].is_null());
|
||||
assert_eq!(delta["choices"][0]["delta"]["content"], "Par");
|
||||
assert!(delta["choices"][0]["delta"]["reasoning_content"].is_null());
|
||||
assert_eq!(terminal["choices"][0]["delta"]["content"], "is");
|
||||
assert!(terminal["choices"][0]["delta"]["reasoning_content"].is_null());
|
||||
assert_eq!(terminal["choices"][0]["finish_reason"], "stop");
|
||||
assert_eq!(usage["usage"]["completion_tokens"], 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_chat_waits_for_backend_output_before_role() {
|
||||
let (choice, tx) = chat_submitted(0);
|
||||
let stream = chat_event_stream(
|
||||
vec![choice],
|
||||
response_processor(None, 1),
|
||||
wire_context(false),
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
|
||||
assert!(stream.next().now_or_never().is_none());
|
||||
|
||||
tx.send(chunk("Paris", false)).await.unwrap();
|
||||
let role = serde_json::to_value(stream.next().await.unwrap().unwrap()).unwrap();
|
||||
let delta = serde_json::to_value(stream.next().await.unwrap().unwrap()).unwrap();
|
||||
assert_eq!(role["choices"][0]["delta"]["role"], "assistant");
|
||||
assert_eq!(delta["choices"][0]["delta"]["content"], "Paris");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_chat_stops_all_choices_after_error() {
|
||||
let (choice0, tx0) = chat_submitted(0);
|
||||
let (choice1, tx1) = chat_submitted(1);
|
||||
let stream = chat_event_stream(
|
||||
vec![choice0, choice1],
|
||||
response_processor(None, 2),
|
||||
wire_context(true),
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
|
||||
tx0.send(Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(429)),
|
||||
message: "out of memory".into(),
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let error = stream.next().await.unwrap().unwrap_err();
|
||||
assert_eq!(
|
||||
error.kind,
|
||||
crate::ResponseErrorKind::Upstream(crate::UpstreamErrorCode::Http(429))
|
||||
);
|
||||
assert_eq!(error.message, "out of memory");
|
||||
|
||||
// The other choice may already be ready, but it must not be polled after
|
||||
// the aggregate request has emitted an error.
|
||||
tx1.send(chunk("late", true)).await.unwrap();
|
||||
let remaining = stream.collect::<Vec<_>>().await;
|
||||
assert_eq!(remaining.len(), 1);
|
||||
assert!(
|
||||
remaining
|
||||
.into_iter()
|
||||
.all(|chunk| chunk.unwrap().choices.is_empty())
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,693 @@
|
||||
//! OpenAI completion preparation, response aggregation, and typed chunks.
|
||||
|
||||
use crate::engine::response::{collect_output, merge_indexed};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::{
|
||||
completion_usage,
|
||||
protocol::{
|
||||
CompletionRequest, lower_text_completion_request, lower_token_ids_completion_request,
|
||||
text_completion_prompts, token_ids_completion_prompts,
|
||||
},
|
||||
unix_seconds_u32,
|
||||
};
|
||||
use crate::{
|
||||
GenerateRequest, GenerationFinishReason, GenerationOutput, GenerationOutputExtras,
|
||||
GenerationStream, MatchedStop, RendererService, ResponseError, engine::TokenDecoder,
|
||||
};
|
||||
use dynamo_protocols::types::{CompletionUsage, Prompt};
|
||||
use futures::StreamExt;
|
||||
use serde::Serialize;
|
||||
|
||||
pub(crate) struct SubmittedChoice {
|
||||
pub(crate) index: usize,
|
||||
pub(crate) prompt_index: usize,
|
||||
pub(crate) echo: String,
|
||||
pub(crate) events: GenerationStream,
|
||||
}
|
||||
|
||||
pub(crate) fn attach_streams(
|
||||
metadata: Vec<(usize, usize, String)>,
|
||||
streams: Vec<GenerationStream>,
|
||||
) -> Vec<SubmittedChoice> {
|
||||
metadata
|
||||
.into_iter()
|
||||
.zip(streams)
|
||||
.map(|((index, prompt_index, echo), events)| SubmittedChoice {
|
||||
index,
|
||||
prompt_index,
|
||||
echo,
|
||||
events,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
enum MatchedStopWire {
|
||||
Token(i64),
|
||||
Text(String),
|
||||
Tokens(Vec<i64>),
|
||||
}
|
||||
|
||||
#[derive(Debug, PartialEq, Serialize)]
|
||||
struct CompletionLogprobsWire {
|
||||
tokens: Vec<String>,
|
||||
token_logprobs: Vec<Option<f64>>,
|
||||
top_logprobs: Vec<Option<BTreeMap<String, f64>>>,
|
||||
text_offset: Vec<i32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct CompletionChoiceWire {
|
||||
text: String,
|
||||
index: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
logprobs: Option<CompletionLogprobsWire>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
finish_reason: Option<String>,
|
||||
matched_stop: Option<MatchedStopWire>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct CompletionResponseWire {
|
||||
id: String,
|
||||
choices: Vec<CompletionChoiceWire>,
|
||||
created: u32,
|
||||
model: String,
|
||||
object: &'static str,
|
||||
usage: Option<CompletionUsage>,
|
||||
}
|
||||
|
||||
struct CompletionResponseContext {
|
||||
metadata: Vec<(usize, usize, String)>,
|
||||
response_id: String,
|
||||
model: String,
|
||||
created: u32,
|
||||
echo: bool,
|
||||
want_logprobs: bool,
|
||||
include_usage: bool,
|
||||
continuous_usage: bool,
|
||||
}
|
||||
|
||||
pub(crate) async fn prepare_request(
|
||||
renderer: &RendererService,
|
||||
request: &CompletionRequest,
|
||||
) -> Result<(String, Vec<GenerateRequest>), ResponseError> {
|
||||
if matches!(&request.prompt, Prompt::String(_) | Prompt::StringArray(_)) {
|
||||
let (response_id, requests) = lower_text_completion_request(renderer.config(), request)?;
|
||||
let requests = renderer.prepare_text_request_groups(requests).await?;
|
||||
Ok((response_id, requests))
|
||||
} else {
|
||||
let (response_id, requests) =
|
||||
lower_token_ids_completion_request(renderer.config(), request)?;
|
||||
let requests = renderer.prepare_token_ids_requests(requests)?;
|
||||
Ok((response_id, requests))
|
||||
}
|
||||
}
|
||||
|
||||
// Called after request preparation has validated the prompt and choice count.
|
||||
fn prepare_response(
|
||||
renderer: &RendererService,
|
||||
tokenizer: &TokenDecoder,
|
||||
request: &CompletionRequest,
|
||||
response_id: String,
|
||||
choice_count: usize,
|
||||
) -> Result<CompletionResponseContext, ResponseError> {
|
||||
let echo = request.echo.unwrap_or(false);
|
||||
let n = request.n.unwrap_or(1) as usize;
|
||||
// Echo uses the original input, even when preprocessing truncates engine input IDs.
|
||||
let prompt_echoes = if !echo {
|
||||
vec![String::new(); choice_count / n]
|
||||
} else if matches!(&request.prompt, Prompt::String(_) | Prompt::StringArray(_)) {
|
||||
text_completion_prompts(&request.prompt).map_err(crate::RendererError::from)?
|
||||
} else {
|
||||
token_ids_completion_prompts(&request.prompt)
|
||||
.map_err(crate::RendererError::from)?
|
||||
.into_iter()
|
||||
.map(|ids| tokenizer.detokenize_prompt(ids))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
};
|
||||
let metadata = prompt_echoes
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.flat_map(|(prompt_index, echo)| {
|
||||
(0..n).map(move |choice| (prompt_index * n + choice, prompt_index, echo.clone()))
|
||||
})
|
||||
.collect();
|
||||
Ok(CompletionResponseContext {
|
||||
metadata,
|
||||
response_id,
|
||||
model: request.model.clone(),
|
||||
created: unix_seconds_u32(),
|
||||
echo,
|
||||
want_logprobs: request.logprobs.is_some(),
|
||||
include_usage: request
|
||||
.stream_options
|
||||
.as_ref()
|
||||
.is_some_and(|options| options.include_usage)
|
||||
|| renderer.config().stream_response_default_include_usage,
|
||||
continuous_usage: request
|
||||
.stream_options
|
||||
.as_ref()
|
||||
.is_some_and(|options| options.continuous_usage_stats),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn unary_completion(
|
||||
submitted: Vec<SubmittedChoice>,
|
||||
response_id: String,
|
||||
model: String,
|
||||
created: u32,
|
||||
echo: bool,
|
||||
want_logprobs: bool,
|
||||
) -> Result<CompletionResponseWire, ResponseError> {
|
||||
// Every request is already submitted, so draining in choice order does not
|
||||
// serialize generation. The non-streaming native path sends one terminal
|
||||
// result, and the accumulator also tolerates intermediate frames.
|
||||
let mut choices = Vec::with_capacity(submitted.len());
|
||||
let mut prompt_tokens = BTreeMap::<usize, u32>::new();
|
||||
let mut completion_tokens = 0u64;
|
||||
|
||||
for choice in submitted {
|
||||
let output = collect_output(choice.events).await?;
|
||||
|
||||
prompt_tokens
|
||||
.entry(choice.prompt_index)
|
||||
.or_insert(output.prompt_tokens);
|
||||
completion_tokens = completion_tokens.saturating_add(output.completion_tokens);
|
||||
let response_choice = completion_choice(
|
||||
choice.index,
|
||||
if echo {
|
||||
choice.echo + &output.text
|
||||
} else {
|
||||
output.text.clone()
|
||||
},
|
||||
&output,
|
||||
want_logprobs,
|
||||
echo,
|
||||
);
|
||||
choices.push(response_choice);
|
||||
}
|
||||
|
||||
let prompt_tokens = prompt_tokens
|
||||
.values()
|
||||
.copied()
|
||||
.fold(0u32, u32::saturating_add);
|
||||
let usage = completion_usage(
|
||||
prompt_tokens,
|
||||
u32::try_from(completion_tokens).unwrap_or(u32::MAX),
|
||||
);
|
||||
|
||||
Ok(CompletionResponseWire {
|
||||
id: response_id,
|
||||
choices,
|
||||
created,
|
||||
model,
|
||||
object: "text_completion",
|
||||
usage: Some(usage),
|
||||
})
|
||||
}
|
||||
|
||||
fn completion_choice(
|
||||
index: usize,
|
||||
text: String,
|
||||
output: &GenerationOutput,
|
||||
want_logprobs: bool,
|
||||
include_input_logprobs: bool,
|
||||
) -> CompletionChoiceWire {
|
||||
let reason = output.finish_reason.as_ref();
|
||||
let finish_reason = match reason {
|
||||
Some(GenerationFinishReason::Stop(_)) => Some("stop".into()),
|
||||
Some(GenerationFinishReason::Length) => Some("length".into()),
|
||||
Some(GenerationFinishReason::ContentFilter) => Some("content_filter".into()),
|
||||
Some(GenerationFinishReason::Abort) => Some("abort".into()),
|
||||
Some(GenerationFinishReason::Other(other)) => Some(other.clone()),
|
||||
None => None,
|
||||
};
|
||||
let matched_stop = reason
|
||||
.and_then(|reason| match reason {
|
||||
GenerationFinishReason::Stop(matched) => matched.as_ref(),
|
||||
_ => None,
|
||||
})
|
||||
.map(|matched| match matched {
|
||||
MatchedStop::Token(id) => MatchedStopWire::Token(*id),
|
||||
MatchedStop::Text(value) => MatchedStopWire::Text(value.clone()),
|
||||
// Python's OpenAI schema supports an integer or string here, not a
|
||||
// multi-token list. Preserve the native value rather than dropping it.
|
||||
MatchedStop::Tokens(ids) => MatchedStopWire::Tokens(ids.clone()),
|
||||
});
|
||||
CompletionChoiceWire {
|
||||
text,
|
||||
index: u32::try_from(index).unwrap_or(u32::MAX),
|
||||
logprobs: want_logprobs
|
||||
.then(|| completion_logprobs(output.extras.as_deref(), include_input_logprobs)),
|
||||
finish_reason,
|
||||
matched_stop,
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) fn completion_event_stream(
|
||||
submitted: Vec<SubmittedChoice>,
|
||||
response_id: String,
|
||||
model: String,
|
||||
created: u32,
|
||||
echo: bool,
|
||||
want_logprobs: bool,
|
||||
include_usage: bool,
|
||||
continuous_usage: bool,
|
||||
) -> impl futures::Stream<Item = Result<CompletionResponseWire, ResponseError>> {
|
||||
async_stream::stream! {
|
||||
let count = submitted.len();
|
||||
let mut prompt_indexes = Vec::with_capacity(count);
|
||||
let mut echoes = Vec::with_capacity(count);
|
||||
let mut first_chunks = vec![true; count];
|
||||
let mut prompt_tokens_by_prompt = BTreeMap::<usize, u32>::new();
|
||||
let mut completion_tokens_by_choice = vec![0u64; count];
|
||||
let mut streams = Vec::with_capacity(count);
|
||||
|
||||
for choice in submitted {
|
||||
prompt_indexes.push(choice.prompt_index);
|
||||
echoes.push(choice.echo);
|
||||
streams.push(choice.events);
|
||||
}
|
||||
let mut events = merge_indexed(streams);
|
||||
|
||||
while let Some((index, item)) = events.next().await {
|
||||
let output = match item {
|
||||
Ok(output) => output,
|
||||
Err(error) => {
|
||||
yield Err(error);
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
prompt_tokens_by_prompt
|
||||
.entry(prompt_indexes[index])
|
||||
.or_insert(output.prompt_tokens);
|
||||
completion_tokens_by_choice[index] = completion_tokens_by_choice[index]
|
||||
.saturating_add(output.completion_tokens);
|
||||
let first = std::mem::replace(&mut first_chunks[index], false);
|
||||
let text = if echo && first {
|
||||
echoes[index].clone() + &output.text
|
||||
} else {
|
||||
output.text.clone()
|
||||
};
|
||||
let chunk_usage = continuous_usage.then(|| {
|
||||
completion_usage(
|
||||
output.prompt_tokens,
|
||||
u32::try_from(completion_tokens_by_choice[index]).unwrap_or(u32::MAX),
|
||||
)
|
||||
});
|
||||
let choice = completion_choice(
|
||||
index,
|
||||
text,
|
||||
&output,
|
||||
want_logprobs,
|
||||
echo && first,
|
||||
);
|
||||
let chunk = CompletionResponseWire {
|
||||
id: response_id.clone(),
|
||||
choices: vec![choice],
|
||||
created,
|
||||
model: model.clone(),
|
||||
object: "text_completion",
|
||||
usage: chunk_usage,
|
||||
};
|
||||
yield Ok(chunk);
|
||||
}
|
||||
|
||||
if include_usage {
|
||||
let prompt_tokens = prompt_tokens_by_prompt
|
||||
.values()
|
||||
.copied()
|
||||
.fold(0u32, u32::saturating_add);
|
||||
let completion_tokens = completion_tokens_by_choice
|
||||
.into_iter()
|
||||
.fold(0u64, u64::saturating_add);
|
||||
let final_chunk = CompletionResponseWire {
|
||||
id: response_id,
|
||||
choices: vec![],
|
||||
created,
|
||||
model,
|
||||
object: "text_completion",
|
||||
usage: Some(completion_usage(
|
||||
prompt_tokens,
|
||||
u32::try_from(completion_tokens).unwrap_or(u32::MAX),
|
||||
)),
|
||||
};
|
||||
yield Ok(final_chunk);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn completion_logprobs(
|
||||
extras: Option<&GenerationOutputExtras>,
|
||||
include_input: bool,
|
||||
) -> CompletionLogprobsWire {
|
||||
let mut result = CompletionLogprobsWire {
|
||||
tokens: Vec::new(),
|
||||
token_logprobs: Vec::new(),
|
||||
top_logprobs: Vec::new(),
|
||||
text_offset: Vec::new(),
|
||||
};
|
||||
let Some(extras) = extras else {
|
||||
return result;
|
||||
};
|
||||
if include_input {
|
||||
append_logprobs(&mut result, &extras.input_logprobs);
|
||||
}
|
||||
append_logprobs(&mut result, &extras.output_logprobs);
|
||||
result
|
||||
}
|
||||
|
||||
fn append_logprobs(result: &mut CompletionLogprobsWire, positions: &[crate::PositionLogprobs]) {
|
||||
for position in positions {
|
||||
let selected = &position.token;
|
||||
result.tokens.push(
|
||||
selected
|
||||
.text
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("token_id:{}", selected.token_id)),
|
||||
);
|
||||
// Python exposes the engine's f32 values as double-precision JSON numbers.
|
||||
result.token_logprobs.push(selected.logprob.map(f64::from));
|
||||
result.text_offset.push(-1);
|
||||
if position.top.is_empty() {
|
||||
result.top_logprobs.push(None);
|
||||
continue;
|
||||
}
|
||||
let mut top = BTreeMap::new();
|
||||
for candidate in &position.top {
|
||||
let Some(logprob) = candidate.logprob else {
|
||||
continue;
|
||||
};
|
||||
top.insert(
|
||||
candidate
|
||||
.text
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("token_id:{}", candidate.token_id)),
|
||||
f64::from(logprob),
|
||||
);
|
||||
}
|
||||
result.top_logprobs.push(Some(top));
|
||||
}
|
||||
}
|
||||
|
||||
impl super::OpenAIService {
|
||||
pub(crate) async fn complete(
|
||||
&self,
|
||||
request: CompletionRequest,
|
||||
) -> Result<
|
||||
super::OperationResponse<CompletionResponseWire, CompletionResponseWire>,
|
||||
ResponseError,
|
||||
> {
|
||||
use super::OperationResponse;
|
||||
let stream = request.stream.unwrap_or(false);
|
||||
let (response_id, requests) = prepare_request(&self.renderer, &request).await?;
|
||||
let context = prepare_response(
|
||||
&self.renderer,
|
||||
&self.generation.decoder,
|
||||
&request,
|
||||
response_id,
|
||||
requests.len(),
|
||||
)?;
|
||||
let streams = match self.generation.generate_many(requests).await {
|
||||
Ok(streams) => streams,
|
||||
Err(error) if stream => {
|
||||
return Ok(OperationResponse::Stream(
|
||||
futures::stream::once(async { Err(error) }).boxed(),
|
||||
));
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let submitted = attach_streams(context.metadata, streams);
|
||||
if stream {
|
||||
Ok(OperationResponse::Stream(
|
||||
completion_event_stream(
|
||||
submitted,
|
||||
context.response_id,
|
||||
context.model,
|
||||
context.created,
|
||||
context.echo,
|
||||
context.want_logprobs,
|
||||
context.include_usage,
|
||||
context.continuous_usage,
|
||||
)
|
||||
.boxed(),
|
||||
))
|
||||
} else {
|
||||
unary_completion(
|
||||
submitted,
|
||||
context.response_id,
|
||||
context.model,
|
||||
context.created,
|
||||
context.echo,
|
||||
context.want_logprobs,
|
||||
)
|
||||
.await
|
||||
.map(OperationResponse::Unary)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
completion_event_stream, completion_logprobs, prepare_request, prepare_response,
|
||||
unary_completion,
|
||||
};
|
||||
use crate::GenerationOutputExtras;
|
||||
use crate::engine::{TokenDecoder, test_utils::tiny_tokenizer};
|
||||
use crate::openai::test_utils::{chunk, renderer_config, submitted};
|
||||
use crate::{DynamoTokenizer, PositionLogprobs, RendererService, ResponseError, TokenLogprob};
|
||||
use futures::StreamExt;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[tokio::test]
|
||||
async fn completion_response_preserves_batched_echo_before_truncation() {
|
||||
let tokenizer = tiny_tokenizer();
|
||||
let prompts = ["hello", "world"];
|
||||
let token_ids =
|
||||
prompts.map(|prompt| tokenizer.encode(prompt).unwrap().token_ids().to_vec());
|
||||
for truncate in [false, true] {
|
||||
let mut config = renderer_config();
|
||||
if truncate {
|
||||
config.limits.context_len = 2;
|
||||
config.limits.allow_auto_truncate = true;
|
||||
assert!(token_ids.iter().all(|ids| ids.len() > 2));
|
||||
}
|
||||
let renderer = RendererService::with_tokenizer(
|
||||
config,
|
||||
Arc::new(DynamoTokenizer::new(tokenizer.clone(), tokenizer.clone())),
|
||||
1,
|
||||
1,
|
||||
);
|
||||
for tokenized in [false, true] {
|
||||
for echo in [false, true] {
|
||||
let prompt = if tokenized {
|
||||
serde_json::json!(token_ids)
|
||||
} else {
|
||||
serde_json::json!(prompts)
|
||||
};
|
||||
let request = serde_json::from_value(serde_json::json!({
|
||||
"model": "model", "prompt": prompt, "n": 2, "echo": echo,
|
||||
"rid": ["prompt-a", "prompt-b"], "max_tokens": 4, "logprobs": 0
|
||||
}))
|
||||
.unwrap();
|
||||
let (response_id, requests) =
|
||||
prepare_request(&renderer, &request).await.unwrap();
|
||||
let context = prepare_response(
|
||||
&renderer,
|
||||
&TokenDecoder::new(tokenizer.clone()),
|
||||
&request,
|
||||
response_id,
|
||||
requests.len(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(requests.len(), 4);
|
||||
assert_eq!(context.metadata.len(), 4);
|
||||
assert_eq!(context.echo, echo);
|
||||
for (index, (request, metadata)) in
|
||||
requests.iter().zip(&context.metadata).enumerate()
|
||||
{
|
||||
let prompt_index = index / 2;
|
||||
let expected_echo = if !echo {
|
||||
String::new()
|
||||
} else if tokenized {
|
||||
String::from(tokenizer.decode(&token_ids[prompt_index], true).unwrap())
|
||||
} else {
|
||||
prompts[prompt_index].to_owned()
|
||||
};
|
||||
assert_eq!(metadata, &(index, prompt_index, expected_echo));
|
||||
let mut expected_ids = token_ids[prompt_index]
|
||||
.iter()
|
||||
.map(|&id| id as i32)
|
||||
.collect::<Vec<_>>();
|
||||
if truncate {
|
||||
expected_ids.truncate(2);
|
||||
}
|
||||
assert_eq!(request.input_ids, expected_ids);
|
||||
assert_eq!(request.logprob_start_len, if echo { 0 } else { -1 });
|
||||
assert_eq!(
|
||||
request.rid,
|
||||
format!(
|
||||
"prompt-{}-{}",
|
||||
if prompt_index == 0 { "a" } else { "b" },
|
||||
index % 2
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serialized_logprobs_preserve_python_float_values() {
|
||||
let selected = -1.586831_f32;
|
||||
let alternative = -2.7182817_f32;
|
||||
let extras = GenerationOutputExtras {
|
||||
output_logprobs: vec![PositionLogprobs {
|
||||
token: TokenLogprob {
|
||||
logprob: Some(selected),
|
||||
token_id: 7,
|
||||
text: Some("x".into()),
|
||||
},
|
||||
top: vec![TokenLogprob {
|
||||
logprob: Some(alternative),
|
||||
token_id: 8,
|
||||
text: Some("y".into()),
|
||||
}],
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
// Exercise the wire serializer: to_value widens f32 before encoding it.
|
||||
let json = serde_json::to_string(&completion_logprobs(Some(&extras), false)).unwrap();
|
||||
let wire: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(
|
||||
wire["token_logprobs"][0].as_f64(),
|
||||
Some(f64::from(selected))
|
||||
);
|
||||
assert_eq!(
|
||||
wire["top_logprobs"][0]["y"].as_f64(),
|
||||
Some(f64::from(alternative))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn zero_top_logprobs_keeps_selected_token_and_empty_top_map() {
|
||||
let extras = GenerationOutputExtras {
|
||||
output_logprobs: vec![PositionLogprobs {
|
||||
token: TokenLogprob {
|
||||
logprob: Some(-0.25),
|
||||
token_id: 7,
|
||||
text: Some("x".into()),
|
||||
},
|
||||
top: Vec::new(),
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
let logprobs = completion_logprobs(Some(&extras), false);
|
||||
assert_eq!(logprobs.tokens, ["x"]);
|
||||
assert_eq!(logprobs.token_logprobs, [Some(-0.25)]);
|
||||
assert_eq!(logprobs.top_logprobs, [None]);
|
||||
assert_eq!(logprobs.text_offset, [-1]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unary_fold_orders_choices_and_counts_each_prompt_once() {
|
||||
let (choice0, tx0) = submitted(0, 0);
|
||||
let (choice1, tx1) = submitted(1, 0);
|
||||
tx0.send(chunk("a", false)).await.unwrap();
|
||||
tx0.send(chunk("b", true)).await.unwrap();
|
||||
tx1.send(chunk("x", false)).await.unwrap();
|
||||
tx1.send(chunk("y", true)).await.unwrap();
|
||||
|
||||
let response = unary_completion(
|
||||
vec![choice0, choice1],
|
||||
"cmpl-test".into(),
|
||||
"model".into(),
|
||||
1,
|
||||
false,
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
let value = serde_json::to_value(response.unwrap()).unwrap();
|
||||
assert_eq!(value["choices"][0]["text"], "ab");
|
||||
assert_eq!(value["choices"][1]["text"], "xy");
|
||||
assert_eq!(value["choices"][0]["matched_stop"], "</s>");
|
||||
assert!(value.get("system_fingerprint").is_none());
|
||||
assert_eq!(value["usage"]["prompt_tokens"], 5);
|
||||
assert_eq!(value["usage"]["completion_tokens"], 4);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_uses_deltas_then_usage() {
|
||||
let (choice, tx) = submitted(0, 0);
|
||||
tx.send(chunk("a", false)).await.unwrap();
|
||||
tx.send(chunk("b", true)).await.unwrap();
|
||||
|
||||
let stream = completion_event_stream(
|
||||
vec![choice],
|
||||
"cmpl-test".into(),
|
||||
"model".into(),
|
||||
1,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
let frames: Vec<_> = stream
|
||||
.map(|chunk| serde_json::to_value(chunk.unwrap()).unwrap())
|
||||
.collect()
|
||||
.await;
|
||||
assert_eq!(frames.len(), 3);
|
||||
let first = &frames[0];
|
||||
let terminal = &frames[1];
|
||||
let usage = &frames[2];
|
||||
assert_eq!(first["choices"][0]["text"], "a");
|
||||
assert_eq!(terminal["choices"][0]["text"], "b");
|
||||
assert_eq!(terminal["choices"][0]["finish_reason"], "stop");
|
||||
assert!(usage["choices"].as_array().unwrap().is_empty());
|
||||
assert_eq!(usage["usage"]["prompt_tokens"], 5);
|
||||
assert_eq!(usage["usage"]["completion_tokens"], 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_stops_all_choices_after_error() {
|
||||
let (choice0, tx0) = submitted(0, 0);
|
||||
let (choice1, tx1) = submitted(1, 0);
|
||||
let stream = completion_event_stream(
|
||||
vec![choice0, choice1],
|
||||
"cmpl-test".into(),
|
||||
"model".into(),
|
||||
1,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
futures::pin_mut!(stream);
|
||||
|
||||
tx0.send(Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Unavailable,
|
||||
message: "out of memory".into(),
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let error = stream.next().await.unwrap().unwrap_err();
|
||||
assert_eq!(error.kind, crate::ResponseErrorKind::Unavailable);
|
||||
|
||||
tx1.send(chunk("late", true)).await.unwrap();
|
||||
let remaining = stream.collect::<Vec<_>>().await;
|
||||
assert_eq!(remaining.len(), 1);
|
||||
assert!(
|
||||
remaining
|
||||
.into_iter()
|
||||
.all(|chunk| chunk.unwrap().choices.is_empty())
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
//! OpenAI request preparation and typed response construction.
|
||||
|
||||
use crate::ResponseError;
|
||||
use dynamo_protocols::types::CompletionUsage;
|
||||
|
||||
pub(crate) mod chat;
|
||||
pub(crate) mod completions;
|
||||
pub(crate) mod protocol;
|
||||
pub(crate) mod render;
|
||||
pub(crate) mod tokenize;
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) mod test_utils;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
pub(super) fn unix_seconds_u32() -> u32 {
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.map(|duration| u32::try_from(duration.as_secs()).unwrap_or(u32::MAX))
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
pub(super) fn completion_usage(prompt_tokens: u32, completion_tokens: u32) -> CompletionUsage {
|
||||
CompletionUsage {
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
total_tokens: prompt_tokens.saturating_add(completion_tokens),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Typed route result; transport adapters supply framing and status policy.
|
||||
pub(crate) enum OperationResponse<U, C> {
|
||||
Unary(U),
|
||||
Stream(futures::stream::BoxStream<'static, Result<C, ResponseError>>),
|
||||
}
|
||||
|
||||
pub(crate) struct OpenAIService {
|
||||
pub(crate) renderer: std::sync::Arc<crate::RendererService>,
|
||||
generation: crate::engine::GenerationService,
|
||||
}
|
||||
|
||||
impl OpenAIService {
|
||||
pub(crate) fn new(
|
||||
renderer: std::sync::Arc<crate::RendererService>,
|
||||
generation: crate::engine::GenerationService,
|
||||
) -> Self {
|
||||
Self {
|
||||
renderer,
|
||||
generation,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn error_payload(
|
||||
code: u16,
|
||||
message: impl Into<String>,
|
||||
error_type: &str,
|
||||
) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"error": {
|
||||
"object": "error", "message": message.into(), "type": error_type,
|
||||
"param": null, "code": code,
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,784 @@
|
||||
//! OpenAI wire types lowered into renderer-owned requests.
|
||||
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
|
||||
use dynamo_protocols::types::{
|
||||
ChatCompletionAudio, ChatCompletionFunctionCall, ChatCompletionFunctions,
|
||||
ChatCompletionRequestMessage, ChatCompletionStreamOptions, ChatCompletionTool,
|
||||
ChatCompletionToolChoiceOption, PredictionContent, Prompt, ResponseFormat, ServiceTier, Stop,
|
||||
WebSearchOptions,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::preprocessing::{GenerateRequestIdentity, TextRequestGroup};
|
||||
use crate::{
|
||||
ChatRequest, GenerateRequestMetadata, GenerationOptions, OneOrMany, ReasoningEffort,
|
||||
RendererConfig, RendererError, SamplingDefaults, SamplingParams, SamplingParamsOverrides,
|
||||
TokenIds, TokenIdsRequest,
|
||||
};
|
||||
|
||||
const MAX_OPENAI_CHOICES: usize = 4096;
|
||||
|
||||
#[derive(Deserialize, PartialEq)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
enum ResponseModality {
|
||||
Text,
|
||||
Audio,
|
||||
}
|
||||
|
||||
fn reject_unsupported_fields(fields: &HashMap<String, Value>) -> Result<(), String> {
|
||||
if fields.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let mut names = fields.keys().cloned().collect::<Vec<_>>();
|
||||
names.sort_unstable();
|
||||
Err(format!(
|
||||
"unsupported request field{}: {}",
|
||||
if names.len() == 1 { "" } else { "s" },
|
||||
names.join(", ")
|
||||
))
|
||||
}
|
||||
|
||||
/// SGLang's OpenAI-compatible chat-completions request.
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct ChatCompletionRequest {
|
||||
pub messages: Vec<ChatCompletionRequestMessage>,
|
||||
pub model: String,
|
||||
#[serde(default)]
|
||||
pub mm_processor_kwargs: Option<Value>,
|
||||
#[serde(default)]
|
||||
pub store: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub reasoning_effort: Option<ReasoningEffort>,
|
||||
#[serde(default)]
|
||||
pub reasoning: Option<Value>,
|
||||
#[serde(default)]
|
||||
pub metadata: Option<Value>,
|
||||
#[serde(default)]
|
||||
pub frequency_penalty: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub logit_bias: Option<HashMap<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub logprobs: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub top_logprobs: Option<u8>,
|
||||
#[serde(default)]
|
||||
pub max_tokens: Option<u32>,
|
||||
#[serde(default)]
|
||||
pub max_completion_tokens: Option<u32>,
|
||||
#[serde(default)]
|
||||
pub n: Option<u8>,
|
||||
#[serde(default)]
|
||||
modalities: Option<Vec<ResponseModality>>,
|
||||
#[serde(default)]
|
||||
pub prediction: Option<PredictionContent>,
|
||||
#[serde(default)]
|
||||
pub audio: Option<ChatCompletionAudio>,
|
||||
#[serde(default)]
|
||||
pub presence_penalty: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub response_format: Option<ResponseFormat>,
|
||||
#[serde(default)]
|
||||
pub seed: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub service_tier: Option<ServiceTier>,
|
||||
#[serde(default)]
|
||||
pub stop: Option<Stop>,
|
||||
#[serde(default)]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub stream_options: Option<ChatCompletionStreamOptions>,
|
||||
#[serde(default)]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub top_p: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub tools: Option<Vec<ChatCompletionTool>>,
|
||||
#[serde(default)]
|
||||
pub tool_choice: Option<ChatCompletionToolChoiceOption>,
|
||||
#[serde(default)]
|
||||
pub parallel_tool_calls: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub user: Option<String>,
|
||||
#[serde(default)]
|
||||
pub function_call: Option<ChatCompletionFunctionCall>,
|
||||
#[serde(default)]
|
||||
pub functions: Option<Vec<ChatCompletionFunctions>>,
|
||||
#[serde(default)]
|
||||
pub web_search_options: Option<WebSearchOptions>,
|
||||
#[serde(default)]
|
||||
pub chat_template_kwargs: Option<HashMap<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub continue_final_message: bool,
|
||||
#[serde(flatten)]
|
||||
pub sampling_overrides: SamplingParamsOverrides,
|
||||
#[serde(flatten)]
|
||||
pub extensions: RequestExtensions,
|
||||
#[serde(flatten)]
|
||||
pub unsupported_fields: HashMap<String, Value>,
|
||||
}
|
||||
|
||||
/// SGLang's OpenAI-compatible legacy-completions request.
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CompletionRequest {
|
||||
pub model: String,
|
||||
pub prompt: Prompt,
|
||||
#[serde(default)]
|
||||
pub prompt_embeds: Option<String>,
|
||||
#[serde(default)]
|
||||
pub suffix: Option<String>,
|
||||
#[serde(default)]
|
||||
pub max_tokens: Option<u32>,
|
||||
#[serde(default)]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub top_p: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub n: Option<u8>,
|
||||
#[serde(default)]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub stream_options: Option<ChatCompletionStreamOptions>,
|
||||
#[serde(default)]
|
||||
pub logprobs: Option<u8>,
|
||||
#[serde(default)]
|
||||
pub echo: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub stop: Option<Stop>,
|
||||
#[serde(default)]
|
||||
pub presence_penalty: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub frequency_penalty: Option<f32>,
|
||||
#[serde(default)]
|
||||
pub best_of: Option<u8>,
|
||||
#[serde(default)]
|
||||
pub logit_bias: Option<HashMap<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub user: Option<String>,
|
||||
#[serde(default)]
|
||||
pub seed: Option<i64>,
|
||||
#[serde(flatten)]
|
||||
pub sampling_overrides: SamplingParamsOverrides,
|
||||
#[serde(flatten)]
|
||||
pub extensions: RequestExtensions,
|
||||
#[serde(flatten)]
|
||||
pub unsupported_fields: HashMap<String, Value>,
|
||||
}
|
||||
|
||||
/// SGLang extensions carried by the OpenAI-compatible request contract.
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct RequestExtensions {
|
||||
#[serde(default)]
|
||||
pub return_meta_info: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub rid: Option<OneOrMany<String>>,
|
||||
#[serde(default)]
|
||||
pub cache_salt: Option<OneOrMany<String>>,
|
||||
#[serde(default)]
|
||||
pub extra_key: Option<OneOrMany<String>>,
|
||||
#[serde(default)]
|
||||
pub priority: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub bootstrap_host: Option<OneOrMany<String>>,
|
||||
#[serde(default)]
|
||||
pub bootstrap_port: Option<OneOrMany<Option<i64>>>,
|
||||
#[serde(default)]
|
||||
pub bootstrap_room: Option<OneOrMany<i64>>,
|
||||
#[serde(default)]
|
||||
pub routed_dp_rank: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub disagg_prefill_dp_rank: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub data_parallel_rank: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub session_id: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub session_params: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub lora_path: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub custom_logit_processor: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub image_data: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub video_data: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub audio_data: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub mm_hashes: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ExpandedRequestContext {
|
||||
request_id: String,
|
||||
metadata: GenerateRequestMetadata,
|
||||
}
|
||||
|
||||
impl RequestExtensions {
|
||||
fn validate(&self) -> Result<(), String> {
|
||||
for (name, value) in [
|
||||
("session_id", &self.session_id),
|
||||
("session_params", &self.session_params),
|
||||
("lora_path", &self.lora_path),
|
||||
("custom_logit_processor", &self.custom_logit_processor),
|
||||
("image_data", &self.image_data),
|
||||
("video_data", &self.video_data),
|
||||
("audio_data", &self.audio_data),
|
||||
("mm_hashes", &self.mm_hashes),
|
||||
] {
|
||||
if value.is_some() {
|
||||
return Err(format!(
|
||||
"{name} is not supported by the text-only Rust frontend"
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn response_id(&self, prefix: &str) -> String {
|
||||
match self.rid.as_ref() {
|
||||
Some(OneOrMany::One(rid)) => rid.clone(),
|
||||
Some(OneOrMany::Many(rids)) => rids
|
||||
.first()
|
||||
.cloned()
|
||||
.unwrap_or_else(|| generated_response_id(prefix)),
|
||||
None => generated_response_id(prefix),
|
||||
}
|
||||
}
|
||||
|
||||
fn expand(
|
||||
self,
|
||||
model: String,
|
||||
prompt_count: usize,
|
||||
choice_count: usize,
|
||||
response_id: &str,
|
||||
) -> Result<Vec<ExpandedRequestContext>, String> {
|
||||
let list_rids = matches!(&self.rid, Some(OneOrMany::Many(_)));
|
||||
let rids = expand_per_prompt("rid", self.rid, prompt_count)?;
|
||||
if list_rids {
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
for rid in rids.iter().flatten() {
|
||||
if !seen.insert(rid) {
|
||||
return Err(format!("duplicate request ID in rid: {rid}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
let cache_salts = expand_per_prompt("cache_salt", self.cache_salt, prompt_count)?;
|
||||
let extra_keys = expand_per_prompt("extra_key", self.extra_key, prompt_count)?;
|
||||
let bootstrap_hosts =
|
||||
expand_per_prompt("bootstrap_host", self.bootstrap_host, prompt_count)?;
|
||||
let bootstrap_ports =
|
||||
expand_per_prompt("bootstrap_port", self.bootstrap_port, prompt_count)?;
|
||||
let bootstrap_rooms = match self.bootstrap_room {
|
||||
Some(OneOrMany::One(base)) => (0..prompt_count)
|
||||
.map(|prompt_index| {
|
||||
let offset = i64::try_from(prompt_index)
|
||||
.map_err(|_| "bootstrap_room prompt index exceeds i64".to_owned())?;
|
||||
base.checked_add(offset)
|
||||
.map(Some)
|
||||
.ok_or_else(|| "bootstrap_room overflows i64".to_owned())
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
value => expand_per_prompt("bootstrap_room", value, prompt_count)?,
|
||||
};
|
||||
let routed_dp_rank = self.routed_dp_rank.or(self.data_parallel_rank);
|
||||
let total = prompt_count
|
||||
.checked_mul(choice_count)
|
||||
.ok_or_else(|| "prompt count times n overflows usize".to_owned())?;
|
||||
let mut contexts = Vec::with_capacity(total);
|
||||
for prompt_index in 0..prompt_count {
|
||||
for sample_index in 0..choice_count {
|
||||
let index = prompt_index * choice_count + sample_index;
|
||||
let request_id = match (&rids[prompt_index], list_rids) {
|
||||
(Some(rid), true) if choice_count == 1 => rid.clone(),
|
||||
(Some(rid), true) => format!("{rid}-{sample_index}"),
|
||||
_ => format!("{response_id}-{index}"),
|
||||
};
|
||||
contexts.push(ExpandedRequestContext {
|
||||
request_id,
|
||||
metadata: GenerateRequestMetadata {
|
||||
model: Some(model.clone()),
|
||||
cache_salt: cache_salts[prompt_index]
|
||||
.clone()
|
||||
.filter(|value| !value.is_empty()),
|
||||
extra_key: extra_keys[prompt_index]
|
||||
.clone()
|
||||
.filter(|value| !value.is_empty()),
|
||||
priority: self.priority,
|
||||
bootstrap_host: bootstrap_hosts[prompt_index].clone(),
|
||||
bootstrap_port: bootstrap_ports[prompt_index].flatten(),
|
||||
bootstrap_room: bootstrap_rooms[prompt_index],
|
||||
routed_dp_rank,
|
||||
disagg_prefill_dp_rank: self.disagg_prefill_dp_rank,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(contexts)
|
||||
}
|
||||
}
|
||||
|
||||
fn expand_per_prompt<T: Clone>(
|
||||
name: &str,
|
||||
value: Option<OneOrMany<T>>,
|
||||
prompt_count: usize,
|
||||
) -> Result<Vec<Option<T>>, String> {
|
||||
match value {
|
||||
None => Ok(vec![None; prompt_count]),
|
||||
Some(OneOrMany::One(value)) => Ok(vec![Some(value); prompt_count]),
|
||||
Some(OneOrMany::Many(values)) if values.len() == prompt_count => {
|
||||
Ok(values.into_iter().map(Some).collect())
|
||||
}
|
||||
Some(OneOrMany::Many(values)) => Err(format!(
|
||||
"the length of {name} must equal the prompt batch size ({prompt_count}), got {}",
|
||||
values.len()
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn generated_response_id(prefix: &str) -> String {
|
||||
format!("{prefix}-{}", uuid::Uuid::new_v4().simple())
|
||||
}
|
||||
|
||||
/// Lower the OpenAI Chat wire type into the structured internal chat request.
|
||||
/// Chat template rendering and tool constraints deliberately happen later in
|
||||
/// `ChatPreprocessor`, where every transport shares them.
|
||||
pub(crate) fn lower_chat_request(
|
||||
config: &RendererConfig,
|
||||
mut request: ChatCompletionRequest,
|
||||
) -> Result<(String, ChatRequest), RendererError> {
|
||||
normalize_reasoning_inputs(
|
||||
&mut request.reasoning_effort,
|
||||
request.reasoning.take(),
|
||||
&mut request.chat_template_kwargs,
|
||||
)?;
|
||||
// Accepted OpenAI metadata fields do not affect SGLang generation.
|
||||
let _ = (&request.store, &request.metadata, &request.user);
|
||||
reject_unsupported_fields(&request.unsupported_fields)?;
|
||||
request.extensions.validate()?;
|
||||
validate_chat_request(config, &request)?;
|
||||
let response_id = request.extensions.response_id("chatcmpl");
|
||||
let metadata = request
|
||||
.extensions
|
||||
.clone()
|
||||
.expand(request.model.clone(), 1, 1, &response_id)?
|
||||
.pop()
|
||||
.expect("one chat prompt produces one metadata context")
|
||||
.metadata;
|
||||
let mut sampling_params = chat_sampling_params(&request, &config.default_sampling_params)?;
|
||||
request.sampling_overrides.apply(&mut sampling_params);
|
||||
Ok((
|
||||
response_id.clone(),
|
||||
ChatRequest {
|
||||
rid: response_id,
|
||||
model: request.model,
|
||||
messages: request.messages,
|
||||
tools: request.tools,
|
||||
tool_choice: request.tool_choice,
|
||||
response_format: request.response_format,
|
||||
reasoning_effort: request.reasoning_effort,
|
||||
continue_final_message: request.continue_final_message,
|
||||
chat_template_args: request.chat_template_kwargs,
|
||||
sampling_params,
|
||||
choice_count: request.n.unwrap_or(1) as usize,
|
||||
stream: request.stream.unwrap_or(false),
|
||||
return_logprob: request.logprobs.unwrap_or(false),
|
||||
top_logprobs_num: request.top_logprobs.unwrap_or(0) as i64,
|
||||
parallel_tool_calls: request.parallel_tool_calls.unwrap_or(true),
|
||||
metadata,
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_reasoning_inputs(
|
||||
reasoning_effort: &mut Option<ReasoningEffort>,
|
||||
reasoning: Option<Value>,
|
||||
chat_template_kwargs: &mut Option<HashMap<String, Value>>,
|
||||
) -> Result<(), RendererError> {
|
||||
let mut thinking = None;
|
||||
if let Some(Value::Object(reasoning)) = reasoning {
|
||||
let nested_effort = reasoning
|
||||
.get("effort")
|
||||
.filter(|value| !value.is_null())
|
||||
.or_else(|| {
|
||||
reasoning
|
||||
.get("reasoning_effort")
|
||||
.filter(|value| !value.is_null())
|
||||
});
|
||||
if let Some(nested_effort) = nested_effort {
|
||||
*reasoning_effort = Some(
|
||||
serde_json::from_value(nested_effort.clone())
|
||||
.map_err(|error| format!("invalid reasoning effort: {error}"))?,
|
||||
);
|
||||
}
|
||||
|
||||
let enabled = reasoning
|
||||
.get("enabled")
|
||||
.filter(|value| !value.is_null())
|
||||
.or_else(|| reasoning.get("enable"));
|
||||
if enabled.is_some_and(json_truthy) {
|
||||
thinking = Some(true);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(effort) = reasoning_effort.as_ref() {
|
||||
thinking = Some(!effort.disables_thinking());
|
||||
}
|
||||
if let Some(thinking) = thinking {
|
||||
let args = chat_template_kwargs.get_or_insert_with(HashMap::new);
|
||||
args.entry("thinking".into()).or_insert(thinking.into());
|
||||
args.entry("enable_thinking".into())
|
||||
.or_insert(thinking.into());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn json_truthy(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::Null => false,
|
||||
Value::Bool(value) => *value,
|
||||
Value::Number(value) => value.as_f64().is_some_and(|value| value != 0.0),
|
||||
Value::String(value) => matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "y" | "on"
|
||||
),
|
||||
Value::Array(value) => !value.is_empty(),
|
||||
Value::Object(value) => !value.is_empty(),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_chat_request(
|
||||
config: &RendererConfig,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<(), RendererError> {
|
||||
if request.model != config.served_model_name {
|
||||
return Err(format!("The model `{}` does not exist", request.model).into());
|
||||
}
|
||||
if request.n == Some(0) {
|
||||
return Err("n must be at least 1".into());
|
||||
}
|
||||
if request.extensions.return_meta_info == Some(true) {
|
||||
return Err("return_meta_info=true is not supported by the renderer".into());
|
||||
}
|
||||
#[allow(deprecated)]
|
||||
let max_tokens = request.max_completion_tokens.or(request.max_tokens);
|
||||
if max_tokens == Some(0) {
|
||||
return Err("max_completion_tokens must be positive".into());
|
||||
}
|
||||
if request
|
||||
.modalities
|
||||
.as_ref()
|
||||
.is_some_and(|modalities| modalities.contains(&ResponseModality::Audio))
|
||||
|| request.audio.is_some()
|
||||
|| request.prediction.is_some()
|
||||
|| request.web_search_options.is_some()
|
||||
|| request.mm_processor_kwargs.is_some()
|
||||
{
|
||||
return Err(
|
||||
"audio, prediction, web search, and multimodal inputs are not supported".into(),
|
||||
);
|
||||
}
|
||||
#[allow(deprecated)]
|
||||
if request.function_call.is_some() || request.functions.is_some() {
|
||||
return Err(
|
||||
"deprecated function_call/functions are not supported; use tools and tool_choice"
|
||||
.into(),
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
pub fn chat_sampling_params(
|
||||
request: &ChatCompletionRequest,
|
||||
model_defaults: &SamplingDefaults,
|
||||
) -> Result<SamplingParams, String> {
|
||||
let defaults = sampling_params_with_model_defaults(model_defaults);
|
||||
let mut stop = None;
|
||||
let mut stop_token_ids = None;
|
||||
match request.stop.as_ref() {
|
||||
Some(Stop::String(value)) => stop = Some(OneOrMany::One(value.clone())),
|
||||
Some(Stop::StringArray(values)) => stop = Some(OneOrMany::Many(values.clone())),
|
||||
Some(Stop::TokenIdArray(values)) => {
|
||||
stop_token_ids = Some(values.iter().map(|&id| id as i64).collect())
|
||||
}
|
||||
None => {}
|
||||
}
|
||||
let mut logit_bias = BTreeMap::new();
|
||||
if let Some(values) = request.logit_bias.as_ref() {
|
||||
for (token, bias) in values {
|
||||
let bias = bias
|
||||
.as_f64()
|
||||
.ok_or_else(|| format!("logit_bias[{token:?}] must be a number"))?;
|
||||
logit_bias.insert(token.clone(), bias);
|
||||
}
|
||||
}
|
||||
let json_schema = match request.response_format.as_ref() {
|
||||
Some(ResponseFormat::JsonSchema { json_schema }) => Some(json_schema.schema.to_string()),
|
||||
Some(ResponseFormat::JsonObject) => Some(r#"{"type":"object"}"#.into()),
|
||||
_ => None,
|
||||
};
|
||||
|
||||
Ok(SamplingParams {
|
||||
max_new_tokens: request
|
||||
.max_completion_tokens
|
||||
.or(request.max_tokens)
|
||||
.map(i64::from),
|
||||
stop,
|
||||
stop_token_ids,
|
||||
temperature: request
|
||||
.temperature
|
||||
.map(f64::from)
|
||||
.unwrap_or(defaults.temperature),
|
||||
top_p: request.top_p.map(f64::from).unwrap_or(defaults.top_p),
|
||||
frequency_penalty: request.frequency_penalty.unwrap_or(0.0) as f64,
|
||||
presence_penalty: request.presence_penalty.unwrap_or(0.0) as f64,
|
||||
n: 1,
|
||||
logit_bias: (!logit_bias.is_empty()).then_some(logit_bias),
|
||||
sampling_seed: request.seed,
|
||||
json_schema,
|
||||
..defaults
|
||||
})
|
||||
}
|
||||
|
||||
fn sampling_params_with_model_defaults(model_defaults: &SamplingDefaults) -> SamplingParams {
|
||||
let terminals = SamplingParams::default();
|
||||
SamplingParams {
|
||||
temperature: model_defaults.temperature.unwrap_or(terminals.temperature),
|
||||
top_p: model_defaults.top_p.unwrap_or(terminals.top_p),
|
||||
top_k: model_defaults.top_k.unwrap_or(terminals.top_k),
|
||||
min_p: model_defaults.min_p.unwrap_or(terminals.min_p),
|
||||
repetition_penalty: model_defaults
|
||||
.repetition_penalty
|
||||
.unwrap_or(terminals.repetition_penalty),
|
||||
..terminals
|
||||
}
|
||||
}
|
||||
|
||||
/// Lower a textual OpenAI completion into text-only internal requests.
|
||||
pub(crate) fn lower_text_completion_request(
|
||||
config: &RendererConfig,
|
||||
request: &CompletionRequest,
|
||||
) -> Result<(String, Vec<TextRequestGroup>), RendererError> {
|
||||
// Accepted OpenAI request attribution does not affect generation.
|
||||
let _ = &request.user;
|
||||
reject_unsupported_fields(&request.unsupported_fields)?;
|
||||
request.extensions.validate()?;
|
||||
let prompts = text_completion_prompts(&request.prompt)?;
|
||||
let prompt_count = prompts.len();
|
||||
let (mut sampling, n, _) = completion_lowering_context(config, request, prompt_count)?;
|
||||
request.sampling_overrides.clone().apply(&mut sampling);
|
||||
let response_id = request.extensions.response_id("cmpl");
|
||||
let mut contexts = request
|
||||
.extensions
|
||||
.clone()
|
||||
.expand(request.model.clone(), prompt_count, n, &response_id)?
|
||||
.into_iter();
|
||||
let mut requests = Vec::with_capacity(prompt_count);
|
||||
for prompt in prompts {
|
||||
let mut choices = Vec::with_capacity(n);
|
||||
for _ in 0..n {
|
||||
let context = contexts
|
||||
.next()
|
||||
.expect("metadata expansion matches completion choice count");
|
||||
choices.push(GenerateRequestIdentity {
|
||||
rid: context.request_id,
|
||||
metadata: context.metadata,
|
||||
});
|
||||
}
|
||||
requests.push(TextRequestGroup {
|
||||
prompt: dynamo_renderer::RenderedPrompt::text(prompt),
|
||||
add_special_tokens: true,
|
||||
options: completion_generation_options(request, sampling.clone()),
|
||||
requests: choices,
|
||||
});
|
||||
}
|
||||
Ok((response_id, requests))
|
||||
}
|
||||
|
||||
/// Lower a pre-tokenized OpenAI completion directly into token-ID requests.
|
||||
pub(crate) fn lower_token_ids_completion_request(
|
||||
config: &RendererConfig,
|
||||
request: &CompletionRequest,
|
||||
) -> Result<(String, Vec<TokenIdsRequest>), RendererError> {
|
||||
// Accepted OpenAI request attribution does not affect generation.
|
||||
let _ = &request.user;
|
||||
reject_unsupported_fields(&request.unsupported_fields)?;
|
||||
request.extensions.validate()?;
|
||||
let prompts = token_ids_completion_prompts(&request.prompt)?;
|
||||
let prompt_count = prompts.len();
|
||||
let (mut sampling, n, choice_count) =
|
||||
completion_lowering_context(config, request, prompt_count)?;
|
||||
request.sampling_overrides.clone().apply(&mut sampling);
|
||||
let response_id = request.extensions.response_id("cmpl");
|
||||
let mut contexts = request
|
||||
.extensions
|
||||
.clone()
|
||||
.expand(request.model.clone(), prompt_count, n, &response_id)?
|
||||
.into_iter();
|
||||
let mut requests = Vec::with_capacity(choice_count);
|
||||
for input_ids in prompts {
|
||||
for _ in 0..n {
|
||||
let context = contexts
|
||||
.next()
|
||||
.expect("metadata expansion matches completion choice count");
|
||||
requests.push(
|
||||
TokenIdsRequest::new(
|
||||
context.request_id,
|
||||
input_ids.clone(),
|
||||
completion_generation_options(request, sampling.clone()),
|
||||
)
|
||||
.with_metadata(context.metadata),
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok((response_id, requests))
|
||||
}
|
||||
|
||||
fn completion_lowering_context(
|
||||
config: &RendererConfig,
|
||||
request: &CompletionRequest,
|
||||
prompt_count: usize,
|
||||
) -> Result<(SamplingParams, usize, usize), RendererError> {
|
||||
if request.model != config.served_model_name {
|
||||
return Err(format!("The model `{}` does not exist", request.model).into());
|
||||
}
|
||||
if request.prompt_embeds.is_some() {
|
||||
return Err("prompt_embeds is not supported by the Rust frontend".into());
|
||||
}
|
||||
if request.suffix.is_some() {
|
||||
return Err("suffix is not supported by this model".into());
|
||||
}
|
||||
if request.best_of.is_some_and(|best_of| best_of != 1) {
|
||||
return Err("best_of values greater than 1 are not supported".into());
|
||||
}
|
||||
if request.max_tokens == Some(0) {
|
||||
return Err("max_tokens must be positive".into());
|
||||
}
|
||||
if request.n == Some(0) {
|
||||
return Err("n must be at least 1".into());
|
||||
}
|
||||
let sampling = completion_sampling_params(request, &config.default_sampling_params)?;
|
||||
let n = request.n.unwrap_or(1) as usize;
|
||||
let choice_count = prompt_count
|
||||
.checked_mul(n)
|
||||
.filter(|&count| count <= MAX_OPENAI_CHOICES)
|
||||
.ok_or_else(|| {
|
||||
format!("prompt count times n exceeds the maximum of {MAX_OPENAI_CHOICES}")
|
||||
})?;
|
||||
Ok((sampling, n, choice_count))
|
||||
}
|
||||
|
||||
fn completion_generation_options(
|
||||
request: &CompletionRequest,
|
||||
sampling_params: SamplingParams,
|
||||
) -> GenerationOptions {
|
||||
GenerationOptions {
|
||||
sampling_params,
|
||||
stream: request.stream.unwrap_or(false),
|
||||
return_logprob: request.logprobs.is_some(),
|
||||
logprob_start_len: if request.echo.unwrap_or(false) && request.logprobs.is_some() {
|
||||
0
|
||||
} else {
|
||||
-1
|
||||
},
|
||||
top_logprobs_num: request.logprobs.unwrap_or(0) as i64,
|
||||
return_text_in_logprobs: request.logprobs.map(|_| true),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn text_completion_prompts(prompt: &Prompt) -> Result<Vec<String>, String> {
|
||||
match prompt {
|
||||
Prompt::String(text) => {
|
||||
if text.is_empty() {
|
||||
return Err("Prompt cannot be empty".into());
|
||||
}
|
||||
Ok(vec![text.clone()])
|
||||
}
|
||||
Prompt::StringArray(texts) => {
|
||||
if texts.is_empty() || texts.iter().any(String::is_empty) {
|
||||
return Err("Prompt cannot be empty".into());
|
||||
}
|
||||
Ok(texts.clone())
|
||||
}
|
||||
Prompt::IntegerArray(_) | Prompt::ArrayOfIntegerArray(_) => {
|
||||
Err("text completion lowerer requires a text prompt".into())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn token_ids_completion_prompts(prompt: &Prompt) -> Result<Vec<TokenIds>, String> {
|
||||
match prompt {
|
||||
Prompt::IntegerArray(ids) => Ok(vec![token_prompt_ids(ids)?]),
|
||||
Prompt::ArrayOfIntegerArray(prompts) => {
|
||||
if prompts.is_empty() {
|
||||
return Err("Prompt cannot be empty".into());
|
||||
}
|
||||
prompts.iter().map(|ids| token_prompt_ids(ids)).collect()
|
||||
}
|
||||
Prompt::String(_) | Prompt::StringArray(_) => {
|
||||
Err("token-ID completion lowerer requires a token-ID prompt".into())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn token_prompt_ids(ids: &[u32]) -> Result<TokenIds, String> {
|
||||
if ids.is_empty() {
|
||||
return Err("Prompt cannot be empty".into());
|
||||
}
|
||||
let input_ids = ids
|
||||
.iter()
|
||||
.map(|&id| i32::try_from(id).map_err(|_| format!("Token ID {id} is out of range")))
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
Ok(input_ids)
|
||||
}
|
||||
|
||||
pub fn completion_sampling_params(
|
||||
request: &CompletionRequest,
|
||||
model_defaults: &SamplingDefaults,
|
||||
) -> Result<SamplingParams, String> {
|
||||
let defaults = sampling_params_with_model_defaults(model_defaults);
|
||||
let mut stop = None;
|
||||
let mut stop_token_ids = None;
|
||||
match request.stop.as_ref() {
|
||||
Some(Stop::String(value)) => stop = Some(OneOrMany::One(value.clone())),
|
||||
Some(Stop::StringArray(values)) => stop = Some(OneOrMany::Many(values.clone())),
|
||||
Some(Stop::TokenIdArray(values)) => {
|
||||
stop_token_ids
|
||||
.get_or_insert_with(Vec::new)
|
||||
.extend(values.iter().map(|&id| id as i64));
|
||||
}
|
||||
None => {}
|
||||
}
|
||||
|
||||
let mut logit_bias = BTreeMap::new();
|
||||
if let Some(values) = request.logit_bias.as_ref() {
|
||||
for (token, bias) in values {
|
||||
let bias = bias
|
||||
.as_f64()
|
||||
.ok_or_else(|| format!("logit_bias[{token:?}] must be a number"))?;
|
||||
logit_bias.insert(token.clone(), bias);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(SamplingParams {
|
||||
max_new_tokens: Some(request.max_tokens.unwrap_or(16) as i64),
|
||||
stop,
|
||||
stop_token_ids,
|
||||
temperature: request
|
||||
.temperature
|
||||
.map(f64::from)
|
||||
.unwrap_or(defaults.temperature),
|
||||
top_p: request.top_p.map(f64::from).unwrap_or(defaults.top_p),
|
||||
frequency_penalty: request.frequency_penalty.unwrap_or(0.0) as f64,
|
||||
presence_penalty: request.presence_penalty.unwrap_or(0.0) as f64,
|
||||
// OpenAI `n` is implemented by fan-out: every native request has one
|
||||
// output, avoiding the native path's intentional `n > 1` rejection.
|
||||
n: 1,
|
||||
logit_bias: (!logit_bias.is_empty()).then_some(logit_bias),
|
||||
sampling_seed: request.seed,
|
||||
..defaults
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
//! OpenAI render-only operations, without model execution or HTTP framing.
|
||||
|
||||
use super::protocol::{ChatCompletionRequest, CompletionRequest};
|
||||
use crate::{GenerateRequest, RendererService, ResponseError};
|
||||
|
||||
pub(crate) async fn render_chat(
|
||||
renderer: &RendererService,
|
||||
request: ChatCompletionRequest,
|
||||
) -> Result<GenerateRequest, ResponseError> {
|
||||
if request.n.is_some_and(|n| n > 1) {
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::InvalidRequest,
|
||||
message: "the standalone chat renderer currently requires n=1".into(),
|
||||
});
|
||||
}
|
||||
let (_, mut chat) = super::chat::prepare_request(renderer, request).await?;
|
||||
Ok(chat
|
||||
.requests
|
||||
.pop()
|
||||
.expect("chat generation contains one request"))
|
||||
}
|
||||
|
||||
pub(crate) async fn render_completions(
|
||||
renderer: &RendererService,
|
||||
request: CompletionRequest,
|
||||
) -> Result<Vec<GenerateRequest>, ResponseError> {
|
||||
let (_, requests) = super::completions::prepare_request(renderer, &request).await?;
|
||||
Ok(requests)
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
use crate::{RendererConfig, RendererLimits, SamplingDefaults};
|
||||
use futures::StreamExt;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::{
|
||||
GenerationFinishReason, GenerationOutput, GenerationStream, MatchedStop, ResponseError,
|
||||
};
|
||||
|
||||
use super::completions::SubmittedChoice;
|
||||
|
||||
fn submission() -> (
|
||||
GenerationStream,
|
||||
mpsc::Sender<Result<GenerationOutput, ResponseError>>,
|
||||
) {
|
||||
let (tx, rx) = mpsc::channel::<Result<GenerationOutput, ResponseError>>(8);
|
||||
let events = futures::stream::unfold((rx, false), |(mut rx, finished)| async move {
|
||||
if finished {
|
||||
return None;
|
||||
}
|
||||
rx.recv().await.map(|item| {
|
||||
let finished = match &item {
|
||||
Ok(output) => output.finish_reason.is_some(),
|
||||
Err(_) => true,
|
||||
};
|
||||
(item, (rx, finished))
|
||||
})
|
||||
})
|
||||
.boxed();
|
||||
(events, tx)
|
||||
}
|
||||
|
||||
pub(super) fn chat_submitted(
|
||||
index: usize,
|
||||
) -> (
|
||||
(usize, GenerationStream),
|
||||
mpsc::Sender<Result<GenerationOutput, ResponseError>>,
|
||||
) {
|
||||
let (events, tx) = submission();
|
||||
((index, events), tx)
|
||||
}
|
||||
|
||||
pub(super) fn submitted(
|
||||
index: usize,
|
||||
prompt_index: usize,
|
||||
) -> (
|
||||
SubmittedChoice,
|
||||
mpsc::Sender<Result<GenerationOutput, ResponseError>>,
|
||||
) {
|
||||
let (events, tx) = submission();
|
||||
(
|
||||
SubmittedChoice {
|
||||
index,
|
||||
prompt_index,
|
||||
echo: String::new(),
|
||||
events,
|
||||
},
|
||||
tx,
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn chunk(text: &str, done: bool) -> Result<GenerationOutput, ResponseError> {
|
||||
let output = GenerationOutput {
|
||||
text: text.to_owned(),
|
||||
token_ids: vec![1],
|
||||
finish_reason: done
|
||||
.then(|| GenerationFinishReason::Stop(Some(MatchedStop::Text("</s>".into())))),
|
||||
prompt_tokens: 5,
|
||||
completion_tokens: 1,
|
||||
extras: None,
|
||||
};
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub(crate) fn renderer_config() -> RendererConfig {
|
||||
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: 128,
|
||||
context_len: 128,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,429 @@
|
||||
//! Protocol preparation invariants shared by rendering and inference.
|
||||
|
||||
use super::protocol::{
|
||||
ChatCompletionRequest, CompletionRequest, lower_chat_request, lower_text_completion_request,
|
||||
lower_token_ids_completion_request,
|
||||
};
|
||||
use super::test_utils::renderer_config;
|
||||
use crate::SamplingDefaults;
|
||||
|
||||
#[test]
|
||||
fn chat_lowering_preserves_template_controls_and_metadata() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"rid": "chat-lowering",
|
||||
"chat_template_kwargs": {"enable_thinking": false},
|
||||
"continue_final_message": true,
|
||||
"top_k": 17,
|
||||
"min_p": 0.2,
|
||||
"min_tokens": 3,
|
||||
"stop_regex": "END[0-9]",
|
||||
"ignore_eos": true,
|
||||
"skip_special_tokens": false,
|
||||
"return_meta_info": false,
|
||||
"bootstrap_host": "prefill",
|
||||
"bootstrap_port": 8998,
|
||||
"bootstrap_room": 42
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(request.model, "model");
|
||||
assert_eq!(
|
||||
request
|
||||
.chat_template_kwargs
|
||||
.as_ref()
|
||||
.and_then(|args| args.get("enable_thinking")),
|
||||
Some(&serde_json::Value::Bool(false))
|
||||
);
|
||||
assert!(request.continue_final_message);
|
||||
assert_eq!(request.sampling_overrides.top_k, Some(17));
|
||||
assert_eq!(request.sampling_overrides.min_p, Some(0.2));
|
||||
assert_eq!(request.sampling_overrides.min_tokens, Some(3));
|
||||
assert_eq!(request.sampling_overrides.ignore_eos, Some(true));
|
||||
assert_eq!(request.sampling_overrides.skip_special_tokens, Some(false));
|
||||
assert_eq!(request.extensions.return_meta_info, Some(false));
|
||||
|
||||
let (response_id, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "chat-lowering");
|
||||
assert_eq!(request.metadata.bootstrap_host.as_deref(), Some("prefill"));
|
||||
assert_eq!(request.metadata.bootstrap_port, Some(8998));
|
||||
assert_eq!(request.metadata.bootstrap_room, Some(42));
|
||||
assert_eq!(request.sampling_params.top_k, 17);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_lowering_rejects_return_meta_info_until_supported() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"return_meta_info": true
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = match lower_chat_request(&renderer_config(), request) {
|
||||
Ok(_) => panic!("return_meta_info=true must not be silently ignored"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.to_string().contains("return_meta_info"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_sampling_defaults_follow_request_model_terminal_priority() {
|
||||
let mut config = renderer_config();
|
||||
config.default_sampling_params = SamplingDefaults {
|
||||
temperature: Some(0.6),
|
||||
top_p: Some(0.9),
|
||||
top_k: Some(32),
|
||||
min_p: Some(0.1),
|
||||
repetition_penalty: Some(1.1),
|
||||
};
|
||||
let omitted: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello"
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&config, &omitted).unwrap();
|
||||
let sampling = &requests[0].options.sampling_params;
|
||||
assert_eq!(sampling.temperature, 0.6);
|
||||
assert_eq!(sampling.top_p, 0.9);
|
||||
assert_eq!(sampling.top_k, 32);
|
||||
assert_eq!(sampling.min_p, 0.1);
|
||||
assert_eq!(sampling.repetition_penalty, 1.1);
|
||||
|
||||
let explicit: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello",
|
||||
"temperature": 0.2,
|
||||
"top_p": 0.5,
|
||||
"top_k": 17,
|
||||
"min_p": 0.2,
|
||||
"repetition_penalty": 1.2
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&config, &explicit).unwrap();
|
||||
let sampling = &requests[0].options.sampling_params;
|
||||
assert!((sampling.temperature - 0.2).abs() < 1e-6);
|
||||
assert!((sampling.top_p - 0.5).abs() < 1e-6);
|
||||
assert_eq!(sampling.top_k, 17);
|
||||
assert_eq!(sampling.min_p, 0.2);
|
||||
assert_eq!(sampling.repetition_penalty, 1.2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unsupported_sglang_fields_are_rejected_instead_of_ignored() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"input_ids": [1, 2, 3],
|
||||
"task": "domain"
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = lower_chat_request(&renderer_config(), request)
|
||||
.unwrap_err()
|
||||
.to_string();
|
||||
|
||||
assert_eq!(error, "unsupported request fields: input_ids, task");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_modalities_keep_the_typed_openai_contract() {
|
||||
for modalities in [serde_json::json!("text"), serde_json::json!(["vision"])] {
|
||||
let request = serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"modalities": modalities
|
||||
});
|
||||
assert!(serde_json::from_value::<ChatCompletionRequest>(request).is_err());
|
||||
}
|
||||
|
||||
let text_request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"modalities": ["text"]
|
||||
}))
|
||||
.unwrap();
|
||||
lower_chat_request(&renderer_config(), text_request).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reasoning_inputs_normalize_with_python_precedence() {
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": "high",
|
||||
"reasoning": {"effort": "none", "enabled": true},
|
||||
"chat_template_kwargs": {"thinking": true}
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
let args = request.chat_template_args.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(request.reasoning_effort).unwrap(),
|
||||
serde_json::json!("none")
|
||||
);
|
||||
assert_eq!(args.get("thinking"), Some(&serde_json::json!(true)));
|
||||
assert_eq!(args.get("enable_thinking"), Some(&serde_json::json!(false)));
|
||||
|
||||
let request: ChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": "0.5"
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, request) = lower_chat_request(&renderer_config(), request).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(request.reasoning_effort).unwrap(),
|
||||
serde_json::json!(0.5)
|
||||
);
|
||||
assert_eq!(
|
||||
request
|
||||
.chat_template_args
|
||||
.as_ref()
|
||||
.and_then(|args| args.get("thinking")),
|
||||
Some(&serde_json::json!(true))
|
||||
);
|
||||
|
||||
for invalid in [serde_json::json!(true), serde_json::json!(1.0)] {
|
||||
let request = serde_json::json!({
|
||||
"model": "model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": invalid
|
||||
});
|
||||
assert!(serde_json::from_value::<ChatCompletionRequest>(request).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_completion_lowering_attaches_batched_metadata_in_prompt_major_order() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"n": 2,
|
||||
"rid": ["prompt-a", "prompt-b"],
|
||||
"cache_salt": ["tenant-a", "tenant-b"],
|
||||
"extra_key": ["", "batch"],
|
||||
"bootstrap_host": ["prefill-a", "prefill-b"],
|
||||
"bootstrap_port": [8998, null],
|
||||
"bootstrap_room": [41, 52],
|
||||
"priority": 7,
|
||||
"routed_dp_rank": 2
|
||||
}))
|
||||
.unwrap();
|
||||
let (response_id, requests) =
|
||||
lower_text_completion_request(&renderer_config(), &request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "prompt-a");
|
||||
assert_eq!(
|
||||
requests
|
||||
.iter()
|
||||
.flat_map(|request| request.requests.iter())
|
||||
.map(|request| request.rid.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["prompt-a-0", "prompt-a-1", "prompt-b-0", "prompt-b-1"]
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].requests[0].metadata.cache_salt.as_deref(),
|
||||
Some("tenant-a")
|
||||
);
|
||||
assert_eq!(requests[0].requests[1].metadata.extra_key, None);
|
||||
assert_eq!(
|
||||
requests[1].requests[0].metadata.extra_key.as_deref(),
|
||||
Some("batch")
|
||||
);
|
||||
assert_eq!(requests[0].requests[0].metadata.bootstrap_port, Some(8998));
|
||||
assert_eq!(requests[1].requests[0].metadata.bootstrap_port, None);
|
||||
assert_eq!(requests[0].requests[1].metadata.bootstrap_room, Some(41));
|
||||
assert_eq!(requests[1].requests[1].metadata.bootstrap_room, Some(52));
|
||||
assert_eq!(requests[1].requests[1].metadata.routed_dp_rank, Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_lowering_validates_metadata_lengths_duplicates_and_scalar_rooms() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"rid": ["duplicate", "duplicate"],
|
||||
"cache_salt": ["only-one"]
|
||||
}))
|
||||
.unwrap();
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
assert!(error.to_string().contains("duplicate request ID"));
|
||||
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"cache_salt": ["only-one"]
|
||||
}))
|
||||
.unwrap();
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
assert!(error.to_string().contains("prompt batch size (2)"));
|
||||
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": ["one", "two"],
|
||||
"n": 2,
|
||||
"bootstrap_room": 90
|
||||
}))
|
||||
.unwrap();
|
||||
let (_, requests) = lower_text_completion_request(&renderer_config(), &request).unwrap();
|
||||
assert_eq!(
|
||||
requests
|
||||
.iter()
|
||||
.flat_map(|request| request.requests.iter())
|
||||
.map(|request| request.metadata.bootstrap_room)
|
||||
.collect::<Vec<_>>(),
|
||||
[Some(90), Some(90), Some(91), Some(91)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_lowering_rejects_zero_max_tokens() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": "hello",
|
||||
"max_tokens": 0
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let error = lower_text_completion_request(&renderer_config(), &request).unwrap_err();
|
||||
|
||||
assert_eq!(error.to_string(), "max_tokens must be positive");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_id_completion_lowering_attaches_batched_metadata() {
|
||||
let request: CompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "model",
|
||||
"prompt": [[1, 2], [3]],
|
||||
"n": 2,
|
||||
"rid": ["tokens-a", "tokens-b"],
|
||||
"bootstrap_host": ["prefill-a", "prefill-b"],
|
||||
"bootstrap_port": [8998, 8999],
|
||||
"bootstrap_room": [41, 52]
|
||||
}))
|
||||
.unwrap();
|
||||
let (response_id, requests) =
|
||||
lower_token_ids_completion_request(&renderer_config(), &request).unwrap();
|
||||
|
||||
assert_eq!(response_id, "tokens-a");
|
||||
assert_eq!(requests[2].rid, "tokens-b-0");
|
||||
assert_eq!(requests[2].input_ids, [3]);
|
||||
assert_eq!(
|
||||
requests[2].metadata.bootstrap_host.as_deref(),
|
||||
Some("prefill-b")
|
||||
);
|
||||
assert_eq!(requests[2].metadata.bootstrap_port, Some(8999));
|
||||
assert_eq!(requests[3].metadata.bootstrap_room, Some(52));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_operations_decode_tokens_without_http() {
|
||||
use super::{OpenAIService, OperationResponse};
|
||||
use crate::engine::{
|
||||
GenerateTransport, GenerationService, TokenDecoder, TokenDelta, TokenStream,
|
||||
};
|
||||
use crate::{
|
||||
DynamoTokenizer, GenerateRequest, GenerationFinishReason, RendererService, ResponseError,
|
||||
};
|
||||
use futures::{StreamExt, future::BoxFuture};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
struct MemoryTransport(Mutex<Vec<GenerateRequest>>);
|
||||
impl GenerateTransport for MemoryTransport {
|
||||
fn generate(
|
||||
&self,
|
||||
request: GenerateRequest,
|
||||
) -> BoxFuture<'_, Result<TokenStream, ResponseError>> {
|
||||
Box::pin(async move {
|
||||
self.0.lock().unwrap().push(request);
|
||||
Ok(futures::stream::iter([Ok(TokenDelta {
|
||||
token_ids: vec![104],
|
||||
prompt_tokens: 5,
|
||||
completion_tokens: 1,
|
||||
finish_reason: Some(GenerationFinishReason::Length),
|
||||
..Default::default()
|
||||
})])
|
||||
.boxed())
|
||||
})
|
||||
}
|
||||
}
|
||||
async fn values<U: serde::Serialize, C: serde::Serialize>(
|
||||
result: OperationResponse<U, C>,
|
||||
) -> Vec<serde_json::Value> {
|
||||
match result {
|
||||
OperationResponse::Unary(value) => vec![serde_json::to_value(value).unwrap()],
|
||||
OperationResponse::Stream(stream) => {
|
||||
stream
|
||||
.map(|value| serde_json::to_value(value.unwrap()).unwrap())
|
||||
.collect()
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let tokenizer = crate::engine::test_utils::tiny_tokenizer();
|
||||
let prompt_ids = tokenizer.encode("hello").unwrap().token_ids().to_vec();
|
||||
let transport = Arc::new(MemoryTransport(Mutex::new(Vec::new())));
|
||||
let renderer = Arc::new(RendererService::with_tokenizer(
|
||||
renderer_config(),
|
||||
Arc::new(DynamoTokenizer::new(tokenizer.clone(), tokenizer.clone())),
|
||||
1,
|
||||
1,
|
||||
));
|
||||
let service = OpenAIService::new(
|
||||
renderer,
|
||||
GenerationService::new(transport.clone(), TokenDecoder::new(tokenizer)),
|
||||
);
|
||||
for chat in [false, true] {
|
||||
for stream in [false, true] {
|
||||
let mut body =
|
||||
serde_json::json!({"model": "model", "n": 2, "max_tokens": 4, "stream": stream});
|
||||
let responses = if chat {
|
||||
body["messages"] = serde_json::json!([{"role": "user", "content": "hello"}]);
|
||||
values(
|
||||
service
|
||||
.chat(serde_json::from_value(body).unwrap())
|
||||
.await
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
body["prompt"] = serde_json::json!(prompt_ids);
|
||||
body["echo"] = serde_json::json!(true);
|
||||
values(
|
||||
service
|
||||
.complete(serde_json::from_value(body).unwrap())
|
||||
.await
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
};
|
||||
let mut texts = [String::new(), String::new()];
|
||||
let mut finished = [false; 2];
|
||||
for response in responses {
|
||||
for choice in response["choices"].as_array().unwrap() {
|
||||
let index = choice["index"].as_u64().unwrap() as usize;
|
||||
let text = if chat {
|
||||
&choice[if stream { "delta" } else { "message" }]["content"]
|
||||
} else {
|
||||
&choice["text"]
|
||||
};
|
||||
texts[index].push_str(text.as_str().unwrap_or_default());
|
||||
if let Some(reason) = choice["finish_reason"].as_str() {
|
||||
assert_eq!(reason, "length");
|
||||
finished[index] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
assert_eq!(texts, [if chat { "h" } else { "helloh" }; 2]);
|
||||
assert_eq!(finished, [true; 2]);
|
||||
}
|
||||
}
|
||||
let requests = transport.0.lock().unwrap();
|
||||
assert_eq!(requests.len(), 8);
|
||||
assert!(requests.iter().all(|request| !request.input_ids.is_empty()));
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
//! SGLang-compatible prompt and chat tokenization.
|
||||
|
||||
use dynamo_protocols::types::{
|
||||
ChatCompletionRequestMessage, ChatCompletionTool, ChatCompletionToolChoiceOption,
|
||||
};
|
||||
use futures::future::try_join_all;
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{ChatRequest, OneOrMany, ReasoningEffort, RendererService, ResponseError};
|
||||
|
||||
use super::protocol::normalize_reasoning_inputs;
|
||||
|
||||
pub(crate) async fn tokenize(
|
||||
renderer: &RendererService,
|
||||
mut request: TokenizeRequest,
|
||||
) -> Result<Value, ResponseError> {
|
||||
let has_prompt = request.prompt.is_some();
|
||||
let has_messages = request.messages.is_some();
|
||||
if has_prompt == has_messages {
|
||||
return Err(ResponseError {
|
||||
kind: crate::ResponseErrorKind::InvalidRequest,
|
||||
message: "Exactly one of 'prompt' or 'messages' must be provided.".into(),
|
||||
});
|
||||
}
|
||||
let (tokens, count) = match request.prompt.take() {
|
||||
Some(prompt) => {
|
||||
let add_special_tokens = request.add_special_tokens;
|
||||
match prompt {
|
||||
OneOrMany::One(text) => {
|
||||
let tokens = renderer.tokenize_prompt(text, add_special_tokens).await?;
|
||||
(json!(tokens), json!(tokens.len()))
|
||||
}
|
||||
OneOrMany::Many(texts) => {
|
||||
let tokens = try_join_all(
|
||||
texts
|
||||
.into_iter()
|
||||
.map(|text| renderer.tokenize_prompt(text, add_special_tokens)),
|
||||
)
|
||||
.await?;
|
||||
let count = tokens.iter().map(Vec::len).collect::<Vec<_>>();
|
||||
(json!(tokens), json!(count))
|
||||
}
|
||||
}
|
||||
}
|
||||
None => {
|
||||
let request = request.into_chat(&renderer.config().served_model_name)?;
|
||||
let tokens = renderer.tokenize_chat(request).await?;
|
||||
(json!(tokens), json!(tokens.len()))
|
||||
}
|
||||
};
|
||||
Ok(json!({
|
||||
"tokens": tokens,
|
||||
"count": count,
|
||||
"max_model_len": renderer.config().limits.context_len,
|
||||
}))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct TokenizeRequest {
|
||||
#[serde(default)]
|
||||
prompt: Option<OneOrMany<String>>,
|
||||
#[serde(default)]
|
||||
messages: Option<Vec<ChatCompletionRequestMessage>>,
|
||||
#[serde(default = "default_true")]
|
||||
add_special_tokens: bool,
|
||||
#[serde(default)]
|
||||
model: Option<String>,
|
||||
#[serde(default)]
|
||||
tools: Option<Vec<ChatCompletionTool>>,
|
||||
#[serde(default)]
|
||||
tool_choice: Option<ChatCompletionToolChoiceOption>,
|
||||
#[serde(default)]
|
||||
reasoning_effort: Option<ReasoningEffort>,
|
||||
#[serde(default)]
|
||||
reasoning: Option<Value>,
|
||||
#[serde(default)]
|
||||
continue_final_message: bool,
|
||||
#[serde(default)]
|
||||
chat_template_kwargs: Option<std::collections::HashMap<String, Value>>,
|
||||
}
|
||||
|
||||
impl TokenizeRequest {
|
||||
fn into_chat(mut self, served_model: &str) -> Result<ChatRequest, crate::RendererError> {
|
||||
normalize_reasoning_inputs(
|
||||
&mut self.reasoning_effort,
|
||||
self.reasoning.take(),
|
||||
&mut self.chat_template_kwargs,
|
||||
)?;
|
||||
let model = self.model.unwrap_or_else(|| served_model.to_owned());
|
||||
if model != served_model {
|
||||
return Err(format!("The model `{model}` does not exist").into());
|
||||
}
|
||||
Ok(ChatRequest {
|
||||
rid: "tokenize".into(),
|
||||
model,
|
||||
messages: self
|
||||
.messages
|
||||
.take()
|
||||
.expect("chat tokenization request has messages"),
|
||||
tools: self.tools,
|
||||
tool_choice: self.tool_choice,
|
||||
response_format: None,
|
||||
reasoning_effort: self.reasoning_effort,
|
||||
continue_final_message: self.continue_final_message,
|
||||
chat_template_args: self.chat_template_kwargs,
|
||||
sampling_params: Default::default(),
|
||||
choice_count: 1,
|
||||
stream: false,
|
||||
return_logprob: false,
|
||||
top_logprobs_num: 0,
|
||||
parallel_tool_calls: true,
|
||||
metadata: crate::GenerateRequestMetadata::default(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
const fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn chat_tokenization_lowers_tokenize_specific_options() {
|
||||
let request: TokenizeRequest = serde_json::from_value(json!({
|
||||
"messages": [{"role": "assistant", "content": "partial"}],
|
||||
"reasoning_effort": "high",
|
||||
"continue_final_message": true,
|
||||
"chat_template_kwargs": {"marker": true}
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let chat = request.into_chat("model").unwrap();
|
||||
|
||||
assert!(chat.continue_final_message);
|
||||
assert_eq!(
|
||||
chat.chat_template_args
|
||||
.as_ref()
|
||||
.and_then(|args| args.get("marker")),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(chat.reasoning_effort).unwrap(),
|
||||
json!("high")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,774 @@
|
||||
//! Request-scoped OpenAI chat output interpretation.
|
||||
//!
|
||||
//! The processor owns parser selection and mutable reasoning/tool state. Its
|
||||
//! input is decoded engine output; its output is typed chat semantics.
|
||||
//! Submission, cancellation, and scheduler transport remain host
|
||||
//! responsibilities. HTTP and future gRPC adapters consume these semantic
|
||||
//! events without reimplementing parser behavior.
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use dynamo_parsers::ToolDefinition;
|
||||
use dynamo_parsers::reasoning::{
|
||||
ReasoningParser as _, ReasoningParserType, ReasoningParserWrapper,
|
||||
};
|
||||
use dynamo_parsers::tool_calling::jail::{Annotated, apply_tool_calling_jail};
|
||||
use dynamo_protocols::types::{
|
||||
ChatChoiceLogprobs, ChatChoiceStream, ChatCompletionMessageContent,
|
||||
ChatCompletionMessageToolCallChunk, ChatCompletionStreamResponseDelta,
|
||||
ChatCompletionToolChoiceOption, CreateChatCompletionStreamResponse, FinishReason, Role,
|
||||
};
|
||||
use futures::{Stream, StreamExt};
|
||||
|
||||
use crate::ResponseError;
|
||||
use crate::preprocessing::dynamo_parser_name;
|
||||
|
||||
/// Engine-neutral terminal reason understood by chat response processing.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ChatFinishReason {
|
||||
Stop,
|
||||
Length,
|
||||
ContentFilter,
|
||||
ToolCalls,
|
||||
}
|
||||
|
||||
/// One decoded engine update after host-specific egress conversion.
|
||||
pub struct DecodedChatEvent {
|
||||
pub choice: usize,
|
||||
pub text: String,
|
||||
pub token_ids: Vec<i32>,
|
||||
pub finish_reason: Option<ChatFinishReason>,
|
||||
pub logprobs: Option<ChatChoiceLogprobs>,
|
||||
pub prompt_tokens: u32,
|
||||
pub completion_tokens: u64,
|
||||
}
|
||||
|
||||
/// One semantic tool-call delta, independent of HTTP or gRPC framing.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ChatToolCallDelta {
|
||||
pub index: u32,
|
||||
pub id: Option<String>,
|
||||
pub name: Option<String>,
|
||||
pub arguments: Option<String>,
|
||||
}
|
||||
|
||||
/// Semantic chat output. Protocol adapters add response metadata and wire
|
||||
/// framing without knowing how reasoning or tool syntax was parsed.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum ChatEvent {
|
||||
Role {
|
||||
choice: usize,
|
||||
},
|
||||
Delta {
|
||||
choice: usize,
|
||||
content: Option<String>,
|
||||
reasoning_content: Option<String>,
|
||||
tool_calls: Option<Vec<ChatToolCallDelta>>,
|
||||
finish_reason: Option<ChatFinishReason>,
|
||||
logprobs: Option<ChatChoiceLogprobs>,
|
||||
},
|
||||
Usage {
|
||||
prompt_tokens: u32,
|
||||
completion_tokens: u64,
|
||||
},
|
||||
}
|
||||
|
||||
/// Mutable parser state for one generated choice.
|
||||
struct ChoiceResponseProcessor {
|
||||
reasoning: ReasoningStreamSplitter,
|
||||
}
|
||||
|
||||
/// Request-scoped chat response processor.
|
||||
///
|
||||
/// Parser names, tool definitions, structural-tag decisions, and mutable
|
||||
/// per-choice state are private so protocol adapters cannot accidentally
|
||||
/// reimplement the semantic contract.
|
||||
pub struct ChatResponseProcessor {
|
||||
tool_parser: Option<String>,
|
||||
tools: Option<Vec<ToolDefinition>>,
|
||||
tool_choice: Option<ChatCompletionToolChoiceOption>,
|
||||
uses_tool_call_structural_tag: bool,
|
||||
parallel_tool_calls: bool,
|
||||
choices: Vec<ChoiceResponseProcessor>,
|
||||
}
|
||||
|
||||
impl ChatResponseProcessor {
|
||||
pub(crate) fn new(
|
||||
tool_parser: Option<String>,
|
||||
reasoning_parser: Option<String>,
|
||||
tools: Option<Vec<ToolDefinition>>,
|
||||
tool_choice: Option<ChatCompletionToolChoiceOption>,
|
||||
uses_tool_call_structural_tag: bool,
|
||||
parallel_tool_calls: bool,
|
||||
choice_count: usize,
|
||||
) -> Self {
|
||||
Self {
|
||||
tool_parser,
|
||||
tools,
|
||||
tool_choice,
|
||||
uses_tool_call_structural_tag,
|
||||
parallel_tool_calls,
|
||||
choices: (0..choice_count)
|
||||
.map(|_| ChoiceResponseProcessor {
|
||||
reasoning: ReasoningStreamSplitter::new(reasoning_parser.as_deref(), None),
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_reasoning_state(mut self, reasoning_state: Option<bool>) -> Self {
|
||||
for choice in &mut self.choices {
|
||||
choice.reasoning.initial_reasoning = reasoning_state;
|
||||
}
|
||||
self
|
||||
}
|
||||
|
||||
/// Interpret decoded output and emit semantic chat events.
|
||||
///
|
||||
/// OpenAI-shaped values are used only as a private adapter to Dynamo's
|
||||
/// stateful tool-call jail. They are removed before events leave this
|
||||
/// crate, so response identity, model metadata, usage policy, and wire
|
||||
/// framing remain outside this semantic processor.
|
||||
pub fn process_stream<S>(
|
||||
mut self,
|
||||
input: S,
|
||||
) -> Pin<Box<dyn Stream<Item = Result<ChatEvent, ResponseError>> + Send>>
|
||||
where
|
||||
S: Stream<Item = Result<DecodedChatEvent, ResponseError>> + Send + 'static,
|
||||
{
|
||||
let count = self.choices.len();
|
||||
let raw = async_stream::stream! {
|
||||
let mut prompt_tokens = 0u32;
|
||||
let mut completion_tokens = 0u64;
|
||||
let mut role_emitted = vec![false; count];
|
||||
|
||||
futures::pin_mut!(input);
|
||||
while let Some(item) = input.next().await {
|
||||
let decoded = match item {
|
||||
Ok(decoded) => decoded,
|
||||
Err(error) => {
|
||||
yield Annotated {
|
||||
data: None,
|
||||
id: None,
|
||||
event: None,
|
||||
comment: None,
|
||||
error: serde_json::to_string(&error).ok(),
|
||||
};
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if prompt_tokens == 0 {
|
||||
prompt_tokens = decoded.prompt_tokens;
|
||||
}
|
||||
completion_tokens = completion_tokens.saturating_add(decoded.completion_tokens);
|
||||
|
||||
if decoded.choice >= count {
|
||||
yield Annotated {
|
||||
data: None,
|
||||
id: None,
|
||||
event: None,
|
||||
comment: None,
|
||||
error: serde_json::to_string(&ResponseError {
|
||||
kind: crate::ResponseErrorKind::Internal,
|
||||
message: format!("output choice {} is out of range", decoded.choice),
|
||||
}).ok(),
|
||||
};
|
||||
continue;
|
||||
}
|
||||
|
||||
if !role_emitted[decoded.choice] {
|
||||
role_emitted[decoded.choice] = true;
|
||||
yield annotated_choices(vec![ChatChoiceStream {
|
||||
index: decoded.choice as u32,
|
||||
delta: chat_delta(None, Some(Role::Assistant), None, None),
|
||||
finish_reason: None,
|
||||
logprobs: None,
|
||||
}]);
|
||||
}
|
||||
|
||||
let choice = &mut self.choices[decoded.choice];
|
||||
let index = decoded.choice as u32;
|
||||
let (reasoning_text, normal_text) =
|
||||
choice.reasoning.split(&decoded.text, &decoded.token_ids);
|
||||
let mut remaining_logprobs = decoded.logprobs;
|
||||
let mut emitted = Vec::with_capacity(3);
|
||||
if !reasoning_text.is_empty() {
|
||||
emitted.push(ChatChoiceStream {
|
||||
index,
|
||||
delta: chat_delta(None, None, None, Some(reasoning_text)),
|
||||
finish_reason: None,
|
||||
logprobs: remaining_logprobs.take(),
|
||||
});
|
||||
}
|
||||
if !normal_text.is_empty() {
|
||||
emitted.push(ChatChoiceStream {
|
||||
index,
|
||||
delta: chat_delta(Some(normal_text), None, None, None),
|
||||
finish_reason: None,
|
||||
logprobs: remaining_logprobs.take(),
|
||||
});
|
||||
}
|
||||
|
||||
if decoded.finish_reason.is_some() {
|
||||
let (reasoning_tail, normal_tail) = choice.reasoning.finish();
|
||||
if !reasoning_tail.is_empty() {
|
||||
emitted.push(ChatChoiceStream {
|
||||
index,
|
||||
delta: chat_delta(None, None, None, Some(reasoning_tail)),
|
||||
finish_reason: None,
|
||||
logprobs: None,
|
||||
});
|
||||
}
|
||||
if !normal_tail.is_empty() {
|
||||
emitted.push(ChatChoiceStream {
|
||||
index,
|
||||
delta: chat_delta(Some(normal_tail), None, None, None),
|
||||
finish_reason: None,
|
||||
logprobs: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let finish_reason = decoded.finish_reason.map(to_dynamo_finish_reason);
|
||||
match emitted.last_mut() {
|
||||
Some(last) => last.finish_reason = finish_reason,
|
||||
None => emitted.push(ChatChoiceStream {
|
||||
index,
|
||||
delta: chat_delta(None, None, None, None),
|
||||
finish_reason,
|
||||
logprobs: remaining_logprobs,
|
||||
}),
|
||||
}
|
||||
yield annotated_choices(emitted);
|
||||
}
|
||||
|
||||
yield annotated_usage(prompt_tokens, completion_tokens);
|
||||
};
|
||||
|
||||
let post_tool_terminal_markers = self.tool_parser.as_deref().map_or(&[][..], |parser| {
|
||||
match dynamo_parser_name(parser) {
|
||||
"qwen25" => &["<|im_end|>"],
|
||||
"glm47" => &["<|user|>", "<|endoftext|>", "<|observation|>"],
|
||||
_ => &[],
|
||||
}
|
||||
});
|
||||
let parsed: Pin<
|
||||
Box<dyn Stream<Item = Annotated<CreateChatCompletionStreamResponse>> + Send>,
|
||||
> = if let Some(parser) = self.tool_parser {
|
||||
Box::pin(apply_tool_calling_jail(
|
||||
Some(dynamo_parser_name(&parser).to_owned()),
|
||||
self.tool_choice,
|
||||
self.tools,
|
||||
self.uses_tool_call_structural_tag,
|
||||
raw,
|
||||
))
|
||||
} else {
|
||||
Box::pin(raw)
|
||||
};
|
||||
let parallel_tool_calls = self.parallel_tool_calls;
|
||||
|
||||
Box::pin(async_stream::stream! {
|
||||
let mut tool_calls_seen = vec![false; count];
|
||||
futures::pin_mut!(parsed);
|
||||
while let Some(mut item) = parsed.next().await {
|
||||
if let Some(response) = item.data.take() {
|
||||
if response.choices.is_empty()
|
||||
&& let Some(usage) = response.usage
|
||||
{
|
||||
yield Ok(ChatEvent::Usage {
|
||||
prompt_tokens: usage.prompt_tokens,
|
||||
completion_tokens: u64::from(usage.completion_tokens),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
for choice in response.choices {
|
||||
let index = choice.index as usize;
|
||||
let had_tool_calls = tool_calls_seen.get(index).copied().unwrap_or(false);
|
||||
let mut tool_calls = choice.delta.tool_calls.map(|calls| {
|
||||
calls.into_iter().map(tool_call_delta).collect::<Vec<_>>()
|
||||
});
|
||||
if !parallel_tool_calls
|
||||
&& let Some(calls) = tool_calls.as_mut()
|
||||
{
|
||||
if had_tool_calls {
|
||||
calls.clear();
|
||||
} else {
|
||||
calls.truncate(1);
|
||||
}
|
||||
if calls.is_empty() {
|
||||
tool_calls = None;
|
||||
}
|
||||
}
|
||||
let emitted_tool_calls = tool_calls.as_ref().is_some_and(|calls| !calls.is_empty());
|
||||
if emitted_tool_calls
|
||||
&& let Some(seen) = tool_calls_seen.get_mut(index)
|
||||
{
|
||||
*seen = true;
|
||||
}
|
||||
let mut content = match choice.delta.content {
|
||||
Some(ChatCompletionMessageContent::Text(text)) => Some(text),
|
||||
_ => None,
|
||||
};
|
||||
if had_tool_calls
|
||||
&& content.as_ref().is_some_and(|text| {
|
||||
post_tool_terminal_markers.contains(&text.trim())
|
||||
})
|
||||
{
|
||||
content = None;
|
||||
}
|
||||
if choice.delta.role.is_some()
|
||||
&& content.is_none()
|
||||
&& choice.delta.reasoning_content.is_none()
|
||||
&& tool_calls.is_none()
|
||||
&& choice.finish_reason.is_none()
|
||||
{
|
||||
yield Ok(ChatEvent::Role { choice: index });
|
||||
continue;
|
||||
}
|
||||
yield Ok(ChatEvent::Delta {
|
||||
choice: index,
|
||||
content,
|
||||
reasoning_content: choice.delta.reasoning_content,
|
||||
tool_calls,
|
||||
finish_reason: choice.finish_reason.map(from_dynamo_finish_reason),
|
||||
logprobs: choice.logprobs,
|
||||
});
|
||||
}
|
||||
} else if let Some(error) = item.error {
|
||||
let error = serde_json::from_str(&error).unwrap_or(ResponseError {
|
||||
kind: crate::ResponseErrorKind::Internal,
|
||||
message: error,
|
||||
});
|
||||
yield Err(error);
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
fn chat_delta(
|
||||
content: Option<String>,
|
||||
role: Option<Role>,
|
||||
tool_calls: Option<Vec<ChatCompletionMessageToolCallChunk>>,
|
||||
reasoning_content: Option<String>,
|
||||
) -> ChatCompletionStreamResponseDelta {
|
||||
ChatCompletionStreamResponseDelta {
|
||||
content: content.map(ChatCompletionMessageContent::Text),
|
||||
function_call: None,
|
||||
tool_calls,
|
||||
role,
|
||||
refusal: None,
|
||||
reasoning_content,
|
||||
}
|
||||
}
|
||||
|
||||
fn annotated_choices(
|
||||
choices: Vec<ChatChoiceStream>,
|
||||
) -> Annotated<CreateChatCompletionStreamResponse> {
|
||||
Annotated {
|
||||
data: Some(CreateChatCompletionStreamResponse {
|
||||
id: String::new(),
|
||||
choices,
|
||||
created: 0,
|
||||
model: String::new(),
|
||||
service_tier: None,
|
||||
system_fingerprint: None,
|
||||
object: String::new(),
|
||||
usage: None,
|
||||
}),
|
||||
id: None,
|
||||
event: None,
|
||||
comment: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn annotated_usage(
|
||||
prompt_tokens: u32,
|
||||
completion_tokens: u64,
|
||||
) -> Annotated<CreateChatCompletionStreamResponse> {
|
||||
Annotated {
|
||||
data: Some(CreateChatCompletionStreamResponse {
|
||||
id: String::new(),
|
||||
choices: Vec::new(),
|
||||
created: 0,
|
||||
model: String::new(),
|
||||
service_tier: None,
|
||||
system_fingerprint: None,
|
||||
object: String::new(),
|
||||
usage: Some(dynamo_protocols::types::CompletionUsage {
|
||||
prompt_tokens,
|
||||
completion_tokens: u32::try_from(completion_tokens).unwrap_or(u32::MAX),
|
||||
total_tokens: prompt_tokens
|
||||
.saturating_add(u32::try_from(completion_tokens).unwrap_or(u32::MAX)),
|
||||
prompt_tokens_details: None,
|
||||
completion_tokens_details: None,
|
||||
}),
|
||||
}),
|
||||
id: None,
|
||||
event: None,
|
||||
comment: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_call_delta(call: ChatCompletionMessageToolCallChunk) -> ChatToolCallDelta {
|
||||
ChatToolCallDelta {
|
||||
index: call.index,
|
||||
id: call.id,
|
||||
name: call
|
||||
.function
|
||||
.as_ref()
|
||||
.and_then(|function| function.name.clone()),
|
||||
arguments: call.function.and_then(|function| function.arguments),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_dynamo_finish_reason(reason: ChatFinishReason) -> FinishReason {
|
||||
match reason {
|
||||
ChatFinishReason::Stop => FinishReason::Stop,
|
||||
ChatFinishReason::Length => FinishReason::Length,
|
||||
ChatFinishReason::ContentFilter => FinishReason::ContentFilter,
|
||||
ChatFinishReason::ToolCalls => FinishReason::ToolCalls,
|
||||
}
|
||||
}
|
||||
|
||||
fn from_dynamo_finish_reason(reason: FinishReason) -> ChatFinishReason {
|
||||
match reason {
|
||||
FinishReason::Stop => ChatFinishReason::Stop,
|
||||
FinishReason::Length => ChatFinishReason::Length,
|
||||
FinishReason::ContentFilter => ChatFinishReason::ContentFilter,
|
||||
FinishReason::ToolCalls | FinishReason::FunctionCall => ChatFinishReason::ToolCalls,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_reasoning_parser(server_name: &str) -> ReasoningParserWrapper {
|
||||
let name = match server_name {
|
||||
"deepseek-r1" | "step3p5" => "deepseek_r1",
|
||||
"kimi_k2" => "kimi_k25",
|
||||
"gpt-oss" => "gpt_oss",
|
||||
"nemotron_3" => "nemotron3",
|
||||
"interns1" => "qwen3",
|
||||
"qwen3-thinking" | "minimax" => "deepseek_r1",
|
||||
_ => server_name,
|
||||
};
|
||||
ReasoningParserType::get_reasoning_parser_from_name(name)
|
||||
}
|
||||
|
||||
struct ReasoningStreamSplitter {
|
||||
name: Option<String>,
|
||||
parser: Option<ReasoningParserWrapper>,
|
||||
initial_reasoning: Option<bool>,
|
||||
}
|
||||
|
||||
impl ReasoningStreamSplitter {
|
||||
fn new(name: Option<&str>, initial_reasoning: Option<bool>) -> Self {
|
||||
Self {
|
||||
name: name.map(str::to_owned),
|
||||
parser: None,
|
||||
initial_reasoning,
|
||||
}
|
||||
}
|
||||
|
||||
fn split(&mut self, text: &str, token_ids: &[i32]) -> (String, String) {
|
||||
let Some(name) = self.name.as_deref() else {
|
||||
return (String::new(), text.to_owned());
|
||||
};
|
||||
let initial_reasoning = self.initial_reasoning;
|
||||
let parser = self.parser.get_or_insert_with(|| {
|
||||
let mut parser = build_reasoning_parser(name);
|
||||
if let Some(initial_reasoning) = initial_reasoning {
|
||||
parser.set_in_reasoning(initial_reasoning);
|
||||
}
|
||||
parser
|
||||
});
|
||||
let token_ids = token_ids
|
||||
.iter()
|
||||
.filter_map(|&id| u32::try_from(id).ok())
|
||||
.collect::<Vec<_>>();
|
||||
let split = parser.parse_reasoning_streaming_incremental(text, &token_ids);
|
||||
(split.reasoning_text, split.normal_text)
|
||||
}
|
||||
|
||||
fn finish(&mut self) -> (String, String) {
|
||||
let Some(parser) = self.parser.as_mut() else {
|
||||
return (String::new(), String::new());
|
||||
};
|
||||
let tail = parser.finish_reasoning_stream();
|
||||
(tail.reasoning_text, tail.normal_text)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use futures::stream;
|
||||
|
||||
fn processor(
|
||||
tool_parser: Option<&str>,
|
||||
reasoning_parser: Option<&str>,
|
||||
choices: usize,
|
||||
) -> ChatResponseProcessor {
|
||||
ChatResponseProcessor::new(
|
||||
tool_parser.map(str::to_owned),
|
||||
reasoning_parser.map(str::to_owned),
|
||||
None,
|
||||
Some(ChatCompletionToolChoiceOption::Auto),
|
||||
false,
|
||||
true,
|
||||
choices,
|
||||
)
|
||||
}
|
||||
|
||||
fn chunk(choice: usize, text: &str, done: bool) -> Result<DecodedChatEvent, ResponseError> {
|
||||
Ok(DecodedChatEvent {
|
||||
choice,
|
||||
text: text.into(),
|
||||
token_ids: vec![],
|
||||
finish_reason: done.then_some(ChatFinishReason::Stop),
|
||||
logprobs: None,
|
||||
prompt_tokens: 5,
|
||||
completion_tokens: 1,
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn streaming_processor_emits_semantics_without_wire_metadata() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(None, Some("deepseek-r1"), 1)
|
||||
.process_stream(stream::iter(vec![
|
||||
chunk(0, "<think>be", false),
|
||||
chunk(0, "cause</think>Paris", true),
|
||||
]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
let reasoning = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
reasoning_content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
assert_eq!(reasoning, "because");
|
||||
assert!(events.iter().any(|event| matches!(
|
||||
event,
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text), ..
|
||||
}) if text == "Paris"
|
||||
)));
|
||||
assert!(matches!(
|
||||
events.last(),
|
||||
Some(Ok(ChatEvent::Usage {
|
||||
prompt_tokens: 5,
|
||||
completion_tokens: 2
|
||||
}))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn each_choice_has_isolated_reasoning_state() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(None, Some("deepseek-r1"), 2)
|
||||
.process_stream(stream::iter(vec![
|
||||
chunk(0, "<think>zero", false),
|
||||
chunk(1, "<think>one", false),
|
||||
chunk(0, "</think>A", true),
|
||||
chunk(1, "</think>B", true),
|
||||
]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
let deltas = events.iter().filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
choice,
|
||||
content: Some(content),
|
||||
..
|
||||
}) => Some((*choice, content.as_str())),
|
||||
_ => None,
|
||||
});
|
||||
assert_eq!(deltas.collect::<Vec<_>>(), vec![(0, "A"), (1, "B")]);
|
||||
let roles = events.iter().filter_map(|event| match event {
|
||||
Ok(ChatEvent::Role { choice }) => Some(*choice),
|
||||
_ => None,
|
||||
});
|
||||
assert_eq!(roles.collect::<Vec<_>>(), vec![0, 1]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prompt_injected_reasoning_starts_without_opening_marker() {
|
||||
let events = futures::executor::block_on(
|
||||
ChatResponseProcessor::new(
|
||||
None,
|
||||
Some("glm45".into()),
|
||||
None,
|
||||
Some(ChatCompletionToolChoiceOption::Auto),
|
||||
false,
|
||||
true,
|
||||
1,
|
||||
)
|
||||
.with_reasoning_state(Some(true))
|
||||
.process_stream(stream::iter(vec![chunk(
|
||||
0,
|
||||
"reasoning</think>answer",
|
||||
true,
|
||||
)]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
|
||||
let reasoning = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
reasoning_content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
let content = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
assert_eq!(reasoning, "reasoning");
|
||||
assert_eq!(content, "answer");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_reasoning_state_preserves_parser_default() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(None, Some("deepseek-r1"), 1)
|
||||
.process_stream(stream::iter(vec![chunk(
|
||||
0,
|
||||
"reasoning</think>answer",
|
||||
true,
|
||||
)]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
|
||||
let reasoning = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
reasoning_content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
let content = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
assert_eq!(reasoning, "reasoning");
|
||||
assert_eq!(content, "answer");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen_tool_calls_drop_post_call_special_tokens() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(Some("qwen"), None, 1)
|
||||
.process_stream(stream::iter(vec![chunk(
|
||||
0,
|
||||
"Let me check.\n<tool_call>\n{\"name\":\"get_weather\",\"arguments\":{\"city\":\"Paris\"}}\n</tool_call><|im_end|>",
|
||||
true,
|
||||
)]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
|
||||
let content = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
assert!(content.contains("Let me check."));
|
||||
assert!(!content.contains("<|im_end|>"));
|
||||
assert!(events.iter().any(|event| matches!(
|
||||
event,
|
||||
Ok(ChatEvent::Delta {
|
||||
tool_calls: Some(calls),
|
||||
..
|
||||
}) if calls.iter().any(|call| call.name.as_deref() == Some("get_weather"))
|
||||
)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen_tool_calls_drop_split_terminal_special_tokens() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(Some("qwen25"), None, 1)
|
||||
.process_stream(stream::iter(vec![
|
||||
chunk(
|
||||
0,
|
||||
"<tool_call>\n{\"name\":\"get_weather\",\"arguments\":{}}\n</tool_call>",
|
||||
false,
|
||||
),
|
||||
chunk(0, "<|im_end|>", true),
|
||||
]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
|
||||
assert!(!events.iter().any(|event| matches!(
|
||||
event,
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text),
|
||||
..
|
||||
}) if text.contains("<|im_end|>")
|
||||
)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_tool_calls_drop_post_call_special_tokens() {
|
||||
let events = futures::executor::block_on(
|
||||
processor(Some("glm45"), None, 1)
|
||||
.process_stream(stream::iter(vec![
|
||||
chunk(
|
||||
0,
|
||||
"<tool_call>get_weather\n<arg_key>city</arg_key>\n<arg_value>Paris</arg_value>\n</tool_call>",
|
||||
false,
|
||||
),
|
||||
chunk(0, "Follow-up text", false),
|
||||
chunk(0, "<|user|>", true),
|
||||
]))
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
|
||||
let content = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Ok(ChatEvent::Delta {
|
||||
content: Some(text),
|
||||
..
|
||||
}) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
assert_eq!(content, "Follow-up text");
|
||||
assert!(events.iter().any(|event| matches!(
|
||||
event,
|
||||
Ok(ChatEvent::Delta {
|
||||
tool_calls: Some(calls),
|
||||
..
|
||||
}) if calls.iter().any(|call| call.name.as_deref() == Some("get_weather"))
|
||||
)));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,905 @@
|
||||
//! Transport-neutral chat preprocessing over a canonical OpenAI-compatible
|
||||
//! message vocabulary.
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use dynamo_parsers::parsers::get_tool_parser_map;
|
||||
use dynamo_parsers::{
|
||||
StructuralTagBuilder, StructuralTagSchemaMode, ToolCallFormatBuildContext,
|
||||
ToolChoice as DynamoToolChoice, ToolDefinition, TriggeredTagsConfig,
|
||||
};
|
||||
use dynamo_protocols::types::{
|
||||
ChatCompletionRequestAssistantMessageContent, ChatCompletionRequestMessage, ChatCompletionTool,
|
||||
ChatCompletionToolChoiceOption, ResponseFormat,
|
||||
};
|
||||
use dynamo_renderer::{
|
||||
OAIChatLikeRequest, RenderedPrompt, RenderedSegment, may_be_fix_tool_schema,
|
||||
};
|
||||
use minijinja::Value;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use crate::ChatResponseProcessor;
|
||||
use crate::{
|
||||
ChatFormatter, GenerateRequestMetadata, GenerationOptions, OneOrMany, RendererConfig,
|
||||
RendererError, SamplingParams, TextRequest,
|
||||
};
|
||||
|
||||
use super::{GenerateRequestIdentity, TextRequestGroup};
|
||||
|
||||
/// SGLang reasoning effort, including Inkling's fine-grained numeric form.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum ReasoningEffort {
|
||||
None,
|
||||
Minimal,
|
||||
Low,
|
||||
Medium,
|
||||
High,
|
||||
XHigh,
|
||||
Max,
|
||||
Numeric(f64),
|
||||
}
|
||||
|
||||
impl ReasoningEffort {
|
||||
pub(crate) const fn disables_thinking(&self) -> bool {
|
||||
matches!(self, Self::None)
|
||||
}
|
||||
|
||||
const fn name(&self) -> Option<&'static str> {
|
||||
match self {
|
||||
Self::None => Some("none"),
|
||||
Self::Minimal => Some("minimal"),
|
||||
Self::Low => Some("low"),
|
||||
Self::Medium => Some("medium"),
|
||||
Self::High => Some("high"),
|
||||
Self::XHigh => Some("xhigh"),
|
||||
Self::Max => Some("max"),
|
||||
Self::Numeric(_) => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Serialize for ReasoningEffort {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
match self {
|
||||
Self::Numeric(value) => serializer.serialize_f64(*value),
|
||||
_ => serializer.serialize_str(self.name().expect("named reasoning effort")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for ReasoningEffort {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let value = serde_json::Value::deserialize(deserializer)?;
|
||||
match value {
|
||||
serde_json::Value::String(value) => {
|
||||
let effort = match value.as_str() {
|
||||
"none" => Some(Self::None),
|
||||
"minimal" => Some(Self::Minimal),
|
||||
"low" => Some(Self::Low),
|
||||
"medium" => Some(Self::Medium),
|
||||
"high" => Some(Self::High),
|
||||
"xhigh" => Some(Self::XHigh),
|
||||
"max" => Some(Self::Max),
|
||||
_ => None,
|
||||
};
|
||||
if let Some(effort) = effort {
|
||||
return Ok(effort);
|
||||
}
|
||||
let numeric = value.parse::<f64>().map_err(|_| {
|
||||
serde::de::Error::custom(format!("invalid reasoning effort: {value:?}"))
|
||||
})?;
|
||||
numeric_reasoning_effort(numeric).map_err(serde::de::Error::custom)
|
||||
}
|
||||
serde_json::Value::Number(value) => {
|
||||
let numeric = value.as_f64().ok_or_else(|| {
|
||||
serde::de::Error::custom("reasoning_effort must be a finite number")
|
||||
})?;
|
||||
numeric_reasoning_effort(numeric).map_err(serde::de::Error::custom)
|
||||
}
|
||||
serde_json::Value::Bool(_) => Err(serde::de::Error::custom(
|
||||
"reasoning_effort must not be a boolean",
|
||||
)),
|
||||
_ => Err(serde::de::Error::custom(
|
||||
"reasoning_effort must be a string or number",
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn numeric_reasoning_effort(value: f64) -> Result<ReasoningEffort, String> {
|
||||
if !value.is_finite() || !(0.0..=0.99).contains(&value) {
|
||||
return Err(format!(
|
||||
"reasoning_effort must be a finite number in [0.0, 0.99], got {value}"
|
||||
));
|
||||
}
|
||||
Ok(ReasoningEffort::Numeric(value))
|
||||
}
|
||||
|
||||
/// Renderer-owned normalized chat state.
|
||||
///
|
||||
/// Message and tool values remain Dynamo OpenAI protocol types until
|
||||
/// [`ChatPreprocessor`] applies the model chat template and lowers the request
|
||||
/// to the same [`TextRequest`] consumed by text completions.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ChatRequest {
|
||||
pub rid: String,
|
||||
pub model: String,
|
||||
pub messages: Vec<ChatCompletionRequestMessage>,
|
||||
pub tools: Option<Vec<ChatCompletionTool>>,
|
||||
pub tool_choice: Option<ChatCompletionToolChoiceOption>,
|
||||
pub response_format: Option<ResponseFormat>,
|
||||
pub reasoning_effort: Option<ReasoningEffort>,
|
||||
pub continue_final_message: bool,
|
||||
pub chat_template_args: Option<HashMap<String, serde_json::Value>>,
|
||||
pub sampling_params: SamplingParams,
|
||||
pub choice_count: usize,
|
||||
pub stream: bool,
|
||||
pub return_logprob: bool,
|
||||
pub top_logprobs_num: i64,
|
||||
pub parallel_tool_calls: bool,
|
||||
pub metadata: GenerateRequestMetadata,
|
||||
}
|
||||
|
||||
impl OAIChatLikeRequest for ChatRequest {
|
||||
fn model(&self) -> String {
|
||||
self.model.clone()
|
||||
}
|
||||
|
||||
fn messages(&self) -> Value {
|
||||
Value::from_serialize(
|
||||
serde_json::to_value(&self.messages).expect("chat messages serialize"),
|
||||
)
|
||||
}
|
||||
|
||||
fn typed_messages(&self) -> Option<&[ChatCompletionRequestMessage]> {
|
||||
Some(&self.messages)
|
||||
}
|
||||
|
||||
fn tools(&self) -> Option<Value> {
|
||||
self.tools.as_ref().and_then(|tools| {
|
||||
may_be_fix_tool_schema(serde_json::to_value(tools).expect("chat tools serialize"))
|
||||
})
|
||||
}
|
||||
|
||||
fn tool_choice(&self) -> Option<Value> {
|
||||
self.tool_choice.as_ref().map(Value::from_serialize)
|
||||
}
|
||||
|
||||
fn response_format(&self) -> Option<Value> {
|
||||
self.response_format.as_ref().map(Value::from_serialize)
|
||||
}
|
||||
|
||||
fn reasoning_effort(&self) -> Option<Value> {
|
||||
self.reasoning_effort.as_ref().map(Value::from_serialize)
|
||||
}
|
||||
|
||||
fn should_add_generation_prompt(&self) -> bool {
|
||||
!self.continue_final_message
|
||||
}
|
||||
|
||||
fn chat_template_args(&self) -> Option<&HashMap<String, serde_json::Value>> {
|
||||
self.chat_template_args.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
/// Chat-to-text result plus the state needed to interpret generated output.
|
||||
pub(crate) struct LoweredChat {
|
||||
pub text_requests: Vec<TextRequestGroup>,
|
||||
pub response_processor: ChatResponseProcessor,
|
||||
}
|
||||
|
||||
struct RenderPreparation {
|
||||
require_reasoning: bool,
|
||||
reasoning_state: Option<bool>,
|
||||
tools_enabled: bool,
|
||||
}
|
||||
|
||||
/// Applies structured chat semantics before the shared text generation path.
|
||||
pub struct ChatPreprocessor {
|
||||
formatter: Option<ChatFormatter>,
|
||||
formatter_error: Option<String>,
|
||||
tool_call_parser: Option<String>,
|
||||
reasoning_parser: Option<String>,
|
||||
default_chat_template_kwargs: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
impl ChatPreprocessor {
|
||||
pub(crate) fn new(config: &RendererConfig, formatter: Option<ChatFormatter>) -> Self {
|
||||
Self {
|
||||
formatter,
|
||||
formatter_error: None,
|
||||
tool_call_parser: config.tool_call_parser.clone(),
|
||||
reasoning_parser: config.reasoning_parser.clone(),
|
||||
default_chat_template_kwargs: config.default_chat_template_kwargs.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_formatter_error(mut self, error: Option<String>) -> Self {
|
||||
self.formatter_error = error;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn preprocess(&self, mut request: ChatRequest) -> Result<LoweredChat, RendererError> {
|
||||
let preparation = self.prepare_for_render(&mut request)?;
|
||||
merge_template_stops(&mut request.sampling_params, self.formatter.as_ref());
|
||||
|
||||
let tool_choice = dynamo_tool_choice(&request.tool_choice);
|
||||
let tools = chat_tool_definitions(&request);
|
||||
let parser =
|
||||
resolve_chat_parser(self.tool_call_parser.as_deref(), preparation.tools_enabled)?;
|
||||
if parser.is_some() {
|
||||
request.sampling_params.skip_special_tokens = false;
|
||||
}
|
||||
apply_tool_constraint(
|
||||
&mut request.sampling_params,
|
||||
parser.as_deref(),
|
||||
&tool_choice,
|
||||
&tools,
|
||||
Some(request.parallel_tool_calls),
|
||||
)?;
|
||||
let prompt = self.render(&request)?;
|
||||
let uses_tool_call_structural_tag = request.sampling_params.structural_tag.is_some();
|
||||
|
||||
let options = GenerationOptions {
|
||||
sampling_params: request.sampling_params.clone(),
|
||||
require_reasoning: preparation.require_reasoning,
|
||||
stream: request.stream,
|
||||
return_logprob: request.return_logprob,
|
||||
logprob_start_len: -1,
|
||||
top_logprobs_num: request.top_logprobs_num,
|
||||
return_text_in_logprobs: request.return_logprob.then_some(true),
|
||||
..Default::default()
|
||||
};
|
||||
let mut choices = Vec::with_capacity(request.choice_count);
|
||||
for index in 0..request.choice_count {
|
||||
choices.push(GenerateRequestIdentity {
|
||||
rid: format!("{}-{index}", request.rid),
|
||||
metadata: request.metadata.clone(),
|
||||
});
|
||||
}
|
||||
let text_requests = vec![TextRequestGroup {
|
||||
prompt,
|
||||
add_special_tokens: false,
|
||||
options,
|
||||
requests: choices,
|
||||
}];
|
||||
|
||||
let response_processor = ChatResponseProcessor::new(
|
||||
parser,
|
||||
self.reasoning_parser.clone(),
|
||||
(!tools.is_empty()).then_some(tools),
|
||||
request.tool_choice,
|
||||
uses_tool_call_structural_tag,
|
||||
request.parallel_tool_calls,
|
||||
request.choice_count,
|
||||
)
|
||||
.with_reasoning_state(preparation.reasoning_state);
|
||||
Ok(LoweredChat {
|
||||
text_requests,
|
||||
response_processor,
|
||||
})
|
||||
}
|
||||
|
||||
/// Render chat for tokenization without creating generation/output state.
|
||||
pub fn lower_to_text(&self, mut request: ChatRequest) -> Result<TextRequest, RendererError> {
|
||||
let preparation = self.prepare_for_render(&mut request)?;
|
||||
let prompt = self.render(&request)?;
|
||||
Ok(TextRequest::rendered(
|
||||
request.rid,
|
||||
prompt,
|
||||
false,
|
||||
GenerationOptions {
|
||||
sampling_params: request.sampling_params,
|
||||
require_reasoning: preparation.require_reasoning,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.with_metadata(request.metadata))
|
||||
}
|
||||
|
||||
fn prepare_for_render(
|
||||
&self,
|
||||
request: &mut ChatRequest,
|
||||
) -> Result<RenderPreparation, RendererError> {
|
||||
validate_chat(request)?;
|
||||
self.normalize_template_args(request);
|
||||
let tool_choice = dynamo_tool_choice(&request.tool_choice);
|
||||
let tools_enabled = request
|
||||
.tools
|
||||
.as_ref()
|
||||
.is_some_and(|tools| !tools.is_empty())
|
||||
&& tool_choice != DynamoToolChoice::None;
|
||||
let named_tool_choice = matches!(tool_choice, DynamoToolChoice::Named(_));
|
||||
let thinking = self.formatter.as_ref().and_then(|formatter| {
|
||||
formatter.resolve_thinking(
|
||||
&mut request.chat_template_args,
|
||||
tools_enabled,
|
||||
named_tool_choice,
|
||||
)
|
||||
});
|
||||
Ok(RenderPreparation {
|
||||
require_reasoning: self.reasoning_parser.is_some() && thinking == Some(true),
|
||||
reasoning_state: thinking,
|
||||
tools_enabled,
|
||||
})
|
||||
}
|
||||
|
||||
fn normalize_template_args(&self, request: &mut ChatRequest) {
|
||||
let request_args = request.chat_template_args.take().unwrap_or_default();
|
||||
let mut args = self.default_chat_template_kwargs.clone();
|
||||
if let Some(reasoning_effort) = request.reasoning_effort.as_ref() {
|
||||
args.insert(
|
||||
"reasoning_effort".into(),
|
||||
serde_json::to_value(reasoning_effort).expect("reasoning effort must serialize"),
|
||||
);
|
||||
let thinking = !reasoning_effort.disables_thinking();
|
||||
let has_explicit_toggle = request_args.contains_key("thinking")
|
||||
|| request_args.contains_key("enable_thinking");
|
||||
if !has_explicit_toggle {
|
||||
args.insert("thinking".into(), thinking.into());
|
||||
args.insert("enable_thinking".into(), thinking.into());
|
||||
}
|
||||
}
|
||||
args.extend(request_args);
|
||||
request.chat_template_args = (!args.is_empty()).then_some(args);
|
||||
}
|
||||
|
||||
fn render(&self, request: &ChatRequest) -> Result<RenderedPrompt, RendererError> {
|
||||
let formatter = self.formatter.as_ref().ok_or_else(|| {
|
||||
RendererError::from(
|
||||
self.formatter_error
|
||||
.clone()
|
||||
.unwrap_or_else(|| "this model has no usable chat template".to_owned()),
|
||||
)
|
||||
})?;
|
||||
let mut request = request.clone();
|
||||
let final_message = prepare_continuation(&mut request);
|
||||
let template_args = request.chat_template_args.get_or_insert_with(HashMap::new);
|
||||
template_args.insert(
|
||||
"add_generation_prompt".into(),
|
||||
(!request.continue_final_message).into(),
|
||||
);
|
||||
template_args.insert(
|
||||
"continue_final_message".into(),
|
||||
request.continue_final_message.into(),
|
||||
);
|
||||
let prompt = formatter
|
||||
.render_prompt(&request)
|
||||
.map_err(|error| format!("chat template render failed: {error}"))?;
|
||||
match final_message {
|
||||
Some(final_message) => truncate_continuation(prompt, &final_message),
|
||||
None => Ok(prompt),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const CONTINUE_FINAL_MESSAGE_TAG: &str = "CONTINUE_FINAL_MESSAGE_TAG ";
|
||||
|
||||
fn prepare_continuation(request: &mut ChatRequest) -> Option<String> {
|
||||
if !request.continue_final_message {
|
||||
return None;
|
||||
}
|
||||
let Some(ChatCompletionRequestMessage::Assistant(message)) = request.messages.last_mut() else {
|
||||
request.continue_final_message = false;
|
||||
return None;
|
||||
};
|
||||
let Some(ChatCompletionRequestAssistantMessageContent::Text(text)) = message.content.as_mut()
|
||||
else {
|
||||
request.continue_final_message = false;
|
||||
return None;
|
||||
};
|
||||
let original = text.clone();
|
||||
text.push_str(CONTINUE_FINAL_MESSAGE_TAG);
|
||||
Some(original)
|
||||
}
|
||||
|
||||
fn truncate_continuation(
|
||||
prompt: RenderedPrompt,
|
||||
final_message: &str,
|
||||
) -> Result<RenderedPrompt, RendererError> {
|
||||
let text = prompt.as_str();
|
||||
let tag_location = text
|
||||
.rfind(CONTINUE_FINAL_MESSAGE_TAG.trim_end())
|
||||
.filter(|_| text.contains(final_message.trim()))
|
||||
.ok_or_else(|| {
|
||||
RendererError::from(
|
||||
"continue_final_message is set but the final message does not appear in the rendered prompt",
|
||||
)
|
||||
})?;
|
||||
let truncate_at = if text[tag_location..].starts_with(CONTINUE_FINAL_MESSAGE_TAG) {
|
||||
tag_location
|
||||
} else {
|
||||
text[..tag_location].trim_end().len()
|
||||
};
|
||||
Ok(truncate_rendered_prompt(&prompt, truncate_at))
|
||||
}
|
||||
|
||||
fn truncate_rendered_prompt(prompt: &RenderedPrompt, truncate_at: usize) -> RenderedPrompt {
|
||||
let Some(segments) = prompt.segments() else {
|
||||
return RenderedPrompt::text(prompt.as_str()[..truncate_at].to_owned());
|
||||
};
|
||||
let mut remaining = truncate_at;
|
||||
let mut truncated = Vec::new();
|
||||
for segment in segments {
|
||||
if remaining == 0 {
|
||||
break;
|
||||
}
|
||||
let take = remaining.min(segment.text.len());
|
||||
if take != 0 {
|
||||
truncated.push(RenderedSegment::new(
|
||||
segment.text[..take].to_owned(),
|
||||
segment.allow_special,
|
||||
));
|
||||
}
|
||||
remaining -= take;
|
||||
}
|
||||
RenderedPrompt::segmented(truncated)
|
||||
}
|
||||
|
||||
fn validate_chat(request: &ChatRequest) -> Result<(), RendererError> {
|
||||
if request.messages.is_empty() {
|
||||
return Err("messages cannot be empty".into());
|
||||
}
|
||||
if request.choice_count == 0 {
|
||||
return Err("choice_count must be at least 1".into());
|
||||
}
|
||||
if serde_json::to_value(&request.messages).is_ok_and(|messages| contains_media(&messages)) {
|
||||
return Err("image, audio, video, and file message content is not supported".into());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn contains_media(value: &serde_json::Value) -> bool {
|
||||
match value {
|
||||
serde_json::Value::Array(values) => values.iter().any(contains_media),
|
||||
serde_json::Value::Object(object) => {
|
||||
object.keys().any(|key| {
|
||||
matches!(
|
||||
key.as_str(),
|
||||
"image_url" | "video_url" | "input_audio" | "audio_url" | "file"
|
||||
)
|
||||
}) || object.values().any(contains_media)
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_template_stops(sampling: &mut SamplingParams, formatter: Option<&ChatFormatter>) {
|
||||
let Some(template_stops) = formatter.and_then(ChatFormatter::stop_strs) else {
|
||||
return;
|
||||
};
|
||||
let mut stops = match template_stops {
|
||||
OneOrMany::One(stop) => vec![stop],
|
||||
OneOrMany::Many(stops) => stops,
|
||||
};
|
||||
if let Some(request_stops) = sampling.stop.take() {
|
||||
match request_stops {
|
||||
OneOrMany::One(stop) => stops.push(stop),
|
||||
OneOrMany::Many(request_stops) => stops.extend(request_stops),
|
||||
}
|
||||
}
|
||||
sampling.stop = Some(OneOrMany::Many(stops));
|
||||
}
|
||||
|
||||
fn resolve_chat_parser(
|
||||
configured_parser: Option<&str>,
|
||||
tools_enabled: bool,
|
||||
) -> Result<Option<String>, RendererError> {
|
||||
if tools_enabled && configured_parser.is_none() {
|
||||
return Err("tool calls require --tool-call-parser".into());
|
||||
}
|
||||
Ok(tools_enabled.then(|| configured_parser.expect("checked").to_owned()))
|
||||
}
|
||||
|
||||
fn chat_tool_definitions(request: &ChatRequest) -> Vec<ToolDefinition> {
|
||||
request
|
||||
.tools
|
||||
.iter()
|
||||
.flatten()
|
||||
.map(|tool| ToolDefinition {
|
||||
name: tool.function.name.clone(),
|
||||
parameters: tool.function.parameters.clone(),
|
||||
strict: tool.function.strict,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn dynamo_parser_name(parser: &str) -> &str {
|
||||
match parser {
|
||||
"llama3" => "llama3_json",
|
||||
"qwen" => "qwen25",
|
||||
"glm" | "glm45" => "glm47",
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
|
||||
fn dynamo_tool_choice(choice: &Option<ChatCompletionToolChoiceOption>) -> DynamoToolChoice {
|
||||
match choice {
|
||||
Some(ChatCompletionToolChoiceOption::None) => DynamoToolChoice::None,
|
||||
Some(ChatCompletionToolChoiceOption::Required) => DynamoToolChoice::Required,
|
||||
Some(ChatCompletionToolChoiceOption::Named(choice)) => {
|
||||
DynamoToolChoice::Named(choice.function.name.clone())
|
||||
}
|
||||
Some(ChatCompletionToolChoiceOption::Auto) | None => DynamoToolChoice::Auto,
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_tool_constraint(
|
||||
sampling: &mut SamplingParams,
|
||||
parser: Option<&str>,
|
||||
tool_choice: &DynamoToolChoice,
|
||||
tools: &[ToolDefinition],
|
||||
parallel_tool_calls: Option<bool>,
|
||||
) -> Result<(), String> {
|
||||
if *tool_choice == DynamoToolChoice::None {
|
||||
return Ok(());
|
||||
}
|
||||
if *tool_choice == DynamoToolChoice::Required && tools.is_empty() {
|
||||
return Err("tool_choice is \"required\" but tools is empty".into());
|
||||
}
|
||||
if let DynamoToolChoice::Named(name) = tool_choice
|
||||
&& !tools.iter().any(|tool| &tool.name == name)
|
||||
{
|
||||
return Err(format!(
|
||||
"tool named \"{name}\" in tool_choice is not present in tools"
|
||||
));
|
||||
}
|
||||
|
||||
let Some(parser) = parser else {
|
||||
return Ok(());
|
||||
};
|
||||
let parser = dynamo_parser_name(parser);
|
||||
let config = get_tool_parser_map()
|
||||
.get(parser)
|
||||
.ok_or_else(|| format!("tool-call parser `{parser}` is not supported by Dynamo"))?;
|
||||
let builder = config.structural_tag_builder.clone().or_else(|| {
|
||||
(parser == "llama3_json"
|
||||
&& *tool_choice == DynamoToolChoice::Auto
|
||||
&& tools.iter().any(|tool| tool.strict.unwrap_or(false)))
|
||||
.then(|| {
|
||||
StructuralTagBuilder::TriggeredTags(TriggeredTagsConfig {
|
||||
begin_template: r#"<|python_tag|>{"name":"{name}", "arguments":"#.to_string(),
|
||||
end_template: "}".to_string(),
|
||||
triggers: vec!["<|python_tag|>".to_string()],
|
||||
content_style: Default::default(),
|
||||
tool_call_ban_tokens: Vec::new(),
|
||||
reasoning_end: None,
|
||||
})
|
||||
})
|
||||
});
|
||||
if let Some(builder) = builder
|
||||
&& let Some(tag) = builder
|
||||
.build_tool_call_format(&ToolCallFormatBuildContext {
|
||||
tool_choice,
|
||||
tools,
|
||||
parallel_tool_calls,
|
||||
schema_mode: StructuralTagSchemaMode::Auto,
|
||||
starts_in_reasoning: false,
|
||||
})
|
||||
.map_err(|error| error.to_string())?
|
||||
{
|
||||
sampling.structural_tag = Some(tag.to_string());
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if matches!(
|
||||
tool_choice,
|
||||
DynamoToolChoice::Required | DynamoToolChoice::Named(_)
|
||||
) {
|
||||
let selected = match tool_choice {
|
||||
DynamoToolChoice::Named(name) => tools
|
||||
.iter()
|
||||
.filter(|tool| tool.name == *name)
|
||||
.collect::<Vec<_>>(),
|
||||
_ => tools.iter().collect(),
|
||||
};
|
||||
let schemas = selected
|
||||
.into_iter()
|
||||
.map(|tool| {
|
||||
serde_json::json!({
|
||||
"properties": {
|
||||
"name": {"type": "string", "enum": [tool.name]},
|
||||
"parameters": tool.parameters.clone().unwrap_or_else(|| {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}),
|
||||
},
|
||||
"required": ["name", "parameters"],
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let items = if schemas.len() == 1 {
|
||||
schemas.into_iter().next().expect("one schema")
|
||||
} else {
|
||||
serde_json::json!({"type": "object", "anyOf": schemas})
|
||||
};
|
||||
let mut schema = serde_json::json!({
|
||||
"type": "array",
|
||||
"minItems": 1,
|
||||
"items": items,
|
||||
});
|
||||
if parallel_tool_calls == Some(false) {
|
||||
schema["maxItems"] = serde_json::json!(1);
|
||||
}
|
||||
sampling.json_schema = Some(schema.to_string());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{RendererLimits, SamplingDefaults};
|
||||
use dynamo_protocols::types::{
|
||||
ChatCompletionNamedToolChoice, ChatCompletionToolType, FunctionName,
|
||||
};
|
||||
|
||||
fn tool(name: &str, strict: bool) -> ToolDefinition {
|
||||
ToolDefinition {
|
||||
name: name.into(),
|
||||
parameters: Some(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"]
|
||||
})),
|
||||
strict: Some(strict),
|
||||
}
|
||||
}
|
||||
|
||||
fn chat_request(tool_choice: Option<ChatCompletionToolChoiceOption>) -> ChatRequest {
|
||||
ChatRequest {
|
||||
rid: "chatcmpl-test".into(),
|
||||
model: "model".into(),
|
||||
messages: serde_json::from_value(serde_json::json!([
|
||||
{"role": "user", "content": "hello"}
|
||||
]))
|
||||
.unwrap(),
|
||||
tools: Some(
|
||||
serde_json::from_value(serde_json::json!([{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object"}
|
||||
}
|
||||
}]))
|
||||
.unwrap(),
|
||||
),
|
||||
tool_choice,
|
||||
response_format: None,
|
||||
reasoning_effort: None,
|
||||
continue_final_message: false,
|
||||
chat_template_args: None,
|
||||
sampling_params: SamplingParams::default(),
|
||||
choice_count: 1,
|
||||
stream: false,
|
||||
return_logprob: false,
|
||||
top_logprobs_num: 0,
|
||||
parallel_tool_calls: true,
|
||||
metadata: GenerateRequestMetadata::default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn chat_preprocessor() -> ChatPreprocessor {
|
||||
chat_preprocessor_with(
|
||||
Some("llama3"),
|
||||
None,
|
||||
crate::preprocessing::template::load_chat_formatter(None, None, Some("chatml"))
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn chat_preprocessor_with(
|
||||
tool_call_parser: Option<&str>,
|
||||
reasoning_parser: Option<&str>,
|
||||
formatter: ChatFormatter,
|
||||
) -> ChatPreprocessor {
|
||||
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: tool_call_parser.map(str::to_owned),
|
||||
reasoning_parser: reasoning_parser.map(str::to_owned),
|
||||
default_chat_template_kwargs: Default::default(),
|
||||
stream_response_default_include_usage: false,
|
||||
default_sampling_params: SamplingDefaults::default(),
|
||||
limits: RendererLimits {
|
||||
vocab_size: 128,
|
||||
context_len: 128,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
};
|
||||
ChatPreprocessor::new(&config, Some(formatter))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wire_tool_choices_lower_to_internal_choices() {
|
||||
let named = Some(ChatCompletionToolChoiceOption::Named(
|
||||
ChatCompletionNamedToolChoice {
|
||||
r#type: ChatCompletionToolType::Function,
|
||||
function: FunctionName {
|
||||
name: "get_weather".into(),
|
||||
},
|
||||
},
|
||||
));
|
||||
|
||||
assert!(matches!(dynamo_tool_choice(&None), DynamoToolChoice::Auto));
|
||||
assert!(matches!(
|
||||
dynamo_tool_choice(&Some(ChatCompletionToolChoiceOption::Required)),
|
||||
DynamoToolChoice::Required
|
||||
));
|
||||
assert!(matches!(
|
||||
dynamo_tool_choice(&named),
|
||||
DynamoToolChoice::Named(name) if name == "get_weather"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn required_choice_builds_a_single_call_constraint() {
|
||||
let mut sampling = SamplingParams::default();
|
||||
apply_tool_constraint(
|
||||
&mut sampling,
|
||||
Some("llama3"),
|
||||
&DynamoToolChoice::Required,
|
||||
&[tool("get_weather", false), tool("get_time", false)],
|
||||
Some(false),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let schema: serde_json::Value =
|
||||
serde_json::from_str(sampling.json_schema.as_deref().unwrap()).unwrap();
|
||||
assert_eq!(schema["minItems"], 1);
|
||||
assert_eq!(schema["maxItems"], 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_tool_choices_are_rejected_before_generation() {
|
||||
let mut sampling = SamplingParams::default();
|
||||
assert!(
|
||||
apply_tool_constraint(&mut sampling, None, &DynamoToolChoice::Required, &[], None,)
|
||||
.unwrap_err()
|
||||
.contains("required")
|
||||
);
|
||||
assert!(
|
||||
apply_tool_constraint(
|
||||
&mut sampling,
|
||||
None,
|
||||
&DynamoToolChoice::Named("missing".into()),
|
||||
&[tool("get_weather", false)],
|
||||
None,
|
||||
)
|
||||
.unwrap_err()
|
||||
.contains("missing")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_parsing_preserves_special_tokens_for_output_processing() {
|
||||
let mut request = chat_request(None);
|
||||
request.sampling_params.skip_special_tokens = true;
|
||||
|
||||
let chat = chat_preprocessor().preprocess(request).unwrap();
|
||||
|
||||
assert!(
|
||||
!chat.text_requests[0]
|
||||
.options
|
||||
.sampling_params
|
||||
.skip_special_tokens
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_choice_none_keeps_the_requested_special_token_behavior() {
|
||||
let mut request = chat_request(Some(ChatCompletionToolChoiceOption::None));
|
||||
request.sampling_params.skip_special_tokens = true;
|
||||
|
||||
let chat = chat_preprocessor().preprocess(request).unwrap();
|
||||
|
||||
assert!(
|
||||
chat.text_requests[0]
|
||||
.options
|
||||
.sampling_params
|
||||
.skip_special_tokens
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen_required_tools_forward_effective_template_thinking() {
|
||||
let formatter = crate::preprocessing::template::test_hugging_face_formatter(
|
||||
"{% if enable_thinking is not defined %}{% set enable_thinking = true %}{% endif %}{{ enable_thinking }}",
|
||||
);
|
||||
let preprocessor = chat_preprocessor_with(Some("qwen"), Some("qwen3"), formatter);
|
||||
|
||||
let enabled = preprocessor
|
||||
.preprocess(chat_request(Some(ChatCompletionToolChoiceOption::Required)))
|
||||
.unwrap();
|
||||
assert!(enabled.text_requests[0].options.require_reasoning);
|
||||
|
||||
let mut disabled_request = chat_request(Some(ChatCompletionToolChoiceOption::Required));
|
||||
disabled_request.reasoning_effort = Some(ReasoningEffort::Max);
|
||||
disabled_request.chat_template_args = Some(HashMap::from([(
|
||||
"enable_thinking".into(),
|
||||
serde_json::Value::Bool(false),
|
||||
)]));
|
||||
let disabled = preprocessor.preprocess(disabled_request).unwrap();
|
||||
assert!(!disabled.text_requests[0].options.require_reasoning);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn thinking_policy_uses_the_effective_tool_template() {
|
||||
let formatter = crate::preprocessing::template::test_hugging_face_formatter_from_config(
|
||||
serde_json::json!({
|
||||
"chat_template": [
|
||||
{"default": "{{ enable_thinking | default(false) }}"},
|
||||
{"tool_use": "{{ enable_thinking | default(true) }}"}
|
||||
]
|
||||
}),
|
||||
);
|
||||
let preprocessor = chat_preprocessor_with(Some("qwen"), Some("qwen3"), formatter);
|
||||
|
||||
let mut no_tools = chat_request(None);
|
||||
no_tools.tools = None;
|
||||
assert!(
|
||||
!preprocessor.preprocess(no_tools).unwrap().text_requests[0]
|
||||
.options
|
||||
.require_reasoning
|
||||
);
|
||||
|
||||
let mut empty_tools = chat_request(None);
|
||||
empty_tools.tools = Some(Vec::new());
|
||||
assert!(
|
||||
!preprocessor.preprocess(empty_tools).unwrap().text_requests[0]
|
||||
.options
|
||||
.require_reasoning
|
||||
);
|
||||
|
||||
assert!(
|
||||
!preprocessor
|
||||
.preprocess(chat_request(Some(ChatCompletionToolChoiceOption::None)))
|
||||
.unwrap()
|
||||
.text_requests[0]
|
||||
.options
|
||||
.require_reasoning
|
||||
);
|
||||
assert!(
|
||||
preprocessor
|
||||
.preprocess(chat_request(Some(ChatCompletionToolChoiceOption::Required)))
|
||||
.unwrap()
|
||||
.text_requests[0]
|
||||
.options
|
||||
.require_reasoning
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn always_on_channel_template_requires_reasoning() {
|
||||
let formatter = crate::preprocessing::template::test_hugging_face_formatter(
|
||||
"<|start|>assistant<|channel|>analysis<|message|>",
|
||||
);
|
||||
let preprocessor = chat_preprocessor_with(None, Some("gpt-oss"), formatter);
|
||||
let mut request = chat_request(None);
|
||||
request.tools = None;
|
||||
request.response_format = Some(
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "answer",
|
||||
"schema": {"type": "object"}
|
||||
}
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let lowered = preprocessor.preprocess(request).unwrap();
|
||||
|
||||
assert!(lowered.text_requests[0].options.require_reasoning);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
//! Request processing from protocol-neutral inputs to token-only generation requests.
|
||||
|
||||
mod chat;
|
||||
mod regex;
|
||||
mod request;
|
||||
mod sampling;
|
||||
mod service;
|
||||
mod template;
|
||||
mod tokenizer;
|
||||
|
||||
pub(crate) use chat::{ChatPreprocessor, LoweredChat, dynamo_parser_name};
|
||||
pub use chat::{ChatRequest, ReasoningEffort};
|
||||
pub use request::{
|
||||
GenerateRequest, GenerateRequestMetadata, GenerateSamplingParams, GenerationOptions,
|
||||
TextRequest, TokenIdsRequest,
|
||||
};
|
||||
pub(crate) use request::{GenerateRequestIdentity, TextRequestGroup};
|
||||
pub use sampling::SamplingParams;
|
||||
pub(crate) use sampling::SamplingParamsOverrides;
|
||||
pub use service::{PreparedChat, RendererService};
|
||||
pub(crate) use template::ChatFormatter;
|
||||
#[cfg(test)]
|
||||
pub(crate) fn load_test_chat_formatter(name: &str) -> ChatFormatter {
|
||||
template::load_chat_formatter(None, None, Some(name)).unwrap()
|
||||
}
|
||||
pub use tokenizer::{DynamoTokenizer, TextTokenizer, load_tokenizer};
|
||||
#[cfg(feature = "http")]
|
||||
pub(crate) use tokenizer::{resolve_model_file, resolve_tokenizer_file};
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,282 @@
|
||||
//! Internal and transport request representations.
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use dynamo_renderer::RenderedPrompt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{SamplingParams, TokenIds};
|
||||
|
||||
/// Request-scoped metadata that must survive protocol lowering and prompt
|
||||
/// tokenization before the request is submitted to SGLang `/generate`.
|
||||
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct GenerateRequestMetadata {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub model: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_salt: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extra_key: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub priority: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bootstrap_host: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bootstrap_port: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bootstrap_room: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub routed_dp_rank: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub disagg_prefill_dp_rank: Option<i64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
/// Generation options shared by text and token-ID inputs.
|
||||
pub struct GenerationOptions {
|
||||
pub sampling_params: SamplingParams,
|
||||
/// Delay structured-output constraints until the model finishes reasoning.
|
||||
pub require_reasoning: bool,
|
||||
pub stream: bool,
|
||||
pub return_logprob: bool,
|
||||
pub logprob_start_len: i64,
|
||||
pub top_logprobs_num: i64,
|
||||
pub token_ids_logprob: Option<TokenIds>,
|
||||
pub return_hidden_states: bool,
|
||||
pub return_text_in_logprobs: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
/// Internal text-only generation request before tokenization.
|
||||
///
|
||||
/// Protocol adapters lower textual completions into this type. Structured chat
|
||||
/// reaches it only after [`crate::ChatPreprocessor`] renders the messages.
|
||||
pub struct TextRequest {
|
||||
pub rid: String,
|
||||
pub prompt: RenderedPrompt,
|
||||
pub add_special_tokens: bool,
|
||||
pub options: GenerationOptions,
|
||||
pub metadata: GenerateRequestMetadata,
|
||||
}
|
||||
|
||||
/// One textual prompt shared by one or more generation choices.
|
||||
///
|
||||
/// OpenAI `n` fan-out changes request identity, not the prompt or generation
|
||||
/// options. Keeping those identities alongside one prompt lets preprocessing
|
||||
/// tokenize the prompt once before producing the individual engine requests.
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct TextRequestGroup {
|
||||
pub prompt: RenderedPrompt,
|
||||
pub add_special_tokens: bool,
|
||||
pub options: GenerationOptions,
|
||||
pub requests: Vec<GenerateRequestIdentity>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct GenerateRequestIdentity {
|
||||
pub rid: String,
|
||||
pub metadata: GenerateRequestMetadata,
|
||||
}
|
||||
|
||||
impl From<TextRequest> for TextRequestGroup {
|
||||
fn from(request: TextRequest) -> Self {
|
||||
Self {
|
||||
prompt: request.prompt,
|
||||
add_special_tokens: request.add_special_tokens,
|
||||
options: request.options,
|
||||
requests: vec![GenerateRequestIdentity {
|
||||
rid: request.rid,
|
||||
metadata: request.metadata,
|
||||
}],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TextRequest {
|
||||
pub fn text(
|
||||
rid: impl Into<String>,
|
||||
text: impl Into<String>,
|
||||
add_special_tokens: bool,
|
||||
options: GenerationOptions,
|
||||
) -> Self {
|
||||
Self {
|
||||
rid: rid.into(),
|
||||
prompt: RenderedPrompt::text(text.into()),
|
||||
add_special_tokens,
|
||||
options,
|
||||
metadata: GenerateRequestMetadata::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn rendered(
|
||||
rid: impl Into<String>,
|
||||
prompt: RenderedPrompt,
|
||||
add_special_tokens: bool,
|
||||
options: GenerationOptions,
|
||||
) -> Self {
|
||||
Self {
|
||||
rid: rid.into(),
|
||||
prompt,
|
||||
add_special_tokens,
|
||||
options,
|
||||
metadata: GenerateRequestMetadata::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_metadata(mut self, metadata: GenerateRequestMetadata) -> Self {
|
||||
self.metadata = metadata;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
/// A generation request whose prompt is already represented by token IDs.
|
||||
pub struct TokenIdsRequest {
|
||||
pub rid: String,
|
||||
pub input_ids: TokenIds,
|
||||
pub options: GenerationOptions,
|
||||
pub metadata: GenerateRequestMetadata,
|
||||
}
|
||||
|
||||
impl TokenIdsRequest {
|
||||
pub fn new(rid: impl Into<String>, input_ids: TokenIds, options: GenerationOptions) -> Self {
|
||||
Self {
|
||||
rid: rid.into(),
|
||||
input_ids,
|
||||
options,
|
||||
metadata: GenerateRequestMetadata::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_metadata(mut self, metadata: GenerateRequestMetadata) -> Self {
|
||||
self.metadata = metadata;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Token-only request sent to the model server's `/generate` endpoint.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct GenerateRequest {
|
||||
pub rid: String,
|
||||
#[serde(flatten)]
|
||||
pub metadata: GenerateRequestMetadata,
|
||||
pub input_ids: TokenIds,
|
||||
#[serde(default)]
|
||||
pub require_reasoning: bool,
|
||||
pub sampling_params: GenerateSamplingParams,
|
||||
pub stream: bool,
|
||||
pub return_logprob: bool,
|
||||
pub logprob_start_len: i64,
|
||||
pub top_logprobs_num: i64,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub token_ids_logprob: Option<TokenIds>,
|
||||
pub return_hidden_states: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub return_text_in_logprobs: Option<bool>,
|
||||
}
|
||||
|
||||
impl From<TokenIdsRequest> for GenerateRequest {
|
||||
fn from(request: TokenIdsRequest) -> Self {
|
||||
let options = request.options;
|
||||
Self {
|
||||
rid: request.rid,
|
||||
metadata: request.metadata,
|
||||
input_ids: request.input_ids,
|
||||
require_reasoning: options.require_reasoning,
|
||||
sampling_params: options.sampling_params.into(),
|
||||
stream: options.stream,
|
||||
return_logprob: options.return_logprob,
|
||||
logprob_start_len: options.logprob_start_len,
|
||||
top_logprobs_num: options.top_logprobs_num,
|
||||
token_ids_logprob: options.token_ids_logprob,
|
||||
return_hidden_states: options.return_hidden_states,
|
||||
return_text_in_logprobs: options.return_text_in_logprobs,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn require_reasoning_is_forwarded_as_a_boolean() {
|
||||
let request = |require_reasoning| {
|
||||
GenerateRequest::from(TokenIdsRequest::new(
|
||||
"request",
|
||||
vec![1, 2],
|
||||
GenerationOptions {
|
||||
require_reasoning,
|
||||
..Default::default()
|
||||
},
|
||||
))
|
||||
};
|
||||
|
||||
let enabled = serde_json::to_value(request(true)).unwrap();
|
||||
assert_eq!(enabled["require_reasoning"], true);
|
||||
|
||||
let disabled = serde_json::to_value(request(false)).unwrap();
|
||||
assert_eq!(disabled["require_reasoning"], false);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct GenerateSamplingParams {
|
||||
pub max_new_tokens: Option<i64>,
|
||||
pub stop: Vec<String>,
|
||||
pub stop_token_ids: Option<Vec<i64>>,
|
||||
pub stop_regex: Vec<String>,
|
||||
pub temperature: f64,
|
||||
pub top_p: f64,
|
||||
pub top_k: i64,
|
||||
pub min_p: f64,
|
||||
pub frequency_penalty: f64,
|
||||
pub presence_penalty: f64,
|
||||
pub repetition_penalty: f64,
|
||||
pub min_new_tokens: i64,
|
||||
pub n: i64,
|
||||
pub json_schema: Option<String>,
|
||||
pub regex: Option<String>,
|
||||
pub ebnf: Option<String>,
|
||||
pub structural_tag: Option<String>,
|
||||
pub ignore_eos: bool,
|
||||
pub skip_special_tokens: bool,
|
||||
pub spaces_between_special_tokens: bool,
|
||||
pub no_stop_trim: bool,
|
||||
pub stream_interval: Option<i64>,
|
||||
pub logit_bias: Option<BTreeMap<String, f64>>,
|
||||
pub sampling_seed: Option<i64>,
|
||||
pub custom_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl From<SamplingParams> for GenerateSamplingParams {
|
||||
fn from(params: SamplingParams) -> Self {
|
||||
Self {
|
||||
max_new_tokens: params.max_new_tokens,
|
||||
stop: params.stop_strs,
|
||||
stop_token_ids: params.stop_token_ids,
|
||||
stop_regex: params.stop_regex_strs,
|
||||
temperature: params.temperature,
|
||||
top_p: params.top_p,
|
||||
top_k: params.top_k,
|
||||
min_p: params.min_p,
|
||||
frequency_penalty: params.frequency_penalty,
|
||||
presence_penalty: params.presence_penalty,
|
||||
repetition_penalty: params.repetition_penalty,
|
||||
min_new_tokens: params.min_new_tokens,
|
||||
n: params.n,
|
||||
json_schema: params.json_schema,
|
||||
regex: params.regex,
|
||||
ebnf: params.ebnf,
|
||||
structural_tag: params.structural_tag,
|
||||
ignore_eos: params.ignore_eos,
|
||||
skip_special_tokens: params.skip_special_tokens,
|
||||
spaces_between_special_tokens: params.spaces_between_special_tokens,
|
||||
no_stop_trim: params.no_stop_trim,
|
||||
stream_interval: params.stream_interval,
|
||||
logit_bias: params.logit_bias,
|
||||
sampling_seed: params.sampling_seed,
|
||||
custom_params: params.custom_params,
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,117 @@
|
||||
//! Adapt SGLang's DeepSeek V4 effort profiles to Dynamo's native formatter.
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum DeepSeekV4Profile {
|
||||
Preview,
|
||||
Official,
|
||||
}
|
||||
|
||||
pub(super) fn dynamo_reasoning_effort(
|
||||
profile: DeepSeekV4Profile,
|
||||
effort: Option<&str>,
|
||||
) -> &'static str {
|
||||
match (profile, effort) {
|
||||
(DeepSeekV4Profile::Preview, Some("max")) | (DeepSeekV4Profile::Official, Some("high")) => {
|
||||
"high"
|
||||
}
|
||||
(DeepSeekV4Profile::Official, Some("max")) => "max",
|
||||
// Dynamo's low effort preserves thinking without adding a prefix.
|
||||
_ => "low",
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use dynamo_protocols::types::CreateChatCompletionRequest;
|
||||
use dynamo_renderer::PromptFormatter;
|
||||
use dynamo_renderer::deepseek::v4::DeepSeekV4Formatter;
|
||||
|
||||
use super::super::{ChatFormatter, TemplateArgsRequest};
|
||||
use super::DeepSeekV4Profile;
|
||||
|
||||
#[test]
|
||||
fn deepseek_v4_profiles_map_effort_without_coercing_unsupported_tiers() {
|
||||
fn render(
|
||||
profile: DeepSeekV4Profile,
|
||||
effort: Option<&str>,
|
||||
thinking: Option<bool>,
|
||||
environment_effort: Option<&str>,
|
||||
) -> String {
|
||||
let request: CreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "test",
|
||||
"messages": [
|
||||
{"role": "system", "content": "Be concise."},
|
||||
{"role": "user", "content": "Hello"}
|
||||
]
|
||||
}))
|
||||
.unwrap();
|
||||
let mut args = HashMap::new();
|
||||
if let Some(effort) = effort {
|
||||
args.insert("reasoning_effort".into(), serde_json::json!(effort));
|
||||
}
|
||||
if let Some(thinking) = thinking {
|
||||
args.insert("thinking".into(), serde_json::json!(thinking));
|
||||
}
|
||||
ChatFormatter::DeepSeekV4 {
|
||||
formatter: PromptFormatter::OAI(Arc::new(DeepSeekV4Formatter::new_chat())),
|
||||
profile,
|
||||
environment_effort: environment_effort.map(str::to_owned),
|
||||
}
|
||||
.render(&TemplateArgsRequest {
|
||||
request: &request,
|
||||
args,
|
||||
})
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
let baseline = "<|begin▁of▁sentence|>Be concise.<|User|>Hello<|Assistant|><think>";
|
||||
for (profile, high_prefix, max_prefix) in [
|
||||
(DeepSeekV4Profile::Preview, None, "Absolute maximum"),
|
||||
(
|
||||
DeepSeekV4Profile::Official,
|
||||
Some("Absolute maximum"),
|
||||
"Beyond maximum",
|
||||
),
|
||||
] {
|
||||
for (effort, prefix) in [
|
||||
(None, None),
|
||||
(Some("low"), None),
|
||||
(Some("high"), high_prefix),
|
||||
(Some("max"), Some(max_prefix)),
|
||||
(Some("xhigh"), None),
|
||||
] {
|
||||
let prompt = render(profile, effort, Some(true), None);
|
||||
assert_eq!(
|
||||
prompt.matches("Reasoning Effort:").count(),
|
||||
usize::from(prefix.is_some()),
|
||||
"{profile:?}, {effort:?}: {prompt}"
|
||||
);
|
||||
if let Some(prefix) = prefix {
|
||||
assert!(prompt.starts_with(&format!(
|
||||
"<|begin▁of▁sentence|>Reasoning Effort: {prefix}"
|
||||
)));
|
||||
assert_eq!(
|
||||
prompt.split_once("\n\n").unwrap().1,
|
||||
baseline.strip_prefix("<|begin▁of▁sentence|>").unwrap()
|
||||
);
|
||||
} else {
|
||||
assert_eq!(prompt, baseline);
|
||||
}
|
||||
}
|
||||
let disabled = baseline.replace("<think>", "</think>");
|
||||
assert_eq!(render(profile, None, None, None), disabled);
|
||||
assert_eq!(render(profile, Some("max"), Some(false), None), disabled);
|
||||
assert_eq!(
|
||||
render(profile, None, Some(true), Some("max")),
|
||||
render(profile, Some("max"), Some(true), None)
|
||||
);
|
||||
assert_eq!(
|
||||
render(profile, Some("low"), Some(true), Some("max")),
|
||||
baseline
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,725 @@
|
||||
//! Kimi K2.5 checkpoint-compatible tool declaration preprocessing.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::fmt::Write as _;
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const INDENT: &str = " ";
|
||||
const FIELD_DELIMITER: &str = ",\n";
|
||||
const MAX_RECURSION_DEPTH: usize = 32;
|
||||
|
||||
pub(crate) fn deep_sort(value: &mut Value) {
|
||||
match value {
|
||||
Value::Object(object) => {
|
||||
let mut entries: Vec<_> = std::mem::take(object).into_iter().collect();
|
||||
for (_, value) in &mut entries {
|
||||
deep_sort(value);
|
||||
}
|
||||
entries.sort_by(|left, right| left.0.cmp(&right.0));
|
||||
*object = entries.into_iter().collect::<Map<_, _>>();
|
||||
}
|
||||
Value::Array(array) => {
|
||||
for value in array {
|
||||
deep_sort(value);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn encode_tools_to_typescript(tools: &[Value]) -> Option<String> {
|
||||
if tools.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut functions = Vec::new();
|
||||
for tool in tools {
|
||||
if tool.get("type").and_then(Value::as_str) != Some("function") {
|
||||
continue;
|
||||
}
|
||||
let function = match tool.get("function") {
|
||||
Some(function)
|
||||
if function
|
||||
.as_object()
|
||||
.is_some_and(|object| !object.is_empty()) =>
|
||||
{
|
||||
function
|
||||
}
|
||||
_ => continue,
|
||||
};
|
||||
match encode_function(function) {
|
||||
Some(function) => functions.push(function),
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Kimi K2.5 tool schema is unsupported by the TypeScript encoder; using the checkpoint JSON fallback"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
}
|
||||
}
|
||||
if functions.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(format!(
|
||||
"# Tools\n\n## functions\nnamespace functions {{\n{}\n}}\n",
|
||||
functions.join("\n")
|
||||
))
|
||||
}
|
||||
|
||||
fn encode_function(function: &Value) -> Option<String> {
|
||||
let parameters = function
|
||||
.get("parameters")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| Value::Object(Map::new()));
|
||||
let mut registry = SchemaRegistry::default();
|
||||
let parsed = ObjectType::parse(¶meters, &mut registry);
|
||||
let mut interfaces = Vec::new();
|
||||
|
||||
let root_name = if registry.has_self_ref {
|
||||
let body = parsed
|
||||
.properties
|
||||
.iter()
|
||||
.map(|parameter| parameter.to_typescript(INDENT, ®istry))
|
||||
.collect::<Vec<_>>()
|
||||
.join(FIELD_DELIMITER);
|
||||
let body = if body.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!("\n{body}\n")
|
||||
};
|
||||
interfaces.push(format!("interface parameters {{{body}}}"));
|
||||
Some("parameters")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let definitions = registry
|
||||
.order
|
||||
.iter()
|
||||
.filter_map(|name| {
|
||||
registry
|
||||
.definitions
|
||||
.get(name)
|
||||
.map(|schema| (name.clone(), schema.clone()))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
for (name, schema) in definitions {
|
||||
let object = parse_type(&schema, &mut registry);
|
||||
let mut definition = String::new();
|
||||
if let Some(description) = schema.get("description").and_then(Value::as_str)
|
||||
&& !description.is_empty()
|
||||
{
|
||||
definition.push_str(&format_description(description, ""));
|
||||
definition.push('\n');
|
||||
}
|
||||
definition.push_str(&format!(
|
||||
"interface {name} {}",
|
||||
object.to_typescript("", ®istry)
|
||||
));
|
||||
interfaces.push(definition);
|
||||
}
|
||||
|
||||
if registry.unsupported {
|
||||
return None;
|
||||
}
|
||||
let name = function
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("function");
|
||||
let type_definition = match root_name {
|
||||
Some(root_name) => format!("type {name} = (_: {root_name}) => any;"),
|
||||
None => format!(
|
||||
"type {name} = (_: {}) => any;",
|
||||
parsed.to_typescript("", ®istry)
|
||||
),
|
||||
};
|
||||
let description = function
|
||||
.get("description")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|description| !description.is_empty())
|
||||
.map(|description| format_description(description, ""))
|
||||
.unwrap_or_default();
|
||||
Some(
|
||||
[interfaces.join("\n"), description, type_definition]
|
||||
.into_iter()
|
||||
.filter(|part| !part.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
)
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct SchemaRegistry {
|
||||
definitions: HashMap<String, Value>,
|
||||
order: Vec<String>,
|
||||
has_self_ref: bool,
|
||||
depth: usize,
|
||||
unsupported: bool,
|
||||
}
|
||||
|
||||
impl SchemaRegistry {
|
||||
fn register_definitions(&mut self, definitions: &Value) {
|
||||
if let Some(definitions) = definitions.as_object() {
|
||||
for (name, schema) in definitions {
|
||||
if !self.definitions.contains_key(name) {
|
||||
self.order.push(name.clone());
|
||||
}
|
||||
self.definitions.insert(name.clone(), schema.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_reference(&mut self, reference: &str) -> Option<Value> {
|
||||
if reference == "#" {
|
||||
self.has_self_ref = true;
|
||||
return Some(serde_json::json!({"$self_ref": true}));
|
||||
}
|
||||
if let Some(name) = reference.strip_prefix("#/$defs/")
|
||||
&& let Some(definition) = self.definitions.get(name)
|
||||
{
|
||||
return Some(definition.clone());
|
||||
}
|
||||
self.unsupported = true;
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
enum ParameterType {
|
||||
Scalar(ScalarType),
|
||||
Object(ObjectType),
|
||||
Array(ArrayType),
|
||||
Enum(EnumType),
|
||||
AnyOf(AnyOfType),
|
||||
Union(UnionType),
|
||||
Reference(ReferenceType),
|
||||
}
|
||||
|
||||
impl ParameterType {
|
||||
fn format_docstring(&self, indent: &str) -> String {
|
||||
match self {
|
||||
Self::Scalar(value) => value.base.format_docstring(indent),
|
||||
Self::Object(value) => value.base.format_docstring(indent),
|
||||
Self::Array(value) => value.base.format_docstring(indent),
|
||||
Self::Enum(value) => value.base.format_docstring(indent),
|
||||
Self::AnyOf(value) => value.base.format_docstring(indent),
|
||||
Self::Union(value) => value.base.format_docstring(indent),
|
||||
Self::Reference(value) => value.base.format_docstring(indent),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self, indent: &str, registry: &SchemaRegistry) -> String {
|
||||
match self {
|
||||
Self::Scalar(value) => value.to_typescript(),
|
||||
Self::Object(value) => value.to_typescript(indent, registry),
|
||||
Self::Array(value) => value.to_typescript(indent, registry),
|
||||
Self::Enum(value) => value.to_typescript(),
|
||||
Self::AnyOf(value) => value.to_typescript(indent, registry),
|
||||
Self::Union(value) => value.to_typescript(),
|
||||
Self::Reference(value) => value.to_typescript(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct BaseType {
|
||||
description: String,
|
||||
constraints: Vec<(String, Value)>,
|
||||
}
|
||||
|
||||
impl BaseType {
|
||||
fn new(schema: &Value, allowed_constraints: &[&str]) -> Self {
|
||||
let description = schema
|
||||
.get("description")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.to_owned();
|
||||
let mut constraints = schema
|
||||
.as_object()
|
||||
.map(|object| {
|
||||
object
|
||||
.iter()
|
||||
.filter(|(key, _)| allowed_constraints.contains(&key.as_str()))
|
||||
.map(|(key, value)| (key.clone(), value.clone()))
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
constraints.sort_by(|left, right| left.0.cmp(&right.0));
|
||||
Self {
|
||||
description,
|
||||
constraints,
|
||||
}
|
||||
}
|
||||
|
||||
fn format_docstring(&self, indent: &str) -> String {
|
||||
let mut output = String::new();
|
||||
if !self.description.is_empty() {
|
||||
output.push_str(&format_description(&self.description, indent));
|
||||
output.push('\n');
|
||||
}
|
||||
if !self.constraints.is_empty() {
|
||||
let constraints = self
|
||||
.constraints
|
||||
.iter()
|
||||
.map(|(key, value)| format!("{key}: {}", json_inline(value)))
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ");
|
||||
output.push_str(&format!("{indent}// {constraints}\n"));
|
||||
}
|
||||
output
|
||||
}
|
||||
}
|
||||
|
||||
struct ScalarType {
|
||||
base: BaseType,
|
||||
kind: String,
|
||||
}
|
||||
|
||||
impl ScalarType {
|
||||
fn parse(kind: &str, schema: &Value) -> Self {
|
||||
let constraints = match kind {
|
||||
"string" => &["maxLength", "minLength", "pattern"][..],
|
||||
"number" | "integer" => &["maximum", "minimum"][..],
|
||||
_ => &[],
|
||||
};
|
||||
Self {
|
||||
base: BaseType::new(schema, constraints),
|
||||
kind: kind.to_owned(),
|
||||
}
|
||||
}
|
||||
|
||||
fn any() -> Self {
|
||||
Self {
|
||||
base: BaseType::default(),
|
||||
kind: "any".into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self) -> String {
|
||||
if self.kind == "integer" {
|
||||
"number".into()
|
||||
} else {
|
||||
self.kind.clone()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct Parameter {
|
||||
name: String,
|
||||
kind: ParameterType,
|
||||
optional: bool,
|
||||
default: Option<Value>,
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn to_typescript(&self, indent: &str, registry: &SchemaRegistry) -> String {
|
||||
let mut output = self.kind.format_docstring(indent);
|
||||
if let Some(default) = &self.default {
|
||||
let default = match default {
|
||||
Value::Bool(true) => "True".into(),
|
||||
Value::Bool(false) => "False".into(),
|
||||
Value::Number(_) => default.to_string(),
|
||||
_ => serde_json::to_string(default).unwrap_or_else(|_| "null".into()),
|
||||
};
|
||||
output.push_str(&format!("{indent}// Default: {default}\n"));
|
||||
}
|
||||
let optional = if self.optional { "?" } else { "" };
|
||||
let _ = write!(
|
||||
output,
|
||||
"{indent}{}{optional}: {}",
|
||||
self.name,
|
||||
self.kind.to_typescript(indent, registry)
|
||||
);
|
||||
output
|
||||
}
|
||||
}
|
||||
|
||||
struct ObjectType {
|
||||
base: BaseType,
|
||||
properties: Vec<Parameter>,
|
||||
additional_properties: AdditionalProperties,
|
||||
}
|
||||
|
||||
enum AdditionalProperties {
|
||||
None,
|
||||
True,
|
||||
False,
|
||||
Schema(Box<ParameterType>),
|
||||
}
|
||||
|
||||
impl ObjectType {
|
||||
fn parse(schema: &Value, registry: &mut SchemaRegistry) -> Self {
|
||||
if let Some(definitions) = schema.get("$defs") {
|
||||
registry.register_definitions(definitions);
|
||||
}
|
||||
let additional_properties = match schema.get("additionalProperties") {
|
||||
None => AdditionalProperties::None,
|
||||
Some(Value::Bool(true)) => AdditionalProperties::True,
|
||||
Some(Value::Bool(false)) => AdditionalProperties::False,
|
||||
Some(schema) => AdditionalProperties::Schema(Box::new(parse_type(schema, registry))),
|
||||
};
|
||||
let required = schema
|
||||
.get("required")
|
||||
.and_then(Value::as_array)
|
||||
.map(|values| values.iter().filter_map(Value::as_str).collect::<Vec<_>>())
|
||||
.unwrap_or_default();
|
||||
let properties = schema
|
||||
.get("properties")
|
||||
.and_then(Value::as_object)
|
||||
.map(|properties| {
|
||||
properties
|
||||
.iter()
|
||||
.map(|(name, schema)| Parameter {
|
||||
name: name.clone(),
|
||||
kind: parse_type(schema, registry),
|
||||
optional: !required.contains(&name.as_str()),
|
||||
default: schema
|
||||
.get("default")
|
||||
.filter(|value| !value.is_null())
|
||||
.cloned(),
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
Self {
|
||||
base: BaseType::new(schema, &[]),
|
||||
properties,
|
||||
additional_properties,
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self, indent: &str, registry: &SchemaRegistry) -> String {
|
||||
let mut required = self
|
||||
.properties
|
||||
.iter()
|
||||
.filter(|parameter| !parameter.optional)
|
||||
.collect::<Vec<_>>();
|
||||
let mut optional = self
|
||||
.properties
|
||||
.iter()
|
||||
.filter(|parameter| parameter.optional)
|
||||
.collect::<Vec<_>>();
|
||||
required.sort_by(|left, right| left.name.cmp(&right.name));
|
||||
optional.sort_by(|left, right| left.name.cmp(&right.name));
|
||||
let inner_indent = format!("{indent}{INDENT}");
|
||||
let mut fields = required
|
||||
.into_iter()
|
||||
.chain(optional)
|
||||
.map(|parameter| parameter.to_typescript(&inner_indent, registry))
|
||||
.collect::<Vec<_>>();
|
||||
match &self.additional_properties {
|
||||
AdditionalProperties::None => {}
|
||||
AdditionalProperties::True => fields.push(format!("{inner_indent}[k: string]: any")),
|
||||
AdditionalProperties::False => {
|
||||
fields.push(format!("{inner_indent}[k: string]: never"));
|
||||
}
|
||||
AdditionalProperties::Schema(schema) => fields.push(format!(
|
||||
"{inner_indent}[k: string]: {}",
|
||||
schema.to_typescript(&inner_indent, registry)
|
||||
)),
|
||||
}
|
||||
if fields.is_empty() {
|
||||
"{}".into()
|
||||
} else {
|
||||
format!("{{\n{}\n{indent}}}", fields.join(FIELD_DELIMITER))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct ArrayType {
|
||||
base: BaseType,
|
||||
item: Box<ParameterType>,
|
||||
}
|
||||
|
||||
impl ArrayType {
|
||||
fn parse(schema: &Value, registry: &mut SchemaRegistry) -> Self {
|
||||
let item = schema
|
||||
.get("items")
|
||||
.filter(|item| !item.is_null())
|
||||
.map(|item| parse_type(item, registry))
|
||||
.unwrap_or_else(|| ParameterType::Scalar(ScalarType::any()));
|
||||
Self {
|
||||
base: BaseType::new(schema, &["minItems", "maxItems"]),
|
||||
item: Box::new(item),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self, indent: &str, registry: &SchemaRegistry) -> String {
|
||||
let inner_indent = format!("{indent}{INDENT}");
|
||||
let docstring = self.item.format_docstring(&inner_indent);
|
||||
let item = self.item.to_typescript(&inner_indent, registry);
|
||||
if docstring.is_empty() {
|
||||
format!("Array<{item}>")
|
||||
} else {
|
||||
format!("Array<\n{docstring}{inner_indent}{item}\n{indent}>")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct EnumType {
|
||||
base: BaseType,
|
||||
values: Vec<Value>,
|
||||
}
|
||||
|
||||
impl EnumType {
|
||||
fn parse(schema: &Value) -> Self {
|
||||
Self {
|
||||
base: BaseType::new(schema, &[]),
|
||||
values: schema
|
||||
.get("enum")
|
||||
.and_then(Value::as_array)
|
||||
.cloned()
|
||||
.unwrap_or_default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self) -> String {
|
||||
self.values
|
||||
.iter()
|
||||
.map(|value| match value {
|
||||
Value::String(value) => format!("\"{value}\""),
|
||||
Value::Null => "None".into(),
|
||||
Value::Bool(true) => "True".into(),
|
||||
Value::Bool(false) => "False".into(),
|
||||
value => value.to_string(),
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(" | ")
|
||||
}
|
||||
}
|
||||
|
||||
struct AnyOfType {
|
||||
base: BaseType,
|
||||
branches: Vec<ParameterType>,
|
||||
}
|
||||
|
||||
impl AnyOfType {
|
||||
fn parse(schema: &Value, registry: &mut SchemaRegistry) -> Self {
|
||||
Self {
|
||||
base: BaseType::new(schema, &[]),
|
||||
branches: schema
|
||||
.get("anyOf")
|
||||
.and_then(Value::as_array)
|
||||
.map(|branches| {
|
||||
branches
|
||||
.iter()
|
||||
.map(|branch| parse_type(branch, registry))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self, indent: &str, registry: &SchemaRegistry) -> String {
|
||||
self.branches
|
||||
.iter()
|
||||
.map(|branch| branch.to_typescript(indent, registry))
|
||||
.collect::<Vec<_>>()
|
||||
.join(" | ")
|
||||
}
|
||||
}
|
||||
|
||||
struct UnionType {
|
||||
base: BaseType,
|
||||
kinds: Vec<String>,
|
||||
}
|
||||
|
||||
impl UnionType {
|
||||
fn parse(schema: &Value) -> Self {
|
||||
let kinds = schema
|
||||
.get("type")
|
||||
.and_then(Value::as_array)
|
||||
.map(|kinds| {
|
||||
kinds
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(|kind| match kind {
|
||||
"integer" => "number".into(),
|
||||
"object" => "{}".into(),
|
||||
"array" => "Array<any>".into(),
|
||||
kind => kind.to_owned(),
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
Self {
|
||||
base: BaseType::new(schema, &[]),
|
||||
kinds,
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self) -> String {
|
||||
self.kinds.join(" | ")
|
||||
}
|
||||
}
|
||||
|
||||
struct ReferenceType {
|
||||
base: BaseType,
|
||||
name: String,
|
||||
}
|
||||
|
||||
impl ReferenceType {
|
||||
fn parse(schema: &Value, registry: &mut SchemaRegistry) -> Self {
|
||||
let reference = schema.get("$ref").and_then(Value::as_str).unwrap_or("");
|
||||
let resolved = registry.resolve_reference(reference);
|
||||
let name = match resolved {
|
||||
Some(value) if value.get("$self_ref").and_then(Value::as_bool) == Some(true) => {
|
||||
"parameters".into()
|
||||
}
|
||||
Some(_) => reference.rsplit('/').next().unwrap_or_default().into(),
|
||||
None => "any".into(),
|
||||
};
|
||||
Self {
|
||||
base: BaseType::new(schema, &[]),
|
||||
name,
|
||||
}
|
||||
}
|
||||
|
||||
fn to_typescript(&self) -> String {
|
||||
self.name.clone()
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_type(schema: &Value, registry: &mut SchemaRegistry) -> ParameterType {
|
||||
if registry.depth >= MAX_RECURSION_DEPTH {
|
||||
return ParameterType::Scalar(ScalarType::any());
|
||||
}
|
||||
registry.depth += 1;
|
||||
let result = parse_type_inner(schema, registry);
|
||||
registry.depth -= 1;
|
||||
result
|
||||
}
|
||||
|
||||
fn parse_type_inner(schema: &Value, registry: &mut SchemaRegistry) -> ParameterType {
|
||||
if let Some(schema) = schema.as_bool() {
|
||||
return ParameterType::Scalar(ScalarType {
|
||||
base: BaseType::default(),
|
||||
kind: if schema { "any" } else { "null" }.into(),
|
||||
});
|
||||
}
|
||||
let Some(object) = schema.as_object() else {
|
||||
registry.unsupported = true;
|
||||
return ParameterType::Scalar(ScalarType::any());
|
||||
};
|
||||
if object.contains_key("$ref") {
|
||||
return ParameterType::Reference(ReferenceType::parse(schema, registry));
|
||||
}
|
||||
if object.contains_key("anyOf") {
|
||||
return ParameterType::AnyOf(AnyOfType::parse(schema, registry));
|
||||
}
|
||||
if object.contains_key("enum") {
|
||||
return ParameterType::Enum(EnumType::parse(schema));
|
||||
}
|
||||
if let Some(kind) = object.get("type") {
|
||||
if kind.is_array() {
|
||||
return ParameterType::Union(UnionType::parse(schema));
|
||||
}
|
||||
if let Some(kind) = kind.as_str() {
|
||||
return match kind {
|
||||
"object" => ParameterType::Object(ObjectType::parse(schema, registry)),
|
||||
"array" => ParameterType::Array(ArrayType::parse(schema, registry)),
|
||||
kind => ParameterType::Scalar(ScalarType::parse(kind, schema)),
|
||||
};
|
||||
}
|
||||
}
|
||||
if object.is_empty() {
|
||||
return ParameterType::Scalar(ScalarType::any());
|
||||
}
|
||||
registry.unsupported = true;
|
||||
ParameterType::Scalar(ScalarType::any())
|
||||
}
|
||||
|
||||
fn format_description(description: &str, indent: &str) -> String {
|
||||
description
|
||||
.split('\n')
|
||||
.map(|line| {
|
||||
if line.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!("{indent}// {line}")
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n")
|
||||
}
|
||||
|
||||
fn json_inline(value: &Value) -> String {
|
||||
match value {
|
||||
Value::String(value) => value.clone(),
|
||||
Value::Bool(value) => value.to_string(),
|
||||
Value::Number(value) => value.to_string(),
|
||||
Value::Null => "null".into(),
|
||||
value => serde_json::to_string(value).unwrap_or_default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn recursively_sorts_tool_schema() {
|
||||
let mut value = serde_json::json!({"z": [{"b": 1, "a": 2}], "a": 0});
|
||||
deep_sort(&mut value);
|
||||
assert_eq!(value.to_string(), r#"{"a":0,"z":[{"a":2,"b":1}]}"#);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn encodes_complex_schema_byte_exactly() {
|
||||
let tools = serde_json::json!([{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "weather",
|
||||
"description": "Read weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"units": {"type": "string", "enum": ["c", "f"]},
|
||||
"city": {"type": "string", "description": "City name"}
|
||||
},
|
||||
"required": ["city"]
|
||||
}
|
||||
}
|
||||
}]);
|
||||
assert_eq!(
|
||||
encode_tools_to_typescript(tools.as_array().unwrap()).unwrap(),
|
||||
"# Tools\n\n## functions\nnamespace functions {\n// Read weather\ntype weather = (_: {\n // City name\n city: string,\n units?: \"c\" | \"f\"\n}) => any;\n}\n"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unsupported_schema_uses_json_fallback() {
|
||||
let tools = serde_json::json!([{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "broken",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"value": {"oneOf": [{"type": "string"}]}}
|
||||
}
|
||||
}
|
||||
}]);
|
||||
assert!(encode_tools_to_typescript(tools.as_array().unwrap()).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn null_default_is_omitted_like_checkpoint_python() {
|
||||
let tools = serde_json::json!([{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "optional_value",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {"type": ["string", "null"], "default": null}
|
||||
}
|
||||
}
|
||||
}
|
||||
}]);
|
||||
let encoded = encode_tools_to_typescript(tools.as_array().unwrap()).unwrap();
|
||||
assert_eq!(
|
||||
encoded,
|
||||
"# Tools\n\n## functions\nnamespace functions {\ntype optional_value = (_: {\n value?: string | null\n}) => any;\n}\n"
|
||||
);
|
||||
assert!(!encoded.contains("Default"));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,602 @@
|
||||
//! Tokenizer primitives shared by renderer hosts.
|
||||
|
||||
use crate::{
|
||||
RendererError as Error, RendererLimits, SamplingParams, TextRequest, TokenIds, TokenIdsRequest,
|
||||
};
|
||||
use futures::channel::oneshot;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
enum PoolJob {
|
||||
Tokenize {
|
||||
request: Box<TextRequest>,
|
||||
reply: oneshot::Sender<Result<TokenIdsRequest, Error>>,
|
||||
},
|
||||
Stop,
|
||||
}
|
||||
|
||||
struct TokenizerPoolInner {
|
||||
jobs: flume::Sender<PoolJob>,
|
||||
workers: Mutex<Vec<std::thread::JoinHandle<()>>>,
|
||||
}
|
||||
|
||||
impl Drop for TokenizerPoolInner {
|
||||
fn drop(&mut self) {
|
||||
let workers = self.workers.get_mut().expect("tokenizer workers mutex");
|
||||
for _ in 0..workers.len() {
|
||||
let _ = self.jobs.send(PoolJob::Stop);
|
||||
}
|
||||
for worker in workers.drain(..) {
|
||||
let _ = worker.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Bounded CPU tokenizer pool owned by renderer state.
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PooledTokenizer {
|
||||
inner: Arc<TokenizerPoolInner>,
|
||||
}
|
||||
|
||||
impl PooledTokenizer {
|
||||
pub fn new(
|
||||
tokenizer: Arc<dyn TextTokenizer>,
|
||||
worker_count: usize,
|
||||
queue_capacity: usize,
|
||||
) -> Self {
|
||||
let worker_count = worker_count.max(1);
|
||||
let (jobs, rx) = flume::bounded(queue_capacity.max(1));
|
||||
let mut workers = Vec::with_capacity(worker_count);
|
||||
for index in 0..worker_count {
|
||||
let rx = rx.clone();
|
||||
let tokenizer = tokenizer.clone();
|
||||
workers.push(
|
||||
std::thread::Builder::new()
|
||||
.name(format!("renderer-tokenizer-{index}"))
|
||||
.spawn(move || {
|
||||
while let Ok(job) = rx.recv() {
|
||||
match job {
|
||||
PoolJob::Tokenize { request, reply } => {
|
||||
let result =
|
||||
tokenize_text_request(*request, tokenizer.as_ref());
|
||||
let _ = reply.send(result);
|
||||
}
|
||||
PoolJob::Stop => break,
|
||||
}
|
||||
}
|
||||
})
|
||||
.expect("spawn renderer tokenizer worker"),
|
||||
);
|
||||
}
|
||||
Self {
|
||||
inner: Arc::new(TokenizerPoolInner {
|
||||
jobs,
|
||||
workers: Mutex::new(workers),
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PooledTokenizer {
|
||||
pub(crate) async fn tokenize(&self, request: TextRequest) -> Result<TokenIdsRequest, Error> {
|
||||
let jobs = self.inner.jobs.clone();
|
||||
let (reply, result) = oneshot::channel();
|
||||
jobs.send_async(PoolJob::Tokenize {
|
||||
request: Box::new(request),
|
||||
reply,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
result.await.map_err(|_| Error::WorkerDropped)?
|
||||
}
|
||||
}
|
||||
|
||||
/// Pluggable text→token-ids backend. `Send + Sync` so one instance is shared
|
||||
/// (read-only) across all pinned workers.
|
||||
pub trait TextTokenizer: Send + Sync {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> Result<TokenIds, Error>;
|
||||
|
||||
fn encode_segments(
|
||||
&self,
|
||||
segments: &[dynamo_tokenizers::EncodeSegment<'_>],
|
||||
add_special_tokens: bool,
|
||||
) -> Result<TokenIds, Error> {
|
||||
let text = segments
|
||||
.iter()
|
||||
.map(|segment| segment.text)
|
||||
.collect::<String>();
|
||||
self.encode(&text, add_special_tokens)
|
||||
}
|
||||
}
|
||||
|
||||
/// Load the tokenizer shared (Arc-backed) by the encode pool and detok shards.
|
||||
/// `tokenizer_path` is a tokenizer file, a model dir, or an HF Hub repo id
|
||||
/// (resolved from the local cache — no network).
|
||||
pub fn load_tokenizer(
|
||||
tokenizer_path: Option<&str>,
|
||||
revision: Option<&str>,
|
||||
add_special_tokens: bool,
|
||||
) -> Result<dynamo_tokenizers::Tokenizer, String> {
|
||||
let path =
|
||||
tokenizer_path.ok_or_else(|| "no tokenizer configured: set tokenizer_path".to_string())?;
|
||||
let file = resolve_tokenizer_file(path, revision).ok_or_else(|| {
|
||||
format!(
|
||||
"no supported tokenizer file found for '{path}' (expected tokenizer.json, tiktoken.model, or *.tiktoken)"
|
||||
)
|
||||
})?;
|
||||
let tokenizer = dynamo_tokenizers::Tokenizer::from_file_with_options(
|
||||
&file,
|
||||
dynamo_tokenizers::TokenizerOptions { add_special_tokens },
|
||||
)
|
||||
.map_err(|e| format!("tokenizer load failed ({file}): {e}"))?;
|
||||
tracing::info!(%path, "loaded tokenizer");
|
||||
Ok(tokenizer)
|
||||
}
|
||||
|
||||
/// Resolve the tokenizer source used by the renderer.
|
||||
pub fn resolve_tokenizer_file(path: &str, revision: Option<&str>) -> Option<String> {
|
||||
let input = Path::new(path);
|
||||
if input.is_file() && is_supported_tokenizer_file(input) {
|
||||
return Some(input.to_string_lossy().into_owned());
|
||||
}
|
||||
let directory = model_directory(path, revision)?;
|
||||
discover_tokenizer_in_dir(&directory).map(|path| path.to_string_lossy().into_owned())
|
||||
}
|
||||
|
||||
/// Resolve a dedicated Hugging Face chat-template file when the template is
|
||||
/// not embedded in `tokenizer_config.json`.
|
||||
pub fn resolve_chat_template_file(path: &str, revision: Option<&str>) -> Option<String> {
|
||||
let directory = model_directory(path, revision)?;
|
||||
discover_chat_template_in_dir(&directory).map(|path| path.to_string_lossy().into_owned())
|
||||
}
|
||||
|
||||
fn model_directory(path: &str, revision: Option<&str>) -> Option<PathBuf> {
|
||||
let input = Path::new(path);
|
||||
if input.is_dir() {
|
||||
return Some(input.to_path_buf());
|
||||
}
|
||||
if input.is_file() {
|
||||
return input.parent().map(Path::to_path_buf);
|
||||
}
|
||||
let repo = cache_repo(path, revision);
|
||||
[
|
||||
"config.json",
|
||||
"tokenizer_config.json",
|
||||
"tokenizer.json",
|
||||
"tiktoken.model",
|
||||
]
|
||||
.into_iter()
|
||||
.find_map(|name| repo.get(name))
|
||||
.and_then(|file| file.parent().map(Path::to_path_buf))
|
||||
}
|
||||
|
||||
fn discover_tokenizer_in_dir(directory: &Path) -> Option<PathBuf> {
|
||||
let tokenizer_config = directory.join("tokenizer_config.json");
|
||||
let prefers_tiktoken = std::fs::read_to_string(tokenizer_config)
|
||||
.ok()
|
||||
.and_then(|text| serde_json::from_str::<serde_json::Value>(&text).ok())
|
||||
.and_then(|config| {
|
||||
config
|
||||
.get("tokenizer_class")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(|class| class.to_ascii_lowercase().contains("tiktoken"))
|
||||
})
|
||||
.unwrap_or(false);
|
||||
let hugging_face = directory.join("tokenizer.json");
|
||||
let tiktoken = directory.join("tiktoken.model");
|
||||
let discovered_tiktoken = || {
|
||||
sorted_directory_files(directory).find(|path| {
|
||||
path.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| name.ends_with(".tiktoken"))
|
||||
})
|
||||
};
|
||||
if prefers_tiktoken {
|
||||
tiktoken
|
||||
.is_file()
|
||||
.then_some(tiktoken)
|
||||
.or_else(discovered_tiktoken)
|
||||
.or_else(|| hugging_face.is_file().then_some(hugging_face))
|
||||
} else {
|
||||
hugging_face
|
||||
.is_file()
|
||||
.then_some(hugging_face)
|
||||
.or_else(|| tiktoken.is_file().then_some(tiktoken))
|
||||
.or_else(discovered_tiktoken)
|
||||
}
|
||||
}
|
||||
|
||||
fn discover_chat_template_in_dir(directory: &Path) -> Option<PathBuf> {
|
||||
for name in ["chat_template.json", "chat_template.jinja"] {
|
||||
let candidate = directory.join(name);
|
||||
if candidate.is_file() {
|
||||
return Some(candidate);
|
||||
}
|
||||
}
|
||||
sorted_directory_files(directory).find(|path| {
|
||||
path.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| name.ends_with(".jinja"))
|
||||
})
|
||||
}
|
||||
|
||||
fn sorted_directory_files(directory: &Path) -> impl Iterator<Item = PathBuf> {
|
||||
let mut files = std::fs::read_dir(directory)
|
||||
.ok()
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.flatten()
|
||||
.map(|entry| entry.path())
|
||||
.filter(|path| path.is_file())
|
||||
.collect::<Vec<_>>();
|
||||
files.sort();
|
||||
files.into_iter()
|
||||
}
|
||||
|
||||
fn is_supported_tokenizer_file(path: &Path) -> bool {
|
||||
path.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| {
|
||||
name == "tokenizer.json" || name == "tiktoken.model" || name.ends_with(".tiktoken")
|
||||
})
|
||||
}
|
||||
|
||||
/// Resolve a model file from the tokenizer source: a dir → `dir/<file>`, a file →
|
||||
/// its sibling, else an HF Hub repo id → the local cache. `None` if not found.
|
||||
pub fn resolve_model_file(path: &str, revision: Option<&str>, filename: &str) -> Option<String> {
|
||||
let p = Path::new(path);
|
||||
if p.is_dir() {
|
||||
let f = p.join(filename);
|
||||
return f.is_file().then(|| f.to_string_lossy().into_owned());
|
||||
}
|
||||
if p.is_file() {
|
||||
// `path` is a file (e.g. `tokenizer.json`); look for the sibling.
|
||||
let f = p.parent()?.join(filename);
|
||||
return f.is_file().then(|| f.to_string_lossy().into_owned());
|
||||
}
|
||||
// Not a local path → HF Hub repo id (offline cache lookup).
|
||||
resolve_from_hub_cache(path, revision, filename)
|
||||
}
|
||||
|
||||
/// Locate a file for an HF Hub repo id in the local cache. Offline —
|
||||
/// the scheduler pre-downloads the model. `None` if not cached.
|
||||
fn resolve_from_hub_cache(repo_id: &str, revision: Option<&str>, filename: &str) -> Option<String> {
|
||||
cache_repo(repo_id, revision)
|
||||
.get(filename)
|
||||
.map(|p| p.to_string_lossy().into_owned())
|
||||
}
|
||||
|
||||
fn cache_repo(repo_id: &str, revision: Option<&str>) -> hf_hub::CacheRepo {
|
||||
use hf_hub::{Cache, Repo, RepoType};
|
||||
|
||||
// Python resolves the cache dir as HF_HUB_CACHE > HUGGINGFACE_HUB_CACHE >
|
||||
// HF_HOME/hub > ~/.cache/huggingface/hub; the hf-hub crate only knows
|
||||
// HF_HOME. Honor the explicit cache-dir overrides first, or the Rust
|
||||
// server misses models the Python scheduler already downloaded.
|
||||
let cache = ["HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE"]
|
||||
.iter()
|
||||
.find_map(|var| std::env::var(var).ok())
|
||||
.map(|dir| Cache::new(dir.into()))
|
||||
.unwrap_or_else(Cache::from_env);
|
||||
cache.repo(Repo::with_revision(
|
||||
repo_id.to_string(),
|
||||
RepoType::Model,
|
||||
revision.unwrap_or("main").to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
/// Real tokenizer over two already-loaded dynamo handles. Dynamo fixes
|
||||
/// `add_special_tokens` when loading, so selecting the mode at request time
|
||||
/// requires one handle for each setting.
|
||||
pub struct DynamoTokenizer {
|
||||
without_specials: dynamo_tokenizers::Tokenizer,
|
||||
with_specials: dynamo_tokenizers::Tokenizer,
|
||||
}
|
||||
|
||||
impl DynamoTokenizer {
|
||||
pub fn new(
|
||||
without_specials: dynamo_tokenizers::Tokenizer,
|
||||
with_specials: dynamo_tokenizers::Tokenizer,
|
||||
) -> Self {
|
||||
Self {
|
||||
without_specials,
|
||||
with_specials,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TextTokenizer for DynamoTokenizer {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> Result<TokenIds, Error> {
|
||||
let encoding = if add_special_tokens {
|
||||
&self.with_specials
|
||||
} else {
|
||||
&self.without_specials
|
||||
}
|
||||
.encode(text)
|
||||
.map_err(|e| Error::Tokenize(e.to_string()))?;
|
||||
// Vocab ids are non-negative and fit in i32.
|
||||
Ok(encoding.token_ids().iter().map(|&id| id as i32).collect())
|
||||
}
|
||||
|
||||
fn encode_segments(
|
||||
&self,
|
||||
segments: &[dynamo_tokenizers::EncodeSegment<'_>],
|
||||
add_special_tokens: bool,
|
||||
) -> Result<TokenIds, Error> {
|
||||
let encoding = if add_special_tokens {
|
||||
&self.with_specials
|
||||
} else {
|
||||
&self.without_specials
|
||||
}
|
||||
.encode_segments(segments)
|
||||
.map_err(|error| Error::Tokenize(error.to_string()))?;
|
||||
Ok(encoding.token_ids().iter().map(|&id| id as i32).collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_stop_token_window(sampling_params: &mut SamplingParams, tokenizer: &dyn TextTokenizer) {
|
||||
// Size the scheduler's stop-match window in TOKENS, as Python's
|
||||
// `normalize(tokenizer)` does.
|
||||
if let Some(stop_tokens) = sampling_params
|
||||
.stop_strs
|
||||
.iter()
|
||||
// A stop that won't encode falls back to its byte length rather
|
||||
// than failing the request: still an over-estimate, never an
|
||||
// under-estimate, so the scheduler cannot miss that stop.
|
||||
.map(|stop| {
|
||||
tokenizer
|
||||
.encode(stop, false)
|
||||
.map_or(stop.len(), |ids| ids.len())
|
||||
})
|
||||
.max()
|
||||
{
|
||||
sampling_params.stop_str_max_len = stop_tokens;
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert a text input into the token-ID request consumed by shared
|
||||
/// post-tokenization preparation.
|
||||
pub fn tokenize_text_request(
|
||||
request: TextRequest,
|
||||
tokenizer: &dyn TextTokenizer,
|
||||
) -> Result<TokenIdsRequest, Error> {
|
||||
let TextRequest {
|
||||
rid,
|
||||
prompt,
|
||||
add_special_tokens,
|
||||
mut options,
|
||||
metadata,
|
||||
} = request;
|
||||
resolve_stop_token_window(&mut options.sampling_params, tokenizer);
|
||||
let input_ids = match prompt.encode_segments() {
|
||||
Some(segments) => tokenizer.encode_segments(&segments, add_special_tokens)?,
|
||||
None => tokenizer.encode(prompt.as_str(), add_special_tokens)?,
|
||||
};
|
||||
Ok(TokenIdsRequest {
|
||||
rid,
|
||||
input_ids,
|
||||
options,
|
||||
metadata,
|
||||
})
|
||||
}
|
||||
|
||||
/// Validate fields that must be safe before tokenization or engine submission.
|
||||
pub fn validate_text_request(request: &TextRequest, limits: &RendererLimits) -> Result<(), Error> {
|
||||
validate_request_id(&request.rid)?;
|
||||
if request.prompt.as_str().is_empty() {
|
||||
return Err(Error::Validation("prompt cannot be empty".into()));
|
||||
}
|
||||
let options = &request.options;
|
||||
validate_completion_fields(
|
||||
None,
|
||||
options.token_ids_logprob.as_deref(),
|
||||
options.return_hidden_states,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
/// Validate an already-tokenized request without passing it through the text
|
||||
/// tokenizer path.
|
||||
pub fn validate_token_ids_request(
|
||||
request: &TokenIdsRequest,
|
||||
limits: &RendererLimits,
|
||||
) -> Result<(), Error> {
|
||||
validate_request_id(&request.rid)?;
|
||||
if request.input_ids.is_empty() {
|
||||
return Err(Error::Validation("input_ids cannot be empty".into()));
|
||||
}
|
||||
let options = &request.options;
|
||||
validate_completion_fields(
|
||||
Some(&request.input_ids),
|
||||
options.token_ids_logprob.as_deref(),
|
||||
options.return_hidden_states,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn validate_request_id(rid: &str) -> Result<(), Error> {
|
||||
if rid.len() > 128 {
|
||||
return Err(Error::Validation(format!(
|
||||
"rid is {} bytes, over the 128-byte limit",
|
||||
rid.len()
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Validate the common completion fields before tokenization or engine
|
||||
/// submission. Request identity remains an enclosing host concern.
|
||||
pub fn validate_completion_fields(
|
||||
input_ids: Option<&[i32]>,
|
||||
token_ids_logprob: Option<&[i32]>,
|
||||
return_hidden_states: bool,
|
||||
limits: &RendererLimits,
|
||||
) -> Result<(), Error> {
|
||||
for &id in input_ids.iter().flat_map(|ids| ids.iter()) {
|
||||
if id < 0 || id as u64 >= limits.vocab_size {
|
||||
return Err(Error::Validation(format!(
|
||||
"input_ids contains out-of-vocabulary token id {id}; valid range is [0, {})",
|
||||
limits.vocab_size
|
||||
)));
|
||||
}
|
||||
}
|
||||
for &id in token_ids_logprob.iter().flat_map(|ids| ids.iter()) {
|
||||
if id < 0 || id as u64 >= limits.vocab_size {
|
||||
return Err(Error::Validation(format!(
|
||||
"token_ids_logprob contains out-of-vocabulary token id {id}; valid range is [0, {})",
|
||||
limits.vocab_size
|
||||
)));
|
||||
}
|
||||
}
|
||||
if return_hidden_states && !limits.enable_return_hidden_states {
|
||||
return Err(Error::Validation(
|
||||
"The server is not configured to return the hidden states. Please set `--enable-return-hidden-states` to enable this feature."
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Enforce the model context limit after tokenization.
|
||||
pub fn check_total_tokens(
|
||||
request: &mut TokenIdsRequest,
|
||||
limits: &RendererLimits,
|
||||
) -> Result<(), Error> {
|
||||
let mut input_ids = Some(std::mem::take(&mut request.input_ids));
|
||||
let result =
|
||||
check_completion_token_budget(&mut input_ids, &mut request.options.sampling_params, limits);
|
||||
request.input_ids = input_ids.expect("validated token-ID request retains input_ids");
|
||||
result
|
||||
}
|
||||
|
||||
/// Enforce the context limit over the common token-only completion fields.
|
||||
pub fn check_completion_token_budget(
|
||||
input_ids: &mut Option<TokenIds>,
|
||||
sampling_params: &mut SamplingParams,
|
||||
limits: &RendererLimits,
|
||||
) -> Result<(), Error> {
|
||||
let max_req_len = limits.context_len;
|
||||
let input_len = input_ids.as_ref().map_or(0, Vec::len) as u64 + limits.num_reserved_tokens;
|
||||
if input_len >= max_req_len {
|
||||
if !limits.allow_auto_truncate {
|
||||
return Err(Error::Validation(format!(
|
||||
"The input ({input_len} tokens) is longer than the model's context length ({max_req_len} tokens)."
|
||||
)));
|
||||
}
|
||||
if let Some(ids) = input_ids {
|
||||
ids.truncate(max_req_len as usize);
|
||||
}
|
||||
}
|
||||
let input_len = input_ids.as_ref().map_or(0, Vec::len) as u64 + limits.num_reserved_tokens;
|
||||
let Some(max_new_tokens) = sampling_params.max_new_tokens else {
|
||||
return Ok(());
|
||||
};
|
||||
let total = input_len.saturating_add(max_new_tokens.max(0) as u64);
|
||||
if total <= max_req_len {
|
||||
return Ok(());
|
||||
}
|
||||
if !limits.allow_auto_truncate {
|
||||
return Err(Error::Validation(format!(
|
||||
"Requested token count exceeds the model's maximum context length of {max_req_len} tokens. You requested a total of {total} tokens: {input_len} tokens from the input messages and {max_new_tokens} tokens for the completion. Please reduce the number of tokens in the input messages or the completion to fit within the limit."
|
||||
)));
|
||||
}
|
||||
let clamped = max_req_len.saturating_sub(input_len) as i64;
|
||||
if sampling_params.min_new_tokens > clamped {
|
||||
return Err(Error::Validation(format!(
|
||||
"min_new_tokens must be in [0, max_new_tokens({clamped})], got {}",
|
||||
sampling_params.min_new_tokens
|
||||
)));
|
||||
}
|
||||
sampling_params.max_new_tokens = Some(clamped);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use dynamo_renderer::{RenderedPrompt, RenderedSegment};
|
||||
|
||||
use super::*;
|
||||
use crate::GenerationOptions;
|
||||
|
||||
static NEXT_TEMP_DIR: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
fn temp_model_dir(label: &str) -> PathBuf {
|
||||
let sequence = NEXT_TEMP_DIR.fetch_add(1, Ordering::Relaxed);
|
||||
let path = std::env::temp_dir().join(format!(
|
||||
"sglang-renderer-{label}-{}-{sequence}",
|
||||
std::process::id()
|
||||
));
|
||||
std::fs::create_dir_all(&path).unwrap();
|
||||
path
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_discovery_finds_tiktoken_and_dedicated_chat_template() {
|
||||
let directory = temp_model_dir("model-files");
|
||||
std::fs::write(
|
||||
directory.join("tokenizer_config.json"),
|
||||
r#"{"tokenizer_class":"KimiTikTokenTokenizer"}"#,
|
||||
)
|
||||
.unwrap();
|
||||
std::fs::write(directory.join("tokenizer.json"), "{}").unwrap();
|
||||
std::fs::write(directory.join("tokenizer.tiktoken"), "token").unwrap();
|
||||
std::fs::write(directory.join("chat_template.jinja"), "{{ messages }}").unwrap();
|
||||
|
||||
assert_eq!(
|
||||
resolve_tokenizer_file(directory.to_str().unwrap(), None),
|
||||
Some(
|
||||
directory
|
||||
.join("tokenizer.tiktoken")
|
||||
.to_string_lossy()
|
||||
.into_owned()
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_chat_template_file(directory.to_str().unwrap(), None),
|
||||
Some(
|
||||
directory
|
||||
.join("chat_template.jinja")
|
||||
.to_string_lossy()
|
||||
.into_owned()
|
||||
)
|
||||
);
|
||||
|
||||
std::fs::remove_dir_all(directory).unwrap();
|
||||
}
|
||||
|
||||
struct SegmentTokenizer;
|
||||
|
||||
impl TextTokenizer for SegmentTokenizer {
|
||||
fn encode(&self, _text: &str, _add_special_tokens: bool) -> Result<TokenIds, Error> {
|
||||
Ok(vec![9])
|
||||
}
|
||||
|
||||
fn encode_segments(
|
||||
&self,
|
||||
segments: &[dynamo_tokenizers::EncodeSegment<'_>],
|
||||
_add_special_tokens: bool,
|
||||
) -> Result<TokenIds, Error> {
|
||||
Ok(segments
|
||||
.iter()
|
||||
.map(|segment| if segment.allow_special { 1 } else { 2 })
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rendered_prompt_preserves_segment_boundaries_until_tokenization() {
|
||||
let prompt = RenderedPrompt::segmented(vec![
|
||||
RenderedSegment::new("<control>", true),
|
||||
RenderedSegment::new("user text", false),
|
||||
]);
|
||||
let tokenized = tokenize_text_request(
|
||||
TextRequest::rendered("request", prompt, false, GenerationOptions::default()),
|
||||
&SegmentTokenizer,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(tokenized.input_ids, [1, 2]);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
//! Renderer process state and HTTP listener.
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::{DynamoTokenizer, RendererConfig, RendererService, TextTokenizer, load_tokenizer};
|
||||
|
||||
use crate::engine::{GenerationService, HttpGenerateClient, TokenDecoder};
|
||||
use crate::frontend::http::{hosted_routes, render_only_routes, standalone_routes};
|
||||
use crate::openai::OpenAIService;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct RendererRuntimeConfig {
|
||||
pub http_addr: SocketAddr,
|
||||
pub http_workers: usize,
|
||||
pub tokenizer_workers: usize,
|
||||
pub queue_capacity: usize,
|
||||
/// Optional SGLang engine origin. When absent, inference routes are not mounted.
|
||||
pub engine_url: Option<String>,
|
||||
/// Proxy routes not owned by the renderer to `engine_url`.
|
||||
pub proxy_unhandled_routes: bool,
|
||||
pub renderer: RendererConfig,
|
||||
}
|
||||
|
||||
pub async fn serve(config: RendererRuntimeConfig) -> Result<(), String> {
|
||||
let mode = match (&config.engine_url, config.proxy_unhandled_routes) {
|
||||
(None, false) => "render-only",
|
||||
(Some(_), false) => "serving",
|
||||
(Some(_), true) => "hosted",
|
||||
(None, true) => return Err("proxy_unhandled_routes requires engine_url".to_string()),
|
||||
};
|
||||
let tokenizer_without_specials = load_tokenizer(
|
||||
(!config.renderer.tokenizer_path.is_empty())
|
||||
.then_some(config.renderer.tokenizer_path.as_str()),
|
||||
config.renderer.revision.as_deref(),
|
||||
false,
|
||||
)?;
|
||||
let tokenizer_with_specials = load_tokenizer(
|
||||
(!config.renderer.tokenizer_path.is_empty())
|
||||
.then_some(config.renderer.tokenizer_path.as_str()),
|
||||
config.renderer.revision.as_deref(),
|
||||
true,
|
||||
)?;
|
||||
let encode_tokenizer: Arc<dyn TextTokenizer> = Arc::new(DynamoTokenizer::new(
|
||||
tokenizer_without_specials.clone(),
|
||||
tokenizer_with_specials,
|
||||
));
|
||||
let renderer = Arc::new(RendererService::with_tokenizer(
|
||||
config.renderer,
|
||||
encode_tokenizer,
|
||||
config.tokenizer_workers,
|
||||
config.queue_capacity,
|
||||
));
|
||||
let app = match (config.engine_url, config.proxy_unhandled_routes) {
|
||||
(None, false) => render_only_routes(renderer),
|
||||
(Some(engine_url), false) => {
|
||||
let generate_client = HttpGenerateClient::new(engine_url)?;
|
||||
standalone_routes(
|
||||
OpenAIService::new(
|
||||
renderer,
|
||||
GenerationService::new(
|
||||
Arc::new(generate_client.clone()),
|
||||
TokenDecoder::new(tokenizer_without_specials),
|
||||
),
|
||||
),
|
||||
generate_client,
|
||||
)
|
||||
}
|
||||
(Some(engine_url), true) => {
|
||||
let generate_client = HttpGenerateClient::new(&engine_url)?;
|
||||
hosted_routes(
|
||||
OpenAIService::new(
|
||||
renderer,
|
||||
GenerationService::new(
|
||||
Arc::new(generate_client),
|
||||
TokenDecoder::new(tokenizer_without_specials),
|
||||
),
|
||||
),
|
||||
engine_url,
|
||||
)?
|
||||
}
|
||||
(None, true) => unreachable!("runtime topology was validated above"),
|
||||
};
|
||||
let listener = tokio::net::TcpListener::bind(config.http_addr)
|
||||
.await
|
||||
.map_err(|error| format!("binding renderer on {} failed: {error}", config.http_addr))?;
|
||||
tracing::info!(address = %config.http_addr, mode, "renderer listening");
|
||||
axum::serve(listener, app.into_make_service())
|
||||
.with_graceful_shutdown(async {
|
||||
if let Err(error) = tokio::signal::ctrl_c().await {
|
||||
tracing::error!(%error, "installing renderer shutdown signal failed");
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map_err(|error| format!("renderer HTTP server failed: {error}"))
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//! Renderer request primitives.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
pub type TokenIds = Vec<i32>;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum OneOrMany<T> {
|
||||
One(T),
|
||||
Many(Vec<T>),
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use dynamo_protocols::types::ChatCompletionRequestMessage;
|
||||
use sglang_renderer::{
|
||||
ChatRequest, GenerateRequestMetadata, GenerationOptions, RendererConfig, RendererError,
|
||||
RendererLimits, RendererService, SamplingDefaults, SamplingParams, TextRequest, TextTokenizer,
|
||||
};
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct RecordingTokenizer {
|
||||
prompts: Arc<Mutex<Vec<(String, bool)>>>,
|
||||
}
|
||||
|
||||
impl TextTokenizer for RecordingTokenizer {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> Result<Vec<i32>, RendererError> {
|
||||
self.prompts
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push((text.to_owned(), add_special_tokens));
|
||||
Ok(vec![7])
|
||||
}
|
||||
}
|
||||
|
||||
fn config() -> RendererConfig {
|
||||
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: 128,
|
||||
context_len: 128,
|
||||
num_reserved_tokens: 0,
|
||||
allow_auto_truncate: false,
|
||||
enable_return_hidden_states: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_and_chat_share_the_public_text_preparation_boundary() {
|
||||
let tokenizer = RecordingTokenizer::default();
|
||||
let prompts = tokenizer.prompts.clone();
|
||||
let renderer = RendererService::with_tokenizer(config(), Arc::new(tokenizer), 1, 8);
|
||||
|
||||
let completion = TextRequest::text(
|
||||
"completion-0",
|
||||
"plain completion",
|
||||
true,
|
||||
GenerationOptions {
|
||||
sampling_params: SamplingParams {
|
||||
max_new_tokens: Some(1),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
futures::executor::block_on(renderer.prepare_text_requests(vec![completion])).unwrap();
|
||||
|
||||
let messages: Vec<ChatCompletionRequestMessage> = serde_json::from_value(serde_json::json!([
|
||||
{"role": "user", "content": "hello"}
|
||||
]))
|
||||
.unwrap();
|
||||
let chat = ChatRequest {
|
||||
rid: "chat".into(),
|
||||
model: "model".into(),
|
||||
messages,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
response_format: None,
|
||||
reasoning_effort: None,
|
||||
continue_final_message: false,
|
||||
chat_template_args: None,
|
||||
sampling_params: SamplingParams {
|
||||
max_new_tokens: Some(1),
|
||||
..Default::default()
|
||||
},
|
||||
choice_count: 1,
|
||||
stream: false,
|
||||
return_logprob: false,
|
||||
top_logprobs_num: 0,
|
||||
parallel_tool_calls: true,
|
||||
metadata: GenerateRequestMetadata::default(),
|
||||
};
|
||||
futures::executor::block_on(renderer.prepare_chat(chat)).unwrap();
|
||||
|
||||
let prompts = prompts.lock().unwrap();
|
||||
assert!(
|
||||
prompts
|
||||
.iter()
|
||||
.any(|(text, add_special_tokens)| text == "plain completion" && *add_special_tokens)
|
||||
);
|
||||
assert!(
|
||||
prompts
|
||||
.iter()
|
||||
.any(|(text, add_special_tokens)| text.contains("hello") && !add_special_tokens)
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
// Keep the inner attribute off the first line so shebang lint does not misclassify it.
|
||||
#![cfg(feature = "http")]
|
||||
|
||||
use std::process::{Child, Command, Stdio};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use reqwest::StatusCode;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
struct ChildGuard(Child);
|
||||
|
||||
struct TestDirectory(std::path::PathBuf);
|
||||
|
||||
impl Drop for ChildGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.0.kill();
|
||||
let _ = self.0.wait();
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TestDirectory {
|
||||
fn drop(&mut self) {
|
||||
let _ = std::fs::remove_dir_all(&self.0);
|
||||
}
|
||||
}
|
||||
|
||||
fn available_port() -> u16 {
|
||||
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
listener.local_addr().unwrap().port()
|
||||
}
|
||||
|
||||
async fn wait_until_ready(child: &mut Child, client: &reqwest::Client, health_url: &str) {
|
||||
let deadline = Instant::now() + Duration::from_secs(10);
|
||||
loop {
|
||||
if let Some(status) = child.try_wait().unwrap() {
|
||||
panic!("render-only process exited during startup with {status}");
|
||||
}
|
||||
if client
|
||||
.get(health_url)
|
||||
.send()
|
||||
.await
|
||||
.is_ok_and(|response| response.status() == StatusCode::OK)
|
||||
{
|
||||
return;
|
||||
}
|
||||
assert!(
|
||||
Instant::now() < deadline,
|
||||
"render-only process did not start"
|
||||
);
|
||||
tokio::time::sleep(Duration::from_millis(25)).await;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn binary_starts_without_an_engine_and_serves_only_preprocessing() {
|
||||
let source_tokenizer = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../../experimental/sgl-router/tests/fixtures/tiny_tokenizer.json");
|
||||
let model = TestDirectory(
|
||||
std::env::temp_dir().join(format!("sglang-render-only-{}", uuid::Uuid::new_v4())),
|
||||
);
|
||||
std::fs::create_dir(&model.0).unwrap();
|
||||
std::fs::copy(source_tokenizer, model.0.join("tokenizer.json")).unwrap();
|
||||
let port = available_port();
|
||||
let mut child = ChildGuard(
|
||||
Command::new(env!("CARGO_BIN_EXE_sglang-renderer"))
|
||||
.arg(&model.0)
|
||||
.arg("--tokenizer-path")
|
||||
.arg(&model.0)
|
||||
.arg("--served-model-name")
|
||||
.arg("model")
|
||||
.arg("--resolved-sampling-params")
|
||||
.arg("{}")
|
||||
.arg("--context-length")
|
||||
.arg("64")
|
||||
.arg("--vocab-size")
|
||||
.arg("512")
|
||||
.arg("--host")
|
||||
.arg("127.0.0.1")
|
||||
.arg("--port")
|
||||
.arg(port.to_string())
|
||||
.stdout(Stdio::null())
|
||||
.stderr(Stdio::inherit())
|
||||
.spawn()
|
||||
.unwrap(),
|
||||
);
|
||||
let client = reqwest::Client::builder().no_proxy().build().unwrap();
|
||||
let origin = format!("http://127.0.0.1:{port}");
|
||||
wait_until_ready(&mut child.0, &client, &format!("{origin}/health")).await;
|
||||
|
||||
let tokenized = client
|
||||
.post(format!("{origin}/v1/tokenize"))
|
||||
.json(&json!({"prompt": "hello"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(tokenized.status(), StatusCode::OK);
|
||||
let tokenized: Value = tokenized.json().await.unwrap();
|
||||
assert!(tokenized["count"].as_u64().is_some_and(|count| count > 0));
|
||||
|
||||
let completion = json!({"model": "model", "prompt": "hello"});
|
||||
let rendered = client
|
||||
.post(format!("{origin}/v1/completions/render"))
|
||||
.json(&completion)
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(rendered.status(), StatusCode::OK);
|
||||
let rendered: Value = rendered.json().await.unwrap();
|
||||
assert!(
|
||||
rendered
|
||||
.as_array()
|
||||
.is_some_and(|requests| requests.len() == 1)
|
||||
);
|
||||
|
||||
let inference = client
|
||||
.post(format!("{origin}/v1/completions"))
|
||||
.json(&completion)
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(inference.status(), StatusCode::NOT_FOUND);
|
||||
}
|
||||
Reference in New Issue
Block a user