[rust-renderer] Standalone preprocessing (#36718)
Signed-off-by: Sage Ahrac <sagiahrak@gmail.com> Co-authored-by: Shangming Cai <csmthu@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com> Co-authored-by: Rain Jiang <96632942+rainj-me@users.noreply.github.com>
This commit is contained in:
co-authored by
Shangming Cai
Liangsheng Yin
Rain Jiang
parent
6880a47955
commit
7b1c2ed0a4
@@ -0,0 +1,89 @@
|
|||||||
|
# syntax=docker/dockerfile:1
|
||||||
|
|
||||||
|
# Keep the compiler aligned with rust/rust-toolchain.toml. Pin image indexes
|
||||||
|
# rather than individual architecture manifests so both platforms use this file.
|
||||||
|
FROM rust:1.92.0-slim-bookworm@sha256:f1f73538ebe623fd3673a35aff3df358ae1084c64c55646516e5b17b321b6c9b AS build
|
||||||
|
|
||||||
|
ARG TARGETARCH
|
||||||
|
ARG CARGO_BUILD_JOBS=4
|
||||||
|
ENV RUSTUP_TOOLCHAIN=1.92.0 \
|
||||||
|
CARGO_BUILD_JOBS=${CARGO_BUILD_JOBS} \
|
||||||
|
PCRE2_SYS_STATIC=1
|
||||||
|
|
||||||
|
RUN apt-get update \
|
||||||
|
&& apt-get install -y --no-install-recommends build-essential pkg-config \
|
||||||
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
|
WORKDIR /build
|
||||||
|
COPY rust/Cargo.toml rust/Cargo.lock rust/rust-toolchain.toml rust/
|
||||||
|
# Cargo loads every workspace member even when building only the renderer.
|
||||||
|
COPY rust/sglang-grpc/Cargo.toml rust/sglang-grpc/
|
||||||
|
COPY rust/sglang-grpc/src/ rust/sglang-grpc/src/
|
||||||
|
COPY rust/sglang-mm/Cargo.toml rust/sglang-mm/
|
||||||
|
COPY rust/sglang-mm/src/ rust/sglang-mm/src/
|
||||||
|
COPY rust/sglang-server/Cargo.toml rust/sglang-server/
|
||||||
|
COPY rust/sglang-server/src/ rust/sglang-server/src/
|
||||||
|
COPY rust/sglang-renderer/Cargo.toml rust/sglang-renderer/
|
||||||
|
COPY rust/sglang-renderer/src/ rust/sglang-renderer/src/
|
||||||
|
|
||||||
|
# Avoid rustup downloading development components from rust-toolchain.toml,
|
||||||
|
# but fail if the image's compiler and the workspace toolchain drift apart.
|
||||||
|
RUN channel=$(sed -n 's/^channel = "\([^"]*\)"/\1/p' rust/rust-toolchain.toml) \
|
||||||
|
&& case "${RUSTUP_TOOLCHAIN}" in "$channel"|"$channel".*) ;; *) exit 1 ;; esac
|
||||||
|
|
||||||
|
RUN --mount=type=cache,id=renderer-registry-${TARGETARCH},target=/usr/local/cargo/registry,sharing=locked \
|
||||||
|
--mount=type=cache,id=renderer-git-${TARGETARCH},target=/usr/local/cargo/git,sharing=locked \
|
||||||
|
--mount=type=cache,id=renderer-target-${TARGETARCH},target=/build/rust/target,sharing=locked \
|
||||||
|
cargo build --manifest-path rust/Cargo.toml -p sglang-renderer \
|
||||||
|
--bin sglang-renderer --release --features http --locked \
|
||||||
|
&& install -D rust/target/release/sglang-renderer /out/sglang-renderer
|
||||||
|
|
||||||
|
# Run the existing unit suite in the same Linux toolchain used for the image.
|
||||||
|
# This sibling stage is selected by CI and is not a dependency of the runtime.
|
||||||
|
FROM build AS test
|
||||||
|
COPY rust/sglang-renderer/tests/ rust/sglang-renderer/tests/
|
||||||
|
COPY experimental/sgl-router/tests/fixtures/tiny_tokenizer.json experimental/sgl-router/tests/fixtures/tiny_tokenizer.json
|
||||||
|
RUN --mount=type=cache,id=renderer-registry-${TARGETARCH},target=/usr/local/cargo/registry,sharing=locked \
|
||||||
|
--mount=type=cache,id=renderer-git-${TARGETARCH},target=/usr/local/cargo/git,sharing=locked \
|
||||||
|
--mount=type=cache,id=renderer-target-${TARGETARCH},target=/build/rust/target,sharing=locked \
|
||||||
|
cargo test --manifest-path rust/Cargo.toml -p sglang-renderer --features http --locked
|
||||||
|
|
||||||
|
FROM debian:bookworm-slim@sha256:88200866dfff7ea7f5cbcb6ec7c8a701889efe6fe859fe64d6990e4b07ea4171 AS runtime
|
||||||
|
|
||||||
|
RUN apt-get update \
|
||||||
|
&& apt-get install -y --no-install-recommends ca-certificates libgcc-s1 \
|
||||||
|
&& rm -rf /var/lib/apt/lists/* \
|
||||||
|
&& groupadd --gid 65532 sglang \
|
||||||
|
&& useradd --uid 65532 --gid 65532 --no-log-init --create-home \
|
||||||
|
--home-dir /home/sglang --shell /usr/sbin/nologin sglang \
|
||||||
|
&& mkdir -p /home/sglang/.cache/huggingface \
|
||||||
|
&& chown -R 65532:65532 /home/sglang
|
||||||
|
|
||||||
|
COPY --from=build /out/sglang-renderer /usr/local/bin/sglang-renderer
|
||||||
|
COPY LICENSE /usr/share/licenses/sglang-renderer/LICENSE
|
||||||
|
|
||||||
|
# Metadata changes must not invalidate compilation.
|
||||||
|
ARG SGLANG_BUILD_COMMIT=unknown
|
||||||
|
ARG SGLANG_BUILD_URL=
|
||||||
|
ARG SGLANG_IMAGE_TAG=local/sglang-renderer:dev
|
||||||
|
ENV HOME=/home/sglang \
|
||||||
|
HF_HOME=/home/sglang/.cache/huggingface \
|
||||||
|
SGLANG_BUILD_COMMIT=${SGLANG_BUILD_COMMIT} \
|
||||||
|
SGLANG_BUILD_URL=${SGLANG_BUILD_URL} \
|
||||||
|
SGLANG_IMAGE_TAG=${SGLANG_IMAGE_TAG}
|
||||||
|
LABEL org.opencontainers.image.source="https://github.com/sgl-project/sglang" \
|
||||||
|
org.opencontainers.image.licenses="Apache-2.0" \
|
||||||
|
org.opencontainers.image.revision="${SGLANG_BUILD_COMMIT}" \
|
||||||
|
org.opencontainers.image.version="${SGLANG_IMAGE_TAG}" \
|
||||||
|
org.opencontainers.image.url="${SGLANG_BUILD_URL}" \
|
||||||
|
ai.sglang.build.commit="${SGLANG_BUILD_COMMIT}" \
|
||||||
|
ai.sglang.build.url="${SGLANG_BUILD_URL}" \
|
||||||
|
ai.sglang.image.tag="${SGLANG_IMAGE_TAG}"
|
||||||
|
|
||||||
|
USER 65532:65532
|
||||||
|
WORKDIR /home/sglang
|
||||||
|
EXPOSE 30000
|
||||||
|
# The renderer's existing graceful shutdown handler listens for Ctrl-C.
|
||||||
|
STOPSIGNAL SIGINT
|
||||||
|
ENTRYPOINT ["/usr/local/bin/sglang-renderer"]
|
||||||
|
CMD ["--help"]
|
||||||
Generated
+88
@@ -703,6 +703,16 @@ version = "0.4.2"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "3d52eff69cd5e647efe296129160853a42795992097e8af39800e1060caeea9b"
|
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]]
|
[[package]]
|
||||||
name = "core-foundation-sys"
|
name = "core-foundation-sys"
|
||||||
version = "0.8.7"
|
version = "0.8.7"
|
||||||
@@ -1076,6 +1086,15 @@ version = "1.0.0"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "34aa73646ffb006b8f5147f3dc182bd4bcb190227ce861fc4a4844bf8e3cb2c0"
|
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]]
|
[[package]]
|
||||||
name = "equator"
|
name = "equator"
|
||||||
version = "0.4.2"
|
version = "0.4.2"
|
||||||
@@ -1510,15 +1529,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97"
|
checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"dirs",
|
"dirs",
|
||||||
|
"futures",
|
||||||
"http",
|
"http",
|
||||||
"indicatif",
|
"indicatif",
|
||||||
"libc",
|
"libc",
|
||||||
"log",
|
"log",
|
||||||
|
"num_cpus",
|
||||||
"rand 0.9.5",
|
"rand 0.9.5",
|
||||||
"reqwest",
|
"reqwest",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
"thiserror",
|
"thiserror",
|
||||||
|
"tokio",
|
||||||
"ureq",
|
"ureq",
|
||||||
"windows-sys 0.60.2",
|
"windows-sys 0.60.2",
|
||||||
]
|
]
|
||||||
@@ -1637,9 +1659,11 @@ dependencies = [
|
|||||||
"percent-encoding",
|
"percent-encoding",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"socket2 0.5.10",
|
"socket2 0.5.10",
|
||||||
|
"system-configuration",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
"tracing",
|
"tracing",
|
||||||
|
"windows-registry",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -3293,9 +3317,11 @@ checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"base64 0.22.1",
|
"base64 0.22.1",
|
||||||
"bytes",
|
"bytes",
|
||||||
|
"encoding_rs",
|
||||||
"futures-channel",
|
"futures-channel",
|
||||||
"futures-core",
|
"futures-core",
|
||||||
"futures-util",
|
"futures-util",
|
||||||
|
"h2",
|
||||||
"http",
|
"http",
|
||||||
"http-body",
|
"http-body",
|
||||||
"http-body-util",
|
"http-body-util",
|
||||||
@@ -3304,6 +3330,7 @@ dependencies = [
|
|||||||
"hyper-util",
|
"hyper-util",
|
||||||
"js-sys",
|
"js-sys",
|
||||||
"log",
|
"log",
|
||||||
|
"mime",
|
||||||
"mime_guess",
|
"mime_guess",
|
||||||
"percent-encoding",
|
"percent-encoding",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
@@ -3696,6 +3723,35 @@ dependencies = [
|
|||||||
"ureq",
|
"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]]
|
[[package]]
|
||||||
name = "sglang-server"
|
name = "sglang-server"
|
||||||
version = "0.1.0"
|
version = "0.1.0"
|
||||||
@@ -3959,6 +4015,27 @@ dependencies = [
|
|||||||
"syn 2.0.119",
|
"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]]
|
[[package]]
|
||||||
name = "tagptr"
|
name = "tagptr"
|
||||||
version = "0.2.0"
|
version = "0.2.0"
|
||||||
@@ -4752,6 +4829,17 @@ version = "0.2.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
|
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]]
|
[[package]]
|
||||||
name = "windows-result"
|
name = "windows-result"
|
||||||
version = "0.4.1"
|
version = "0.4.1"
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ resolver = "3"
|
|||||||
members = [
|
members = [
|
||||||
"sglang-grpc",
|
"sglang-grpc",
|
||||||
"sglang-mm",
|
"sglang-mm",
|
||||||
|
"sglang-renderer",
|
||||||
"sglang-server"
|
"sglang-server"
|
||||||
]
|
]
|
||||||
exclude = ["sglang-radix-tree"]
|
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