diff --git a/examples/README.md b/examples/README.md index fa6ce0574c..e3e808360e 100644 --- a/examples/README.md +++ b/examples/README.md @@ -21,6 +21,7 @@ Functions are long-running services that respond to HTTP or gRPC invocations. | [Ray Serve Helm Chart](function-samples/helmchart-samples/ray-serve-sample/) | Helm chart that deploys a Ray Serve application as an NVCF function. | | [Dynamo Operator Sample](function-samples/helmchart-samples/dynamo-operator-sample/) | Helm chart for a vLLM disaggregated router deployed through NVCF. | | [Load Tester Supreme](function-samples/load-tester-supreme/) | HTTP and gRPC echo servers designed for load and throughput testing. | +| [OpenAI-compatible](function-samples/openai-compatible-sample/) | Controllable OpenAI-compatible LLM endpoint target for SDK and load testing. | | [gRPC Streaming ASR Client](function-samples/grpc-streaming-asr-client/) | Invoke a Nemotron ASR Streaming NIM over the NVCF gRPC gateway using bidirectional streaming. | ## Task Samples diff --git a/examples/function-samples/load-tester-supreme/Dockerfile b/examples/function-samples/load-tester-supreme/Dockerfile index 2ab831fb80..536258d1b7 100644 --- a/examples/function-samples/load-tester-supreme/Dockerfile +++ b/examples/function-samples/load-tester-supreme/Dockerfile @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -FROM golang:1.22 AS http-server-build +FROM golang:1.23 AS http-server-build WORKDIR /go/src/http-app COPY http-server . @@ -43,4 +43,4 @@ COPY grpc-server/grpc_echo_server.py /app/ ENV GIN_MODE=release COPY --from=http-server-build /go/bin/http-app /app/ -CMD python3 /app/grpc_echo_server.py & /app/http-app \ No newline at end of file +CMD python3 /app/grpc_echo_server.py & /app/http-app diff --git a/examples/function-samples/load-tester-supreme/README.md b/examples/function-samples/load-tester-supreme/README.md index b24ee15377..f41b5a96c5 100644 --- a/examples/function-samples/load-tester-supreme/README.md +++ b/examples/function-samples/load-tester-supreme/README.md @@ -2,7 +2,7 @@ A dual-protocol echo server (HTTP + gRPC) purpose-built for load and throughput testing of NVCF deployments. Use this instead of the simpler `grpc-echo-sample` -for any real load testing — it ships with 500 gRPC worker threads (vs 10) and +for real load testing. It ships with 500 gRPC worker threads (vs 10) and exposes tunable response behaviour. ## What's included @@ -19,7 +19,7 @@ exposes tunable response behaviour. | Field | Type | Default | Description | |-------|------|---------|-------------| -| `message` | string | — | Payload content to echo back | +| `message` | string | Required for HTTP | Payload content to echo back. gRPC accepts an omitted value and returns an empty string. | | `repeats` | int | 1 | Number of times to repeat the response | | `delay` | float | ~0 | Seconds to sleep between responses | | `size` | int | 0 | Generate a random string of this length instead of echoing `message` (HTTP only) | diff --git a/examples/function-samples/openai-compatible-sample/Dockerfile b/examples/function-samples/openai-compatible-sample/Dockerfile new file mode 100644 index 0000000000..4ebbefa924 --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/Dockerfile @@ -0,0 +1,33 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +FROM golang:1.23 AS build + +WORKDIR /src +COPY http-server/go.mod ./ +RUN go mod download +COPY http-server/ ./ +RUN CGO_ENABLED=0 go build -o /openai-compatible-sample . + +FROM alpine:3.20 + +COPY --from=build /openai-compatible-sample /app/openai-compatible-sample +RUN adduser -D -H app && chown app:app /app/openai-compatible-sample + +EXPOSE 8000 + +USER app + +CMD ["/app/openai-compatible-sample"] diff --git a/examples/function-samples/openai-compatible-sample/README.md b/examples/function-samples/openai-compatible-sample/README.md new file mode 100644 index 0000000000..1006a03bde --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/README.md @@ -0,0 +1,204 @@ +# OpenAI-compatible sample + +A controllable HTTP target for testing NVCF LLM functions with OpenAI client +libraries and load generators. It returns synthetic text and deterministic +embedding vectors. It is not a model server and does not implement the full +OpenAI API. + +## Supported routes + +| Route | Behavior | +|-------|----------| +| `POST /v1/chat/completions` | Chat completions with JSON or SSE output | +| `POST /v1/completions` | Legacy text completions with JSON or SSE output | +| `POST /v1/responses` | Responses API JSON or SSE output | +| `POST /v1/embeddings` | Embeddings for one string or an array of strings | +| `GET /v1/models` | Lists the sample model | +| `GET /v1/models/{id}` | Returns metadata for a model ID | +| `GET /health` | Health check returning 200 | + +The sample accepts normal OpenAI request envelopes. It reads the standard +fields needed to select streaming and embedding output, but does not inspect +prompt or message content for benchmark tuning. Images, files, tools, +multimodal input, stored responses, retrieval, and token-array embeddings are +not implemented. + +## Benchmark controls + +Benchmark controls use either HTTP headers or top-level JSON body fields. Body +controls may appear anywhere in the JSON object. The normal OpenAI fields, +including `model`, `stream`, `input`, `messages`, and `prompt`, are still +decoded normally. + +If any request header starts with `X-Load-Tester-`, all body controls are +ignored. Header controls and defaults apply instead. The default output chunk +is `xxxx`. Text controls apply to Chat Completions, Responses, and legacy +Completions. Queue delay, TTFT, status injection, and concurrency limits apply +to all POST routes. + +| Control | Header | JSON body field | Default | Behavior | +|---------|--------|-----------------|---------|----------| +| Queue delay | `X-Load-Tester-Queue-Delay-Ms` | `x_load_tester_queue_delay_ms` | `0` | Delay before processing the request. | +| TTFT | `X-Load-Tester-TTFT-Ms` | `x_load_tester_ttft_ms` | `0` | Delay before the first response byte. | +| TTFT jitter | `X-Load-Tester-TTFT-Jitter-Ms` | `x_load_tester_ttft_jitter_ms` | `0` | Random extra delay from 0 through this value. | +| ITL | `X-Load-Tester-ITL-Ms` | `x_load_tester_itl_ms` | `0` | Delay between streamed output chunks. | +| ITL jitter | `X-Load-Tester-ITL-Jitter-Ms` | `x_load_tester_itl_jitter_ms` | `0` | Random extra delay between chunks. | +| Chunk text | `X-Load-Tester-Chunk` | `x_load_tester_chunk` | `xxxx` | Text returned in each output chunk. | +| Chunk bytes | `X-Load-Tester-Chunk-Bytes` | `x_load_tester_chunk_bytes` | `0` | Generate a random chunk of this byte length. | +| Output chunks | `X-Load-Tester-Output-Chunks` | `x_load_tester_output_chunks` | `1` | Number of text chunks to return, capped by the startup limit. | +| Status injection | `X-Load-Tester-Status-Code` | `x_load_tester_status_code` | unset | Return an OpenAI-shaped HTTP error. | +| Stream error | `X-Load-Tester-Stream-Error-After-Chunks` | `x_load_tester_stream_error_after_chunks` | unset | End a stream with an OpenAI-shaped error after this many chunks. | +| Stream truncate | `X-Load-Tester-Stream-Truncate-After-Chunks` | `x_load_tester_stream_truncate_after_chunks` | unset | Close a stream without its completion event after this many chunks. | +| Concurrency limit | `X-Load-Tester-Max-Concurrency` | `x_load_tester_max_concurrency` | `0` | Return 429 when the global in-flight request count exceeds this value. | + +Timing values are non-negative integer milliseconds and are capped at five +minutes. Output is capped at 1 MiB and the startup chunk limit. The +`LOAD_TESTER_MAX_OUTPUT_CHUNKS` environment variable defaults to 6000 and +accepts values from 1 through 60000. It is read once when the server starts. +`Chunk` and `Chunk-Bytes` cannot be combined. Stream error and truncate controls +are mutually exclusive. +For a deterministic concurrency test, send the same concurrency limit header +on every request in the load test. + +## Build and run + +```bash +cd examples/function-samples/openai-compatible-sample +docker build --platform linux/amd64 -t openai-compatible-sample . +docker run --rm -p 18000:8000 openai-compatible-sample +``` + +## Smoke test + +```bash +curl --request POST \ + --url http://localhost:18000/v1/chat/completions \ + --header 'Content-Type: application/json' \ + --header 'X-Load-Tester-Chunk: token' \ + --header 'X-Load-Tester-Output-Chunks: 3' \ + --data '{ + "model": "test-model", + "messages": [{"role": "user", "content": "hello"}] + }' + +curl --request POST \ + --url http://localhost:18000/v1/responses \ + --header 'Content-Type: application/json' \ + --data '{ + "model": "test-model", + "input": "hello", + "stream": true, + "x_load_tester_ttft_ms": 200, + "x_load_tester_itl_ms": 50, + "x_load_tester_chunk": "token", + "x_load_tester_output_chunks": 3 + }' + +curl --request POST \ + --url http://localhost:18000/v1/responses \ + --header 'Content-Type: application/json' \ + --header 'X-Load-Tester-TTFT-Ms: 200' \ + --header 'X-Load-Tester-ITL-Ms: 50' \ + --header 'X-Load-Tester-Chunk: token' \ + --header 'X-Load-Tester-Output-Chunks: 3' \ + --data '{ + "model": "test-model", + "input": "hello", + "stream": true + }' + +curl --request POST \ + --url http://localhost:18000/v1/embeddings \ + --header 'Content-Type: application/json' \ + --data '{"model": "test-model", "input": ["one", "two"], "encoding_format": "float"}' +``` + +## OpenAI Python client + +Use a normal OpenAI client with the sample URL as its `base_url`: + +```python +from openai import OpenAI + +client = OpenAI( + api_key="not-needed", + base_url="http://localhost:18000/v1", + _strict_response_validation=True, +) + +response = client.responses.create( + model="test-model", + input="hello", + extra_headers={ + "X-Load-Tester-Chunk": "token", + "X-Load-Tester-Output-Chunks": "3", + }, +) +assert response.output_text == "tokentokentoken" + +chat = client.chat.completions.create( + model="test-model", + messages=[{"role": "user", "content": "hello"}], +) +assert chat.choices[0].message.content == "xxxx" + +embedding = client.embeddings.create( + model="test-model", + input=["one", "two"], + encoding_format="float", +) +assert len(embedding.data) == 2 +``` + +Run the included client compatibility check after starting the Go server: + +```bash +cd examples/function-samples/openai-compatible-sample/http-server +go run . + +# In another terminal, with the openai package installed: +python3 openai_client_check.py +``` + +## 60-second SSE capacity run + +Use the matching pinned xk6 binary from `examples/load-tests`. The command +opens two approximately 60-second streams per VU at 5 ms ITL. Start the sample +with `LOAD_TESTER_MAX_OUTPUT_CHUNKS=12000` so it accepts the requested output +shape. Raise the generator file-descriptor limit before using high concurrency. + +```bash +cd examples/load-tests +ulimit -n 65536 + +./k6 run functions/oai_compatible_responses_sse_load_test.js \ + -e OAI_COMPAT_URL=$OAI_COMPAT_URL \ + -e OPENAI_RESPONSES_PROFILE=calibration \ + -e OPENAI_RESPONSES_VUS=1024 \ + -e OPENAI_RESPONSES_ITERATIONS=2 \ + -e OPENAI_RESPONSES_MAX_DURATION=5m \ + -e OPENAI_RESPONSES_EXPECTED_DELTAS=12000 \ + -e OPENAI_RESPONSES_CALIBRATION_TOLERANCE_MS=1 \ + -e LOAD_TESTER_QUEUE_DELAY_MS=0 \ + -e LOAD_TESTER_TTFT_MS=1 \ + -e LOAD_TESTER_TTFT_JITTER_MS=0 \ + -e LOAD_TESTER_ITL_MS=5 \ + -e LOAD_TESTER_ITL_JITTER_MS=0 \ + -e LOAD_TESTER_CHUNK=xxxx \ + -e LOAD_TESTER_OUTPUT_CHUNKS=12000 +``` + +## NVCF LLM functions + +Expose port 8000 and configure the function inference URL as `/`. Declare the +routes required by the workload, such as: + +```text +/v1/chat/completions +/v1/completions +/v1/responses +/v1/embeddings +``` + +See the [LLM Gateway guide](../../../docs/user/llm-gateway.md) for function +model configuration and invocation flow. diff --git a/examples/function-samples/openai-compatible-sample/http-server/go.mod b/examples/function-samples/openai-compatible-sample/http-server/go.mod new file mode 100644 index 0000000000..f086d6bc58 --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/http-server/go.mod @@ -0,0 +1,3 @@ +module nvcf-openai-compatible-sample + +go 1.23.0 diff --git a/examples/function-samples/openai-compatible-sample/http-server/main.go b/examples/function-samples/openai-compatible-sample/http-server/main.go new file mode 100644 index 0000000000..72faa799c4 --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/http-server/main.go @@ -0,0 +1,1440 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package main + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "math" + "math/rand" + "net/http" + "os" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" +) + +const ( + maxRequestBytes = 10 << 20 + maxOutputBytes = 1 << 20 + defaultMaxOutputChunks = 6000 + hardMaxOutputChunks = 60000 + maxEmbeddingItems = 2048 + maxControlMilliseconds = 5 * 60 * 1000 + maxConcurrencyLimit = 100000 + defaultChunk = "xxxx" + defaultModel = "test-model" + maxOutputChunksEnv = "LOAD_TESTER_MAX_OUTPUT_CHUNKS" + + headerQueueDelay = "X-Load-Tester-Queue-Delay-Ms" + headerTTFT = "X-Load-Tester-TTFT-Ms" + headerTTFTJitter = "X-Load-Tester-TTFT-Jitter-Ms" + headerITL = "X-Load-Tester-ITL-Ms" + headerITLJitter = "X-Load-Tester-ITL-Jitter-Ms" + headerChunk = "X-Load-Tester-Chunk" + headerChunkBytes = "X-Load-Tester-Chunk-Bytes" + headerOutputChunks = "X-Load-Tester-Output-Chunks" + headerStatusCode = "X-Load-Tester-Status-Code" + headerStreamErrorAfter = "X-Load-Tester-Stream-Error-After-Chunks" + headerStreamTruncateAfter = "X-Load-Tester-Stream-Truncate-After-Chunks" + headerMaxConcurrency = "X-Load-Tester-Max-Concurrency" + + bodyQueueDelay = "x_load_tester_queue_delay_ms" + bodyTTFT = "x_load_tester_ttft_ms" + bodyTTFTJitter = "x_load_tester_ttft_jitter_ms" + bodyITL = "x_load_tester_itl_ms" + bodyITLJitter = "x_load_tester_itl_jitter_ms" + bodyChunk = "x_load_tester_chunk" + bodyChunkBytes = "x_load_tester_chunk_bytes" + bodyOutputChunks = "x_load_tester_output_chunks" + bodyStatusCode = "x_load_tester_status_code" + bodyStreamErrorAfter = "x_load_tester_stream_error_after_chunks" + bodyStreamTruncateAfter = "x_load_tester_stream_truncate_after_chunks" + bodyMaxConcurrency = "x_load_tester_max_concurrency" + + loadTesterHeaderPrefix = "x-load-tester-" +) + +var ( + responseSequence atomic.Uint64 + activeRequests atomic.Int64 + randomSource = rand.New(rand.NewSource(time.Now().UnixNano())) + randomSourceMu sync.Mutex +) + +type benchmarkTuning struct { + QueueDelay time.Duration + TTFT time.Duration + TTFTJitter time.Duration + ITL time.Duration + ITLJitter time.Duration + Chunk string + ChunkBytes int + OutputChunks int + StatusCode int + StreamErrorAfter int + StreamTruncateAfter int + MaxConcurrency int +} + +type serverConfig struct { + maxOutputChunks int +} + +type benchmarkBodyControls struct { + QueueDelay json.RawMessage `json:"x_load_tester_queue_delay_ms"` + TTFT json.RawMessage `json:"x_load_tester_ttft_ms"` + TTFTJitter json.RawMessage `json:"x_load_tester_ttft_jitter_ms"` + ITL json.RawMessage `json:"x_load_tester_itl_ms"` + ITLJitter json.RawMessage `json:"x_load_tester_itl_jitter_ms"` + Chunk json.RawMessage `json:"x_load_tester_chunk"` + ChunkBytes json.RawMessage `json:"x_load_tester_chunk_bytes"` + OutputChunks json.RawMessage `json:"x_load_tester_output_chunks"` + StatusCode json.RawMessage `json:"x_load_tester_status_code"` + StreamErrorAfter json.RawMessage `json:"x_load_tester_stream_error_after_chunks"` + StreamTruncateAfter json.RawMessage `json:"x_load_tester_stream_truncate_after_chunks"` + MaxConcurrency json.RawMessage `json:"x_load_tester_max_concurrency"` +} + +type responsesRequest struct { + Model string `json:"model"` + Stream bool `json:"stream"` + Input json.RawMessage `json:"input"` + benchmarkBodyControls +} + +type chatCompletionsRequest struct { + Model string `json:"model"` + Stream bool `json:"stream"` + StreamOptions *chatStreamOptions `json:"stream_options"` + Messages json.RawMessage `json:"messages"` + benchmarkBodyControls +} + +type completionsRequest struct { + Model string `json:"model"` + Stream bool `json:"stream"` + Prompt json.RawMessage `json:"prompt"` + benchmarkBodyControls +} + +type chatStreamOptions struct { + IncludeUsage bool `json:"include_usage"` +} + +type embeddingsRequest struct { + Model string `json:"model"` + Input json.RawMessage `json:"input"` + EncodingFormat string `json:"encoding_format"` + benchmarkBodyControls +} + +type responsesResponse struct { + ID string `json:"id"` + Object string `json:"object"` + CreatedAt int64 `json:"created_at"` + CompletedAt *int64 `json:"completed_at,omitempty"` + Status string `json:"status"` + Model string `json:"model"` + Output []responseOutputItem `json:"output"` + ParallelToolCalls bool `json:"parallel_tool_calls"` + ToolChoice string `json:"tool_choice"` + Tools []any `json:"tools"` + Usage *responsesUsage `json:"usage"` +} + +type responseOutputItem struct { + ID string `json:"id"` + Type string `json:"type"` + Status string `json:"status"` + Role string `json:"role"` + Content []responseOutputText `json:"content"` +} + +type responseOutputText struct { + Type string `json:"type"` + Text string `json:"text"` + Annotations []any `json:"annotations"` +} + +type responsesUsage struct { + InputTokens int `json:"input_tokens"` + InputTokensDetails responseInputTokenDetails `json:"input_tokens_details"` + OutputTokens int `json:"output_tokens"` + OutputTokensDetail responseOutputTokenDetail `json:"output_tokens_details"` + TotalTokens int `json:"total_tokens"` +} + +type responseInputTokenDetails struct { + CachedTokens int `json:"cached_tokens"` +} + +type responseOutputTokenDetail struct { + ReasoningTokens int `json:"reasoning_tokens"` +} + +type completionUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +type chatCompletionResponse struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []chatChoice `json:"choices"` + Usage completionUsage `json:"usage"` +} + +type chatChoice struct { + Index int `json:"index"` + Message chatMessage `json:"message"` + FinishReason string `json:"finish_reason"` +} + +type chatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type chatCompletionChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []chatChunkChoice `json:"choices"` + Usage *completionUsage `json:"usage,omitempty"` +} + +type chatChunkChoice struct { + Index int `json:"index"` + Delta chatDelta `json:"delta"` + FinishReason *string `json:"finish_reason"` +} + +type chatDelta struct { + Role *string `json:"role,omitempty"` + Content *string `json:"content,omitempty"` +} + +type completionResponse struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []completionChoice `json:"choices"` + Usage completionUsage `json:"usage"` +} + +type completionChoice struct { + Text string `json:"text"` + Index int `json:"index"` + Logprobs any `json:"logprobs"` + FinishReason string `json:"finish_reason"` +} + +type completionChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []completionChunkChoice `json:"choices"` +} + +type completionChunkChoice struct { + Text string `json:"text"` + Index int `json:"index"` + Logprobs any `json:"logprobs"` + FinishReason *string `json:"finish_reason"` +} + +type embeddingsResponse struct { + Object string `json:"object"` + Data []embeddingResult `json:"data"` + Model string `json:"model"` + Usage embeddingUsage `json:"usage"` +} + +type embeddingResult struct { + Object string `json:"object"` + Embedding any `json:"embedding"` + Index int `json:"index"` +} + +type embeddingUsage struct { + PromptTokens int `json:"prompt_tokens"` + TotalTokens int `json:"total_tokens"` +} + +type modelsResponse struct { + Object string `json:"object"` + Data []modelInfo `json:"data"` +} + +type modelInfo struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + OwnedBy string `json:"owned_by"` +} + +type apiError struct { + Message string `json:"message"` + Type string `json:"type"` + Param any `json:"param"` + Code any `json:"code"` +} + +type apiErrorResponse struct { + Error apiError `json:"error"` +} + +type streamPacer struct { + timer *time.Timer +} + +type responsesDeltaWriter struct { + w http.ResponseWriter + itemID []byte + fixedChunk []byte + frame []byte + fixedChunkMode bool +} + +func main() { + config, err := loadServerConfig(os.Getenv) + if err != nil { + log.Fatal(err) + } + server := &http.Server{ + Addr: ":8000", + Handler: newRouterWithConfig(config), + ReadHeaderTimeout: 5 * time.Second, + ReadTimeout: 30 * time.Second, + IdleTimeout: 2 * time.Minute, + } + log.Printf("listening on %s with max output chunks %d", server.Addr, config.maxOutputChunks) + log.Fatal(server.ListenAndServe()) +} + +func newRouter() http.Handler { + return newRouterWithConfig(serverConfig{maxOutputChunks: defaultMaxOutputChunks}) +} + +func newRouterWithConfig(config serverConfig) http.Handler { + mux := http.NewServeMux() + mux.HandleFunc("/health", handleHealth) + mux.HandleFunc("/v1/models", handleModels) + mux.HandleFunc("/v1/models/", handleModel) + mux.HandleFunc("/v1/responses", func(w http.ResponseWriter, r *http.Request) { handleResponses(w, r, config) }) + mux.HandleFunc("/v1/chat/completions", func(w http.ResponseWriter, r *http.Request) { handleChatCompletions(w, r, config) }) + mux.HandleFunc("/v1/completions", func(w http.ResponseWriter, r *http.Request) { handleCompletions(w, r, config) }) + mux.HandleFunc("/v1/embeddings", func(w http.ResponseWriter, r *http.Request) { handleEmbeddings(w, r, config) }) + return mux +} + +func handleHealth(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + w.Header().Set("Allow", http.MethodGet) + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + w.WriteHeader(http.StatusOK) +} + +func handleModels(w http.ResponseWriter, r *http.Request) { + if !requireGet(w, r) { + return + } + writeJSON(w, http.StatusOK, modelsResponse{ + Object: "list", + Data: []modelInfo{newModelInfo(defaultModel)}, + }) +} + +func handleModel(w http.ResponseWriter, r *http.Request) { + if !requireGet(w, r) { + return + } + model := strings.TrimPrefix(r.URL.Path, "/v1/models/") + if model == "" { + writeAPIError(w, http.StatusNotFound, "model not found", "") + return + } + writeJSON(w, http.StatusOK, newModelInfo(model)) +} + +func handleResponses(w http.ResponseWriter, r *http.Request, config serverConfig) { + if !requirePost(w, r) { + return + } + var request responsesRequest + if err := decodeJSON(w, r, &request); err != nil { + writeAPIError(w, http.StatusBadRequest, "invalid JSON request body", "") + return + } + tuning, release, ok := startBenchmark(w, r, config, request.benchmarkBodyControls) + if !ok { + return + } + defer release() + if request.Model == "" { + writeAPIError(w, http.StatusBadRequest, "model is required", "model") + return + } + if request.Stream { + streamResponses(r.Context(), w, newResponsesResponse(request.Model, ""), tuning) + return + } + chunks := outputChunks(tuning) + response := newResponsesResponse(request.Model, strings.Join(chunks, "")) + if !waitFor(r.Context(), tuning.TTFT, tuning.TTFTJitter) { + return + } + setResponsesCompleted(&response) + writeJSON(w, http.StatusOK, response) +} + +func handleChatCompletions(w http.ResponseWriter, r *http.Request, config serverConfig) { + if !requirePost(w, r) { + return + } + var request chatCompletionsRequest + if err := decodeJSON(w, r, &request); err != nil { + writeAPIError(w, http.StatusBadRequest, "invalid JSON request body", "") + return + } + tuning, release, ok := startBenchmark(w, r, config, request.benchmarkBodyControls) + if !ok { + return + } + defer release() + if request.Model == "" { + writeAPIError(w, http.StatusBadRequest, "model is required", "model") + return + } + chunks := outputChunks(tuning) + + response := newChatCompletionResponse(request.Model, strings.Join(chunks, "")) + if request.Stream { + includeUsage := request.StreamOptions != nil && request.StreamOptions.IncludeUsage + streamChatCompletion(r.Context(), w, response, chunks, tuning, includeUsage) + return + } + if !waitFor(r.Context(), tuning.TTFT, tuning.TTFTJitter) { + return + } + writeJSON(w, http.StatusOK, response) +} + +func handleCompletions(w http.ResponseWriter, r *http.Request, config serverConfig) { + if !requirePost(w, r) { + return + } + var request completionsRequest + if err := decodeJSON(w, r, &request); err != nil { + writeAPIError(w, http.StatusBadRequest, "invalid JSON request body", "") + return + } + tuning, release, ok := startBenchmark(w, r, config, request.benchmarkBodyControls) + if !ok { + return + } + defer release() + if request.Model == "" { + writeAPIError(w, http.StatusBadRequest, "model is required", "model") + return + } + chunks := outputChunks(tuning) + + response := newCompletionResponse(request.Model, strings.Join(chunks, "")) + if request.Stream { + streamCompletion(r.Context(), w, response, chunks, tuning) + return + } + if !waitFor(r.Context(), tuning.TTFT, tuning.TTFTJitter) { + return + } + writeJSON(w, http.StatusOK, response) +} + +func handleEmbeddings(w http.ResponseWriter, r *http.Request, config serverConfig) { + if !requirePost(w, r) { + return + } + var request embeddingsRequest + if err := decodeJSON(w, r, &request); err != nil { + writeAPIError(w, http.StatusBadRequest, "invalid JSON request body", "") + return + } + tuning, release, ok := startBenchmark(w, r, config, request.benchmarkBodyControls) + if !ok { + return + } + defer release() + if request.Model == "" { + writeAPIError(w, http.StatusBadRequest, "model is required", "model") + return + } + inputs, err := embeddingInputs(request.Input) + if err != nil { + writeAPIError(w, http.StatusBadRequest, err.Error(), "input") + return + } + if request.EncodingFormat == "" { + request.EncodingFormat = "float" + } + if request.EncodingFormat != "float" && request.EncodingFormat != "base64" { + writeAPIError(w, http.StatusBadRequest, "encoding_format must be float or base64", "encoding_format") + return + } + if !waitFor(r.Context(), tuning.TTFT, tuning.TTFTJitter) { + return + } + + data := make([]embeddingResult, len(inputs)) + for index := range inputs { + vector := embeddingVector(index) + var embedding any = vector + if request.EncodingFormat == "base64" { + embedding = base64Embedding(vector) + } + data[index] = embeddingResult{ + Object: "embedding", + Embedding: embedding, + Index: index, + } + } + + tokens := countTokens(strings.Join(inputs, " ")) + writeJSON(w, http.StatusOK, embeddingsResponse{ + Object: "list", + Data: data, + Model: request.Model, + Usage: embeddingUsage{ + PromptTokens: tokens, + TotalTokens: tokens, + }, + }) +} + +func requirePost(w http.ResponseWriter, r *http.Request) bool { + if r.Method == http.MethodPost { + return true + } + w.Header().Set("Allow", http.MethodPost) + writeAPIError(w, http.StatusMethodNotAllowed, "method not allowed", "") + return false +} + +func requireGet(w http.ResponseWriter, r *http.Request) bool { + if r.Method == http.MethodGet { + return true + } + w.Header().Set("Allow", http.MethodGet) + writeAPIError(w, http.StatusMethodNotAllowed, "method not allowed", "") + return false +} + +func startBenchmark(w http.ResponseWriter, r *http.Request, config serverConfig, body benchmarkBodyControls) (benchmarkTuning, func(), bool) { + tuning, err := resolveBenchmarkTuning(r, body, config.maxOutputChunks) + if err != nil { + writeAPIError(w, http.StatusBadRequest, err.Error(), "") + return benchmarkTuning{}, nil, false + } + + active := activeRequests.Add(1) + release := func() { + activeRequests.Add(-1) + } + if tuning.MaxConcurrency > 0 && active > int64(tuning.MaxConcurrency) { + release() + writeAPIError(w, http.StatusTooManyRequests, "injected concurrency limit reached", "") + return benchmarkTuning{}, nil, false + } + if !waitFor(r.Context(), tuning.QueueDelay, 0) { + release() + return benchmarkTuning{}, nil, false + } + if tuning.StatusCode != 0 { + writeAPIErrorWithCode(w, tuning.StatusCode, "injected status response", "", "injected_status") + release() + return benchmarkTuning{}, nil, false + } + return tuning, release, true +} + +func resolveBenchmarkTuning(r *http.Request, body benchmarkBodyControls, maxOutputChunks int) (benchmarkTuning, error) { + if hasLoadTesterHeader(r) { + body = benchmarkBodyControls{} + } + tuning := benchmarkTuning{ + Chunk: defaultChunk, + OutputChunks: 1, + StreamErrorAfter: -1, + StreamTruncateAfter: -1, + } + var err error + for _, setting := range []struct { + header string + body string + raw json.RawMessage + value *time.Duration + }{ + {header: headerQueueDelay, body: bodyQueueDelay, raw: body.QueueDelay, value: &tuning.QueueDelay}, + {header: headerTTFT, body: bodyTTFT, raw: body.TTFT, value: &tuning.TTFT}, + {header: headerTTFTJitter, body: bodyTTFTJitter, raw: body.TTFTJitter, value: &tuning.TTFTJitter}, + {header: headerITL, body: bodyITL, raw: body.ITL, value: &tuning.ITL}, + {header: headerITLJitter, body: bodyITLJitter, raw: body.ITLJitter, value: &tuning.ITLJitter}, + } { + *setting.value, err = durationControl(r, setting.header, setting.body, setting.raw) + if err != nil { + return benchmarkTuning{}, err + } + } + + chunk, hasChunk, chunkName, err := textControl(r, headerChunk, bodyChunk, body.Chunk) + if err != nil { + return benchmarkTuning{}, err + } + if hasChunk { + if chunk == "" { + return benchmarkTuning{}, fmt.Errorf("%s must not be empty", chunkName) + } + tuning.Chunk = chunk + } + var chunkBytesName string + var hasChunkBytes bool + if tuning.ChunkBytes, hasChunkBytes, chunkBytesName, err = integerControl(r, headerChunkBytes, bodyChunkBytes, body.ChunkBytes, 0, 0, maxOutputBytes); err != nil { + return benchmarkTuning{}, err + } + if hasChunkBytes && tuning.ChunkBytes > 0 && hasChunk { + return benchmarkTuning{}, fmt.Errorf("%s and %s cannot be combined", chunkName, chunkBytesName) + } + if tuning.OutputChunks, _, _, err = integerControl(r, headerOutputChunks, bodyOutputChunks, body.OutputChunks, 1, 1, maxOutputChunks); err != nil { + return benchmarkTuning{}, err + } + var statusName string + if tuning.StatusCode, _, statusName, err = integerControl(r, headerStatusCode, bodyStatusCode, body.StatusCode, 0, 0, 599); err != nil { + return benchmarkTuning{}, err + } + if tuning.StatusCode != 0 && tuning.StatusCode < http.StatusBadRequest { + return benchmarkTuning{}, fmt.Errorf("%s must be an HTTP error status", statusName) + } + var streamErrorName string + if tuning.StreamErrorAfter, _, streamErrorName, err = integerControl(r, headerStreamErrorAfter, bodyStreamErrorAfter, body.StreamErrorAfter, -1, -1, tuning.OutputChunks); err != nil { + return benchmarkTuning{}, err + } + var streamTruncateName string + if tuning.StreamTruncateAfter, _, streamTruncateName, err = integerControl(r, headerStreamTruncateAfter, bodyStreamTruncateAfter, body.StreamTruncateAfter, -1, -1, tuning.OutputChunks); err != nil { + return benchmarkTuning{}, err + } + if tuning.StreamErrorAfter >= 0 && tuning.StreamTruncateAfter >= 0 { + return benchmarkTuning{}, fmt.Errorf("%s and %s cannot be combined", streamErrorName, streamTruncateName) + } + if tuning.MaxConcurrency, _, _, err = integerControl(r, headerMaxConcurrency, bodyMaxConcurrency, body.MaxConcurrency, 0, 0, maxConcurrencyLimit); err != nil { + return benchmarkTuning{}, err + } + + chunkBytes := len(tuning.Chunk) + if tuning.ChunkBytes > 0 { + chunkBytes = tuning.ChunkBytes + } + if int64(chunkBytes)*int64(tuning.OutputChunks) > maxOutputBytes { + return benchmarkTuning{}, fmt.Errorf("configured output exceeds %d bytes", maxOutputBytes) + } + return tuning, nil +} + +func loadServerConfig(getenv func(string) string) (serverConfig, error) { + config := serverConfig{maxOutputChunks: defaultMaxOutputChunks} + value := getenv(maxOutputChunksEnv) + if value == "" { + return config, nil + } + parsed, err := strconv.Atoi(value) + if err != nil || parsed < 1 || parsed > hardMaxOutputChunks { + return serverConfig{}, fmt.Errorf("%s must be an integer from 1 through %d", maxOutputChunksEnv, hardMaxOutputChunks) + } + config.maxOutputChunks = parsed + return config, nil +} + +func oneHeader(r *http.Request, name string) (string, bool, error) { + values := r.Header.Values(name) + if len(values) == 0 { + return "", false, nil + } + if len(values) != 1 { + return "", false, fmt.Errorf("%s must have one value", name) + } + return values[0], true, nil +} + +func hasLoadTesterHeader(r *http.Request) bool { + for name := range r.Header { + if strings.HasPrefix(strings.ToLower(name), loadTesterHeaderPrefix) { + return true + } + } + return false +} + +func controlValue(r *http.Request, header, body string, raw json.RawMessage) (string, json.RawMessage, bool, string, bool, error) { + value, present, err := oneHeader(r, header) + if err != nil { + return "", nil, false, header, true, err + } + if present { + return value, nil, true, header, true, nil + } + if len(raw) == 0 { + return "", nil, false, body, false, nil + } + return "", raw, true, body, false, nil +} + +func bodyInteger(raw json.RawMessage) (int64, error) { + decoder := json.NewDecoder(bytes.NewReader(raw)) + decoder.UseNumber() + var value any + if err := decoder.Decode(&value); err != nil { + return 0, err + } + var text string + switch typed := value.(type) { + case json.Number: + text = typed.String() + case string: + text = typed + default: + return 0, errors.New("must be an integer") + } + return strconv.ParseInt(text, 10, 64) +} + +func durationControl(r *http.Request, header, body string, raw json.RawMessage) (time.Duration, error) { + headerValue, bodyValue, present, source, fromHeader, err := controlValue(r, header, body, raw) + if err != nil { + return 0, err + } + if !present { + return 0, nil + } + var milliseconds int64 + if fromHeader { + milliseconds, err = strconv.ParseInt(headerValue, 10, 64) + } else { + milliseconds, err = bodyInteger(bodyValue) + } + if err != nil || milliseconds < 0 || milliseconds > maxControlMilliseconds { + return 0, fmt.Errorf("%s must be an integer from 0 to %d milliseconds", source, maxControlMilliseconds) + } + return time.Duration(milliseconds) * time.Millisecond, nil +} + +func integerControl(r *http.Request, header, body string, raw json.RawMessage, defaultValue, minimum, maximum int) (int, bool, string, error) { + headerValue, bodyValue, present, source, fromHeader, err := controlValue(r, header, body, raw) + if err != nil { + return 0, false, source, err + } + if !present { + return defaultValue, false, source, nil + } + var parsed64 int64 + if fromHeader { + parsed64, err = strconv.ParseInt(headerValue, 10, 64) + } else { + parsed64, err = bodyInteger(bodyValue) + } + if err != nil || parsed64 < int64(minimum) || parsed64 > int64(maximum) { + return 0, true, source, fmt.Errorf("%s must be an integer from %d to %d", source, minimum, maximum) + } + return int(parsed64), true, source, nil +} + +func textControl(r *http.Request, header, body string, raw json.RawMessage) (string, bool, string, error) { + headerValue, bodyValue, present, source, fromHeader, err := controlValue(r, header, body, raw) + if err != nil { + return "", false, source, err + } + if !present { + return "", false, source, nil + } + if fromHeader { + return headerValue, true, source, nil + } + var value string + if err := json.Unmarshal(bodyValue, &value); err != nil { + return "", true, source, fmt.Errorf("%s must be a string", source) + } + return value, true, source, nil +} + +func outputChunks(tuning benchmarkTuning) []string { + chunks := make([]string, tuning.OutputChunks) + for index := range chunks { + chunk := tuning.Chunk + if tuning.ChunkBytes > 0 { + chunk = randomText(tuning.ChunkBytes) + } + chunks[index] = chunk + } + return chunks +} + +func randomText(size int) string { + const alphabet = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + bytes := make([]byte, size) + randomSourceMu.Lock() + defer randomSourceMu.Unlock() + for index := range bytes { + bytes[index] = alphabet[randomSource.Intn(len(alphabet))] + } + return string(bytes) +} + +func waitFor(ctx context.Context, delay, jitter time.Duration) bool { + delay += randomDuration(jitter) + if delay <= 0 { + return true + } + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func (pacer *streamPacer) wait(ctx context.Context, delay, jitter time.Duration) bool { + if ctx.Err() != nil { + return false + } + delay += randomDuration(jitter) + if delay <= 0 { + return true + } + if pacer.timer == nil { + pacer.timer = time.NewTimer(delay) + } else { + if !pacer.timer.Stop() { + select { + case <-pacer.timer.C: + default: + } + } + pacer.timer.Reset(delay) + } + select { + case <-ctx.Done(): + return false + case <-pacer.timer.C: + return true + } +} + +func (pacer *streamPacer) stop() { + pacer.stopTimer() +} + +func (pacer *streamPacer) stopTimer() { + if pacer.timer == nil { + return + } + if !pacer.timer.Stop() { + select { + case <-pacer.timer.C: + default: + } + } + pacer.timer = nil +} + +func randomDuration(max time.Duration) time.Duration { + if max <= 0 { + return 0 + } + randomSourceMu.Lock() + defer randomSourceMu.Unlock() + return time.Duration(randomSource.Int63n(int64(max) + 1)) +} + +func decodeJSON(w http.ResponseWriter, r *http.Request, destination any) error { + r.Body = http.MaxBytesReader(w, r.Body, maxRequestBytes) + decoder := json.NewDecoder(r.Body) + if err := decoder.Decode(destination); err != nil { + return err + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return errors.New("request body must contain one JSON object") + } + return nil +} + +func writeJSON(w http.ResponseWriter, status int, payload any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + if err := json.NewEncoder(w).Encode(payload); err != nil { + log.Printf("response JSON write failed: %v", err) + } +} + +func writeAPIError(w http.ResponseWriter, status int, message, param string) { + writeAPIErrorWithCode(w, status, message, param, "") +} + +func writeAPIErrorWithCode(w http.ResponseWriter, status int, message, param, code string) { + var parameter any + if param != "" { + parameter = param + } + var errorCode any + if code != "" { + errorCode = code + } + writeJSON(w, status, apiErrorResponse{Error: apiError{ + Message: message, + Type: apiErrorType(status), + Param: parameter, + Code: errorCode, + }}) +} + +func apiErrorType(status int) string { + if status == http.StatusTooManyRequests { + return "rate_limit_error" + } + if status >= http.StatusInternalServerError { + return "server_error" + } + return "invalid_request_error" +} + +func embeddingInputs(input json.RawMessage) ([]string, error) { + trimmed := bytes.TrimSpace(input) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil, errors.New("input is required") + } + + var one string + if err := json.Unmarshal(trimmed, &one); err == nil { + return validateEmbeddingInputs([]string{one}) + } + + var many []string + if err := json.Unmarshal(trimmed, &many); err != nil { + return nil, errors.New("input must be a string or array of strings") + } + return validateEmbeddingInputs(many) +} + +func validateEmbeddingInputs(inputs []string) ([]string, error) { + if len(inputs) == 0 { + return nil, errors.New("input must include at least one string") + } + if len(inputs) > maxEmbeddingItems { + return nil, fmt.Errorf("input supports at most %d strings", maxEmbeddingItems) + } + for _, input := range inputs { + if input == "" { + return nil, errors.New("input strings must not be empty") + } + } + return inputs, nil +} + +func newResponsesResponse(model, text string) responsesResponse { + sequence := responseSequence.Add(1) + now := time.Now().Unix() + outputTokens := countTokens(text) + return responsesResponse{ + ID: fmt.Sprintf("resp_%d", sequence), + Object: "response", + CreatedAt: now, + Status: "completed", + Model: model, + ParallelToolCalls: true, + ToolChoice: "auto", + Tools: []any{}, + Output: []responseOutputItem{{ + ID: fmt.Sprintf("msg_%d", sequence), + Type: "message", + Status: "completed", + Role: "assistant", + Content: []responseOutputText{{ + Type: "output_text", + Text: text, + Annotations: []any{}, + }}, + }}, + Usage: newResponsesUsage(outputTokens), + } +} + +func newResponsesUsage(outputTokens int) *responsesUsage { + return &responsesUsage{ + InputTokens: 0, + InputTokensDetails: responseInputTokenDetails{CachedTokens: 0}, + OutputTokens: outputTokens, + OutputTokensDetail: responseOutputTokenDetail{ReasoningTokens: 0}, + TotalTokens: outputTokens, + } +} + +func setResponsesOutput(response *responsesResponse, text string) { + response.Output[0].Content[0].Text = text + response.Usage = newResponsesUsage(countTokens(text)) +} + +func setResponsesCompleted(response *responsesResponse) { + completedAt := time.Now().Unix() + response.CompletedAt = &completedAt +} + +func newChatCompletionResponse(model, text string) chatCompletionResponse { + sequence := responseSequence.Add(1) + outputTokens := countTokens(text) + return chatCompletionResponse{ + ID: fmt.Sprintf("chatcmpl_%d", sequence), + Object: "chat.completion", + Created: time.Now().Unix(), + Model: model, + Choices: []chatChoice{{ + Index: 0, + Message: chatMessage{ + Role: "assistant", + Content: text, + }, + FinishReason: "stop", + }}, + Usage: completionUsage{ + PromptTokens: 0, + CompletionTokens: outputTokens, + TotalTokens: outputTokens, + }, + } +} + +func newCompletionResponse(model, text string) completionResponse { + sequence := responseSequence.Add(1) + outputTokens := countTokens(text) + return completionResponse{ + ID: fmt.Sprintf("cmpl_%d", sequence), + Object: "text_completion", + Created: time.Now().Unix(), + Model: model, + Choices: []completionChoice{{ + Text: text, + Index: 0, + Logprobs: nil, + FinishReason: "stop", + }}, + Usage: completionUsage{ + PromptTokens: 0, + CompletionTokens: outputTokens, + TotalTokens: outputTokens, + }, + } +} + +func newModelInfo(model string) modelInfo { + return modelInfo{ + ID: model, + Object: "model", + Created: time.Now().Unix(), + OwnedBy: "nvidia", + } +} + +func streamResponses(ctx context.Context, w http.ResponseWriter, response responsesResponse, tuning benchmarkTuning) { + setSSEHeaders(w) + if !waitFor(ctx, tuning.TTFT, tuning.TTFTJitter) { + return + } + + inProgress := response + inProgress.Status = "in_progress" + inProgress.CompletedAt = nil + inProgress.Output = []responseOutputItem{} + inProgress.Usage = nil + + item := response.Output[0] + itemInProgress := item + itemInProgress.Status = "in_progress" + itemInProgress.Content = []responseOutputText{} + part := item.Content[0] + emptyPart := part + emptyPart.Text = "" + + events := []struct { + name string + data any + }{ + {"response.created", map[string]any{"type": "response.created", "sequence_number": 0, "response": inProgress}}, + {"response.in_progress", map[string]any{"type": "response.in_progress", "sequence_number": 1, "response": inProgress}}, + {"response.output_item.added", map[string]any{"type": "response.output_item.added", "sequence_number": 2, "output_index": 0, "item": itemInProgress}}, + {"response.content_part.added", map[string]any{"type": "response.content_part.added", "sequence_number": 3, "item_id": item.ID, "output_index": 0, "content_index": 0, "part": emptyPart}}, + } + for _, event := range events { + if err := writeSSEJSON(w, event.name, event.data); err != nil { + log.Printf("Responses event write failed: %v", err) + return + } + } + + if terminate, truncated := streamTermination(tuning, 0); terminate { + if !truncated { + writeResponsesStreamError(w, 4) + } + return + } + deltaWriter := newResponsesDeltaWriter(w, item.ID, tuning.Chunk, tuning.ChunkBytes == 0) + var pacer streamPacer + defer pacer.stop() + var output strings.Builder + for index := 0; index < tuning.OutputChunks; index++ { + if index > 0 && !pacer.wait(ctx, tuning.ITL, tuning.ITLJitter) { + return + } + chunk := tuning.Chunk + if tuning.ChunkBytes > 0 { + chunk = randomText(tuning.ChunkBytes) + } + if err := deltaWriter.write(4+index, chunk); err != nil { + log.Printf("Responses event write failed: %v", err) + return + } + if tuning.ChunkBytes > 0 { + output.WriteString(chunk) + } + if terminate, truncated := streamTermination(tuning, index+1); terminate { + if !truncated { + writeResponsesStreamError(w, 5+index) + } + return + } + } + if ctx.Err() != nil { + return + } + + text := output.String() + if tuning.ChunkBytes == 0 { + text = strings.Repeat(tuning.Chunk, tuning.OutputChunks) + } + setResponsesOutput(&response, text) + item = response.Output[0] + part = item.Content[0] + setResponsesCompleted(&response) + sequence := 4 + tuning.OutputChunks + events = []struct { + name string + data any + }{ + {"response.output_text.done", map[string]any{"type": "response.output_text.done", "sequence_number": sequence, "item_id": item.ID, "output_index": 0, "content_index": 0, "text": part.Text, "logprobs": []any{}}}, + {"response.content_part.done", map[string]any{"type": "response.content_part.done", "sequence_number": sequence + 1, "item_id": item.ID, "output_index": 0, "content_index": 0, "part": part}}, + {"response.output_item.done", map[string]any{"type": "response.output_item.done", "sequence_number": sequence + 2, "output_index": 0, "item": item}}, + {"response.completed", map[string]any{"type": "response.completed", "sequence_number": sequence + 3, "response": response}}, + } + for _, event := range events { + if err := writeSSEJSON(w, event.name, event.data); err != nil { + log.Printf("Responses event write failed: %v", err) + return + } + } +} + +func streamChatCompletion(ctx context.Context, w http.ResponseWriter, response chatCompletionResponse, chunks []string, tuning benchmarkTuning, includeUsage bool) { + setSSEHeaders(w) + if !waitFor(ctx, tuning.TTFT, tuning.TTFTJitter) { + return + } + + role := "assistant" + if err := writeSSEJSON(w, "", chatCompletionChunk{ + ID: response.ID, + Object: "chat.completion.chunk", + Created: response.Created, + Model: response.Model, + Choices: []chatChunkChoice{{Index: 0, Delta: chatDelta{Role: &role}}}, + }); err != nil { + log.Printf("Chat role event write failed: %v", err) + return + } + + if terminate, truncated := streamTermination(tuning, 0); terminate { + if !truncated { + writeLegacyStreamError(w) + } + return + } + var pacer streamPacer + defer pacer.stop() + for index, text := range chunks { + if index > 0 && !pacer.wait(ctx, tuning.ITL, tuning.ITLJitter) { + return + } + if err := writeSSEJSON(w, "", chatCompletionChunk{ + ID: response.ID, + Object: "chat.completion.chunk", + Created: response.Created, + Model: response.Model, + Choices: []chatChunkChoice{{Index: 0, Delta: chatDelta{Content: &text}}}, + }); err != nil { + log.Printf("Chat content event write failed: %v", err) + return + } + if terminate, truncated := streamTermination(tuning, index+1); terminate { + if !truncated { + writeLegacyStreamError(w) + } + return + } + } + if ctx.Err() != nil { + return + } + + stop := "stop" + if err := writeSSEJSON(w, "", chatCompletionChunk{ + ID: response.ID, + Object: "chat.completion.chunk", + Created: response.Created, + Model: response.Model, + Choices: []chatChunkChoice{{Index: 0, Delta: chatDelta{}, FinishReason: &stop}}, + }); err != nil { + log.Printf("Chat stop event write failed: %v", err) + return + } + + if includeUsage { + usage := response.Usage + if err := writeSSEJSON(w, "", chatCompletionChunk{ + ID: response.ID, + Object: "chat.completion.chunk", + Created: response.Created, + Model: response.Model, + Choices: []chatChunkChoice{}, + Usage: &usage, + }); err != nil { + log.Printf("Chat usage event write failed: %v", err) + return + } + } + if err := writeSSEFrame(w, []byte("data: [DONE]\n\n")); err != nil { + log.Printf("Chat completion event write failed: %v", err) + return + } +} + +func streamCompletion(ctx context.Context, w http.ResponseWriter, response completionResponse, chunks []string, tuning benchmarkTuning) { + setSSEHeaders(w) + if !waitFor(ctx, tuning.TTFT, tuning.TTFTJitter) { + return + } + + if terminate, truncated := streamTermination(tuning, 0); terminate { + if !truncated { + writeLegacyStreamError(w) + } + return + } + var pacer streamPacer + defer pacer.stop() + for index, text := range chunks { + if index > 0 && !pacer.wait(ctx, tuning.ITL, tuning.ITLJitter) { + return + } + if err := writeSSEJSON(w, "", completionChunk{ + ID: response.ID, + Object: "text_completion", + Created: response.Created, + Model: response.Model, + Choices: []completionChunkChoice{{Text: text, Index: 0, Logprobs: nil}}, + }); err != nil { + log.Printf("Completion event write failed: %v", err) + return + } + if terminate, truncated := streamTermination(tuning, index+1); terminate { + if !truncated { + writeLegacyStreamError(w) + } + return + } + } + if ctx.Err() != nil { + return + } + + stop := "stop" + if err := writeSSEJSON(w, "", completionChunk{ + ID: response.ID, + Object: "text_completion", + Created: response.Created, + Model: response.Model, + Choices: []completionChunkChoice{{Text: "", Index: 0, Logprobs: nil, FinishReason: &stop}}, + }); err != nil { + log.Printf("Completion stop event write failed: %v", err) + return + } + if err := writeSSEFrame(w, []byte("data: [DONE]\n\n")); err != nil { + log.Printf("Completion event write failed: %v", err) + return + } +} + +func streamTermination(tuning benchmarkTuning, emitted int) (bool, bool) { + if tuning.StreamErrorAfter >= 0 && emitted >= tuning.StreamErrorAfter { + return true, false + } + if tuning.StreamTruncateAfter >= 0 && emitted >= tuning.StreamTruncateAfter { + return true, true + } + return false, false +} + +func writeResponsesStreamError(w http.ResponseWriter, sequence int) { + if err := writeSSEJSON(w, "error", map[string]any{ + "type": "error", + "sequence_number": sequence, + "error": apiError{ + Message: "injected streaming error", + Type: "server_error", + Param: nil, + Code: "injected_stream_error", + }, + }); err != nil { + log.Printf("Responses error event write failed: %v", err) + } +} + +func writeLegacyStreamError(w http.ResponseWriter) { + if err := writeSSEJSON(w, "", apiErrorResponse{Error: apiError{ + Message: "injected streaming error", + Type: "server_error", + Param: nil, + Code: "injected_stream_error", + }}); err != nil { + log.Printf("Streaming error event write failed: %v", err) + } +} + +func setSSEHeaders(w http.ResponseWriter) { + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Content-Type", "text/event-stream") +} + +func writeSSEJSON(w http.ResponseWriter, name string, data any) error { + payload, err := json.Marshal(data) + if err != nil { + return err + } + frame := make([]byte, 0, len(name)+len(payload)+16) + if name != "" { + frame = append(frame, "event: "...) + frame = append(frame, name...) + frame = append(frame, '\n') + } + frame = append(frame, "data: "...) + frame = append(frame, payload...) + frame = append(frame, '\n', '\n') + return writeSSEFrame(w, frame) +} + +func newResponsesDeltaWriter(w http.ResponseWriter, itemID, chunk string, fixedChunkMode bool) responsesDeltaWriter { + itemIDJSON, _ := json.Marshal(itemID) + fixedChunkJSON, _ := json.Marshal(chunk) + return responsesDeltaWriter{ + w: w, + itemID: itemIDJSON, + fixedChunk: fixedChunkJSON, + frame: make([]byte, 0, len(itemIDJSON)+len(fixedChunkJSON)+128), + fixedChunkMode: fixedChunkMode, + } +} + +func (writer *responsesDeltaWriter) write(sequence int, chunk string) error { + frame := writer.frame[:0] + frame = append(frame, "event: response.output_text.delta\ndata: {\"content_index\":0,\"delta\":"...) + if writer.fixedChunkMode { + frame = append(frame, writer.fixedChunk...) + } else { + chunkJSON, err := json.Marshal(chunk) + if err != nil { + return err + } + frame = append(frame, chunkJSON...) + } + frame = append(frame, ",\"item_id\":"...) + frame = append(frame, writer.itemID...) + frame = append(frame, ",\"logprobs\":[],\"output_index\":0,\"sequence_number\":"...) + frame = strconv.AppendInt(frame, int64(sequence), 10) + frame = append(frame, ",\"type\":\"response.output_text.delta\"}\n\n"...) + writer.frame = frame + return writeSSEFrame(writer.w, frame) +} + +func writeSSEFrame(w http.ResponseWriter, frame []byte) error { + for len(frame) > 0 { + written, err := w.Write(frame) + if err != nil { + return err + } + if written == 0 { + return io.ErrShortWrite + } + frame = frame[written:] + } + flush(w) + return nil +} + +func flush(w http.ResponseWriter) { + if flusher, ok := w.(http.Flusher); ok { + flusher.Flush() + } +} + +func embeddingVector(index int) []float64 { + value := float64(index) + return []float64{value, value + 0.125, -value - 0.25} +} + +func base64Embedding(vector []float64) string { + bytes := make([]byte, 4*len(vector)) + for index, value := range vector { + binary.LittleEndian.PutUint32(bytes[index*4:], math.Float32bits(float32(value))) + } + return base64.StdEncoding.EncodeToString(bytes) +} + +func countTokens(text string) int { + return len(strings.Fields(text)) +} diff --git a/examples/function-samples/openai-compatible-sample/http-server/main_test.go b/examples/function-samples/openai-compatible-sample/http-server/main_test.go new file mode 100644 index 0000000000..072b7f4cea --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/http-server/main_test.go @@ -0,0 +1,1159 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package main + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" +) + +func TestResponsesReturnsStrictCompatibleResponse(t *testing.T) { + recorder := postJSON(t, "/v1/responses", `{"model":"test-model","input":"hello world"}`) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + + body := responseMap(t, recorder) + for _, field := range []string{ + "id", + "object", + "created_at", + "completed_at", + "parallel_tool_calls", + "tool_choice", + "tools", + "usage", + } { + if _, ok := body[field]; !ok { + t.Fatalf("response missing %q: %s", field, recorder.Body.String()) + } + } + if got := body["object"]; got != "response" { + t.Fatalf("object = %#v, want response", got) + } + if body["parallel_tool_calls"] != true { + t.Fatalf("parallel_tool_calls = %#v, want true", body["parallel_tool_calls"]) + } + + usage := body["usage"].(map[string]any) + for _, field := range []string{"input_tokens_details", "output_tokens_details"} { + if _, ok := usage[field]; !ok { + t.Fatalf("usage missing %q: %s", field, recorder.Body.String()) + } + } + + output := body["output"].([]any) + item := output[0].(map[string]any) + content := item["content"].([]any) + if got := content[0].(map[string]any)["text"]; got != defaultChunk { + t.Fatalf("output text = %#v, want %s", got, defaultChunk) + } +} + +func TestResponsesStreamUsesHeaderControls(t *testing.T) { + started := time.Now() + recorder := postJSONWithHeaders(t, "/v1/responses", `{ + "model":"test-model", + "input":"hello", + "stream":true, + "repeats":99, + "delay":99, + "size":999 + }`, map[string]string{ + headerTTFT: "20", + headerITL: "10", + headerChunk: "header", + headerOutputChunks: "3", + }) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if elapsed := time.Since(started); elapsed < 35*time.Millisecond { + t.Fatalf("stream elapsed = %s, want at least 35ms", elapsed) + } + if contentType := recorder.Header().Get("Content-Type"); !strings.HasPrefix(contentType, "text/event-stream") { + t.Fatalf("Content-Type = %q, want text/event-stream", contentType) + } + body := recorder.Body.String() + if got := strings.Count(body, `"delta":"header"`); got != 3 { + t.Fatalf("body chunk count = %d, want 3: %s", got, body) + } + if !strings.Contains(body, `"text":"headerheaderheader"`) { + t.Fatalf("stream did not complete with repeated header chunks: %s", body) + } + + lastIndex := -1 + for _, event := range []string{ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + } { + index := strings.Index(body, "event: "+event) + if index == -1 { + t.Fatalf("stream missing %q: %s", event, body) + } + if index <= lastIndex { + t.Fatalf("event %q arrived out of order: %s", event, body) + } + lastIndex = index + } +} + +func TestBodyControlsWorkBeforeAndAfterNormalFields(t *testing.T) { + for _, test := range []struct { + name string + path string + body string + want string + }{ + { + name: "responses before input", + path: "/v1/responses", + body: `{"x_load_tester_ttft_ms":0,"x_load_tester_chunk":"before-response","x_load_tester_output_chunks":2,"model":"test-model","input":"hello"}`, + want: `"text":"before-responsebefore-response"`, + }, + { + name: "responses after input", + path: "/v1/responses", + body: `{"model":"test-model","input":"hello","x_load_tester_ttft_ms":0,"x_load_tester_chunk":"after-response","x_load_tester_output_chunks":2}`, + want: `"text":"after-responseafter-response"`, + }, + { + name: "chat before messages", + path: "/v1/chat/completions", + body: `{"x_load_tester_ttft_ms":0,"x_load_tester_chunk":"before-chat","x_load_tester_output_chunks":2,"model":"test-model","messages":[{"role":"user","content":"hello"}]}`, + want: `"content":"before-chatbefore-chat"`, + }, + { + name: "chat after messages", + path: "/v1/chat/completions", + body: `{"model":"test-model","messages":[{"role":"user","content":"hello"}],"x_load_tester_ttft_ms":0,"x_load_tester_chunk":"after-chat","x_load_tester_output_chunks":2}`, + want: `"content":"after-chatafter-chat"`, + }, + { + name: "completion before prompt", + path: "/v1/completions", + body: `{"x_load_tester_ttft_ms":0,"x_load_tester_chunk":"before-completion","x_load_tester_output_chunks":2,"model":"test-model","prompt":"hello"}`, + want: `"text":"before-completionbefore-completion"`, + }, + { + name: "completion after prompt", + path: "/v1/completions", + body: `{"model":"test-model","prompt":"hello","x_load_tester_ttft_ms":0,"x_load_tester_chunk":"after-completion","x_load_tester_output_chunks":2}`, + want: `"text":"after-completionafter-completion"`, + }, + } { + t.Run(test.name, func(t *testing.T) { + recorder := postJSON(t, test.path, test.body) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if !strings.Contains(recorder.Body.String(), test.want) { + t.Fatalf("response does not contain %s: %s", test.want, recorder.Body.String()) + } + }) + } +} + +func TestAnyLoadTesterHeaderDisablesBodyControls(t *testing.T) { + recorder := postJSONWithHeaders(t, "/v1/responses", `{ + "model":"test-model", + "input":"hello", + "stream":true, + "x_load_tester_ttft_ms":"not-an-integer", + "x_load_tester_chunk":"body", + "x_load_tester_output_chunks":2, + "x_load_tester_status_code":503 + }`, map[string]string{"X-Load-Tester-Ignored": "1"}) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if got := strings.Count(recorder.Body.String(), `"delta":"`+defaultChunk+`"`); got != 1 { + t.Fatalf("delta count = %d, want 1: %s", got, recorder.Body.String()) + } + if !strings.Contains(recorder.Body.String(), "event: response.completed") { + t.Fatalf("stream did not complete: %s", recorder.Body.String()) + } +} + +func TestResponsesStreamPreservesSSESemanticsAndFlushes(t *testing.T) { + chunk := "quote\" slash\\ newline\n" + recorder := newFlushCountingRecorder() + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/responses", strings.NewReader(`{"model":"test-model","input":"hello","stream":true}`)) + request.Header.Set("Content-Type", "application/json") + request.Header.Set(headerChunk, chunk) + request.Header.Set(headerOutputChunks, "3") + newRouter().ServeHTTP(recorder, request) + + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + events := parseResponsesSSEEvents(t, recorder.Body.String()) + wantNames := []string{ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.delta", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + } + if len(events) != len(wantNames) { + t.Fatalf("event count = %d, want %d: %s", len(events), len(wantNames), recorder.Body.String()) + } + if recorder.flushes != len(events) { + t.Fatalf("flushes = %d, want %d", recorder.flushes, len(events)) + } + + var deltaText strings.Builder + for index, event := range events { + if event.name != wantNames[index] { + t.Fatalf("event %d = %q, want %q", index, event.name, wantNames[index]) + } + if got := sseSequence(t, event); got != index { + t.Fatalf("event %d sequence = %d, want %d", index, got, index) + } + if event.name == "response.output_text.delta" { + delta, ok := event.data["delta"].(string) + if !ok { + t.Fatalf("delta = %#v, want string", event.data["delta"]) + } + if delta != chunk { + t.Fatalf("delta = %q, want %q", delta, chunk) + } + deltaText.WriteString(delta) + } + } + + wantText := strings.Repeat(chunk, 3) + if got := eventText(t, events[7], "text"); got != wantText { + t.Fatalf("output_text.done text = %q, want %q", got, wantText) + } + if got := nestedEventText(t, events[8], "part"); got != wantText { + t.Fatalf("content_part.done text = %q, want %q", got, wantText) + } + completed := events[10].data["response"].(map[string]any) + if got := nestedResponseText(t, completed); got != wantText { + t.Fatalf("completed response text = %q, want %q", got, wantText) + } + if got := completed["status"]; got != "completed" { + t.Fatalf("completed status = %#v, want completed", got) + } + usage := completed["usage"].(map[string]any) + if got := int(usage["output_tokens"].(float64)); got != countTokens(deltaText.String()) { + t.Fatalf("output tokens = %d, want %d", got, countTokens(deltaText.String())) + } +} + +func TestResponsesStreamRandomChunksMatchCompletedText(t *testing.T) { + recorder := postJSONWithHeaders(t, "/v1/responses", `{"model":"test-model","input":"hello","stream":true}`, map[string]string{ + headerChunkBytes: "7", + headerOutputChunks: "3", + }) + events := parseResponsesSSEEvents(t, recorder.Body.String()) + var deltaText strings.Builder + for _, event := range events { + if event.name != "response.output_text.delta" { + continue + } + delta := event.data["delta"].(string) + if len(delta) != 7 { + t.Fatalf("random chunk length = %d, want 7", len(delta)) + } + deltaText.WriteString(delta) + } + if got := eventText(t, events[len(events)-4], "text"); got != deltaText.String() { + t.Fatalf("output_text.done text = %q, want %q", got, deltaText.String()) + } + completed := events[len(events)-1].data["response"].(map[string]any) + if got := nestedResponseText(t, completed); got != deltaText.String() { + t.Fatalf("completed response text = %q, want %q", got, deltaText.String()) + } +} + +func TestResponsesStreamCancellationDoesNotComplete(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + recorder := newFlushCountingRecorder() + recorder.onWrite = func(frame []byte) { + if bytes.Contains(frame, []byte("event: response.output_text.delta\n")) { + cancel() + } + } + done := make(chan struct{}) + go func() { + defer close(done) + streamResponses(ctx, recorder, newResponsesResponse("test-model", ""), benchmarkTuning{ + ITL: time.Hour, + Chunk: defaultChunk, + OutputChunks: 2, + StreamErrorAfter: -1, + StreamTruncateAfter: -1, + }) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("stream did not stop after context cancellation") + } + if strings.Contains(recorder.Body.String(), "event: response.completed") { + t.Fatalf("cancelled stream completed: %s", recorder.Body.String()) + } +} + +func TestResponsesStreamCancellationDuringFinalDeltaDoesNotComplete(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + recorder := newFlushCountingRecorder() + recorder.onWrite = func(frame []byte) { + if bytes.Contains(frame, []byte("event: response.output_text.delta\n")) { + cancel() + } + } + + streamResponses(ctx, recorder, newResponsesResponse("test-model", ""), benchmarkTuning{ + Chunk: defaultChunk, + OutputChunks: 1, + StreamErrorAfter: -1, + StreamTruncateAfter: -1, + }) + + if strings.Contains(recorder.Body.String(), "event: response.completed") { + t.Fatalf("cancelled final delta completed: %s", recorder.Body.String()) + } +} + +func TestResponsesStreamWaitsAfterSlowDeltaWrite(t *testing.T) { + const itl = 25 * time.Millisecond + recorder := &delayedDeltaRecorder{ + flushCountingRecorder: newFlushCountingRecorder(), + delay: 50 * time.Millisecond, + } + + streamResponses(context.Background(), recorder, newResponsesResponse("test-model", ""), benchmarkTuning{ + ITL: itl, + Chunk: defaultChunk, + OutputChunks: 3, + StreamErrorAfter: -1, + StreamTruncateAfter: -1, + }) + + if len(recorder.deltaWrites) != 3 { + t.Fatalf("delta writes = %d, want 3", len(recorder.deltaWrites)) + } + if elapsed := recorder.deltaWrites[2].Sub(recorder.deltaWrites[1]); elapsed < itl-5*time.Millisecond { + t.Fatalf("post-write ITL = %s, want at least %s", elapsed, itl-5*time.Millisecond) + } +} + +func TestLoadServerConfig(t *testing.T) { + for _, test := range []struct { + name string + value string + want int + wantErr bool + }{ + {name: "unset defaults", want: defaultMaxOutputChunks}, + {name: "valid override", value: "12000", want: 12000}, + {name: "non numeric rejected", value: "many", wantErr: true}, + {name: "zero rejected", value: "0", wantErr: true}, + {name: "above hard maximum rejected", value: "60001", wantErr: true}, + } { + t.Run(test.name, func(t *testing.T) { + config, err := loadServerConfig(func(name string) string { + if name == maxOutputChunksEnv { + return test.value + } + return "" + }) + if test.wantErr { + if err == nil { + t.Fatal("loadServerConfig() error = nil, want error") + } + return + } + if err != nil { + t.Fatalf("loadServerConfig() error = %v", err) + } + if config.maxOutputChunks != test.want { + t.Fatalf("max output chunks = %d, want %d", config.maxOutputChunks, test.want) + } + }) + } +} + +func TestOutputChunksLimit(t *testing.T) { + for _, test := range []struct { + name string + value string + want bool + }{ + {name: "default maximum accepted", value: "6000", want: true}, + {name: "above default maximum rejected", value: "6001", want: false}, + } { + t.Run(test.name, func(t *testing.T) { + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/responses", nil) + request.Header.Set(headerOutputChunks, test.value) + tuning, err := resolveBenchmarkTuning(request, benchmarkBodyControls{}, defaultMaxOutputChunks) + if test.want { + if err != nil { + t.Fatalf("resolveBenchmarkTuning() error = %v", err) + } + if tuning.OutputChunks != defaultMaxOutputChunks { + t.Fatalf("output chunks = %d, want %d", tuning.OutputChunks, defaultMaxOutputChunks) + } + return + } + if err == nil { + t.Fatal("resolveBenchmarkTuning() error = nil, want error") + } + }) + } +} + +func TestConfiguredOutputChunksLimit(t *testing.T) { + config, err := loadServerConfig(func(name string) string { + if name == maxOutputChunksEnv { + return "12000" + } + return "" + }) + if err != nil { + t.Fatalf("loadServerConfig() error = %v", err) + } + + for _, test := range []struct { + name string + output string + wantStatus int + }{ + {name: "configured maximum accepted", output: "12000", wantStatus: http.StatusOK}, + {name: "above configured maximum rejected", output: "12001", wantStatus: http.StatusBadRequest}, + } { + t.Run(test.name, func(t *testing.T) { + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/responses", strings.NewReader(`{"model":"test-model"}`)) + request.Header.Set("Content-Type", "application/json") + request.Header.Set(headerOutputChunks, test.output) + recorder := httptest.NewRecorder() + newRouterWithConfig(config).ServeHTTP(recorder, request) + if recorder.Code != test.wantStatus { + t.Fatalf("status = %d, want %d: %s", recorder.Code, test.wantStatus, recorder.Body.String()) + } + }) + } +} + +func TestWriteSSEFrameHandlesPartialWrites(t *testing.T) { + writer := &partialResponseWriter{header: make(http.Header), maxWrite: 3} + if err := writeSSEFrame(writer, []byte("event: test\ndata: {}\n\n")); err != nil { + t.Fatalf("writeSSEFrame() error = %v", err) + } + if got, want := writer.body.String(), "event: test\ndata: {}\n\n"; got != want { + t.Fatalf("body = %q, want %q", got, want) + } + if writer.flushes != 1 { + t.Fatalf("flushes = %d, want 1", writer.flushes) + } +} + +func BenchmarkResponsesFixedChunkStream(b *testing.B) { + tuning := benchmarkTuning{ + Chunk: defaultChunk, + OutputChunks: 1000, + StreamErrorAfter: -1, + StreamTruncateAfter: -1, + } + writer := &discardResponseWriter{header: make(http.Header)} + b.SetBytes(int64(tuning.OutputChunks * len(tuning.Chunk))) + b.ReportAllocs() + b.ResetTimer() + for index := 0; index < b.N; index++ { + streamResponses(context.Background(), writer, newResponsesResponse("test-model", ""), tuning) + } +} + +func TestOpenAITextEndpointsIgnoreBodyTuning(t *testing.T) { + for _, test := range []struct { + path string + body string + }{ + { + path: "/v1/responses", + body: `{"model":"test-model","input":null,"repeats":99,"delay":99,"size":999}`, + }, + { + path: "/v1/chat/completions", + body: `{"model":"test-model","messages":[],"repeats":99,"delay":99,"size":999}`, + }, + { + path: "/v1/completions", + body: `{"model":"test-model","prompt":"ignored","repeats":99,"delay":99,"size":999}`, + }, + } { + t.Run(test.path, func(t *testing.T) { + recorder := postJSON(t, test.path, test.body) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if !strings.Contains(recorder.Body.String(), defaultChunk) { + t.Fatalf("response does not contain default chunk: %s", recorder.Body.String()) + } + }) + } +} + +func TestChunkBytesControlsOutput(t *testing.T) { + recorder := postJSONWithHeaders(t, "/v1/chat/completions", `{ + "model":"test-model", + "messages":[{"role":"user","content":"hello"}] + }`, map[string]string{ + headerChunkBytes: "7", + headerOutputChunks: "2", + }) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body := responseMap(t, recorder) + choices := body["choices"].([]any) + message := choices[0].(map[string]any)["message"].(map[string]any) + if got := len(message["content"].(string)); got != 14 { + t.Fatalf("content length = %d, want 14", got) + } +} + +func TestBodyControlMapping(t *testing.T) { + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/responses", nil) + tuning, err := resolveBenchmarkTuning(request, benchmarkBodyControls{ + QueueDelay: json.RawMessage(`1`), + TTFT: json.RawMessage(`2`), + TTFTJitter: json.RawMessage(`3`), + ITL: json.RawMessage(`4`), + ITLJitter: json.RawMessage(`5`), + Chunk: json.RawMessage(`"body"`), + OutputChunks: json.RawMessage(`6`), + StatusCode: json.RawMessage(`503`), + StreamErrorAfter: json.RawMessage(`1`), + MaxConcurrency: json.RawMessage(`2`), + }, defaultMaxOutputChunks) + if err != nil { + t.Fatalf("resolveBenchmarkTuning() error = %v", err) + } + if tuning.QueueDelay != time.Millisecond || tuning.TTFT != 2*time.Millisecond || tuning.TTFTJitter != 3*time.Millisecond || tuning.ITL != 4*time.Millisecond || tuning.ITLJitter != 5*time.Millisecond { + t.Fatalf("timing controls = %#v, want body values", tuning) + } + if tuning.Chunk != "body" || tuning.OutputChunks != 6 || tuning.StatusCode != http.StatusServiceUnavailable || tuning.StreamErrorAfter != 1 || tuning.MaxConcurrency != 2 { + t.Fatalf("body controls = %#v, want configured values", tuning) + } +} + +func TestBodyChunkBytesControlsOutput(t *testing.T) { + recorder := postJSON(t, "/v1/chat/completions", `{ + "model":"test-model", + "messages":[{"role":"user","content":"hello"}], + "x_load_tester_chunk_bytes":7, + "x_load_tester_output_chunks":2 + }`) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body := responseMap(t, recorder) + choices := body["choices"].([]any) + message := choices[0].(map[string]any)["message"].(map[string]any) + if got := len(message["content"].(string)); got != 14 { + t.Fatalf("content length = %d, want 14", got) + } +} + +func TestRejectsInvalidLoadTesterBodyControls(t *testing.T) { + for name, body := range map[string]string{ + "negative itl": `{"model":"test-model","input":"hello","x_load_tester_itl_ms":-1}`, + "empty chunk": `{"model":"test-model","input":"hello","x_load_tester_chunk":""}`, + "zero chunks": `{"model":"test-model","input":"hello","x_load_tester_output_chunks":0}`, + "success status": `{"model":"test-model","input":"hello","x_load_tester_status_code":200}`, + "combined chunk values": `{"model":"test-model","input":"hello","x_load_tester_chunk":"text","x_load_tester_chunk_bytes":4}`, + "output over cap": `{"model":"test-model","input":"hello","x_load_tester_chunk_bytes":1048576,"x_load_tester_output_chunks":2}`, + "competing stream stops": `{"model":"test-model","input":"hello","x_load_tester_stream_error_after_chunks":0,"x_load_tester_stream_truncate_after_chunks":0}`, + } { + t.Run(name, func(t *testing.T) { + recorder := postJSON(t, "/v1/responses", body) + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusBadRequest, recorder.Body.String()) + } + }) + } +} + +func TestRejectsInvalidLoadTesterHeaders(t *testing.T) { + for name, headers := range map[string]map[string]string{ + "negative itl": { + headerITL: "-1", + }, + "empty chunk": { + headerChunk: "", + }, + "zero chunks": { + headerOutputChunks: "0", + }, + "success status": { + headerStatusCode: "200", + }, + "combined chunk values": { + headerChunk: "text", + headerChunkBytes: "4", + }, + "output over cap": { + headerChunkBytes: strconv.Itoa(maxOutputBytes), + headerOutputChunks: "2", + }, + "competing stream stops": { + headerStreamErrorAfter: "0", + headerStreamTruncateAfter: "0", + }, + } { + t.Run(name, func(t *testing.T) { + recorder := postJSONWithHeaders(t, "/v1/responses", `{"model":"test-model","input":"hello"}`, headers) + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusBadRequest, recorder.Body.String()) + } + }) + } +} + +func TestInjectedStatusAndStreamFailures(t *testing.T) { + recorder := postJSONWithHeaders(t, "/v1/responses", `{}`, map[string]string{ + headerStatusCode: "503", + }) + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusServiceUnavailable, recorder.Body.String()) + } + if got := responseMap(t, recorder)["error"].(map[string]any)["type"]; got != "server_error" { + t.Fatalf("error type = %#v, want server_error", got) + } + + recorder = postJSONWithHeaders(t, "/v1/responses", `{"model":"test-model","input":"hello","stream":true}`, map[string]string{ + headerOutputChunks: "2", + headerStreamErrorAfter: "1", + }) + body := recorder.Body.String() + if !strings.Contains(body, "event: error") { + t.Fatalf("stream does not contain error event: %s", body) + } + if strings.Contains(body, "event: response.completed") { + t.Fatalf("failed stream unexpectedly completed: %s", body) + } + + recorder = postJSONWithHeaders(t, "/v1/responses", `{"model":"test-model","input":"hello","stream":true}`, map[string]string{ + headerOutputChunks: "2", + headerStreamTruncateAfter: "1", + }) + body = recorder.Body.String() + if strings.Contains(body, "event: error") || strings.Contains(body, "event: response.completed") { + t.Fatalf("truncated Responses stream unexpectedly terminated: %s", body) + } + + recorder = postJSONWithHeaders(t, "/v1/chat/completions", `{"model":"test-model","messages":[],"stream":true}`, map[string]string{ + headerOutputChunks: "2", + headerStreamTruncateAfter: "1", + }) + body = recorder.Body.String() + if strings.Contains(body, "data: [DONE]") || strings.Contains(body, `"finish_reason":"stop"`) { + t.Fatalf("truncated stream unexpectedly completed: %s", body) + } +} + +func TestBodyInjectedStatusAndStreamFailures(t *testing.T) { + recorder := postJSON(t, "/v1/responses", `{ + "model":"test-model", + "input":"hello", + "x_load_tester_status_code":503 + }`) + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusServiceUnavailable, recorder.Body.String()) + } + + recorder = postJSON(t, "/v1/embeddings", `{ + "model":"test-model", + "input":"hello", + "x_load_tester_status_code":503 + }`) + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("embedding status = %d, want %d: %s", recorder.Code, http.StatusServiceUnavailable, recorder.Body.String()) + } + + recorder = postJSON(t, "/v1/responses", `{ + "model":"test-model", + "input":"hello", + "stream":true, + "x_load_tester_output_chunks":2, + "x_load_tester_stream_error_after_chunks":1 + }`) + body := recorder.Body.String() + if !strings.Contains(body, "event: error") || strings.Contains(body, "event: response.completed") { + t.Fatalf("unexpected failed stream: %s", body) + } + + recorder = postJSON(t, "/v1/responses", `{ + "model":"test-model", + "input":"hello", + "stream":true, + "x_load_tester_output_chunks":2, + "x_load_tester_stream_truncate_after_chunks":1 + }`) + body = recorder.Body.String() + if strings.Contains(body, "event: error") || strings.Contains(body, "event: response.completed") { + t.Fatalf("unexpected truncated stream: %s", body) + } +} + +func TestConcurrencyLimit(t *testing.T) { + activeRequests.Store(0) + firstDone := make(chan *httptest.ResponseRecorder, 1) + requestFinished := make(chan struct{}) + t.Cleanup(func() { + select { + case <-requestFinished: + case <-time.After(time.Second): + } + activeRequests.Store(0) + }) + router := newRouter() + firstRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"test-model"}`)) + firstRequest.Header.Set("Content-Type", "application/json") + firstRequest.Header.Set(headerMaxConcurrency, "1") + firstRequest.Header.Set(headerTTFT, "100") + go func() { + defer close(requestFinished) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, firstRequest) + firstDone <- recorder + }() + + deadline := time.Now().Add(time.Second) + for activeRequests.Load() != 1 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if activeRequests.Load() != 1 { + t.Fatal("first request did not acquire the concurrency slot") + } + + secondRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"test-model"}`)) + secondRequest.Header.Set("Content-Type", "application/json") + secondRequest.Header.Set(headerMaxConcurrency, "1") + secondRecorder := httptest.NewRecorder() + router.ServeHTTP(secondRecorder, secondRequest) + if secondRecorder.Code != http.StatusTooManyRequests { + t.Fatalf("status = %d, want %d: %s", secondRecorder.Code, http.StatusTooManyRequests, secondRecorder.Body.String()) + } + if firstRecorder := <-firstDone; firstRecorder.Code != http.StatusOK { + t.Fatalf("first status = %d, want %d: %s", firstRecorder.Code, http.StatusOK, firstRecorder.Body.String()) + } + if activeRequests.Load() != 0 { + t.Fatalf("active requests = %d, want 0", activeRequests.Load()) + } +} + +func TestBodyConcurrencyLimit(t *testing.T) { + activeRequests.Store(0) + firstDone := make(chan *httptest.ResponseRecorder, 1) + requestFinished := make(chan struct{}) + t.Cleanup(func() { + select { + case <-requestFinished: + case <-time.After(time.Second): + } + activeRequests.Store(0) + }) + router := newRouter() + body := `{"model":"test-model","x_load_tester_max_concurrency":1,"x_load_tester_ttft_ms":100}` + firstRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + firstRequest.Header.Set("Content-Type", "application/json") + go func() { + defer close(requestFinished) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, firstRequest) + firstDone <- recorder + }() + + deadline := time.Now().Add(time.Second) + for activeRequests.Load() != 1 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if activeRequests.Load() != 1 { + t.Fatal("first request did not acquire the concurrency slot") + } + + secondRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + secondRequest.Header.Set("Content-Type", "application/json") + secondRecorder := httptest.NewRecorder() + router.ServeHTTP(secondRecorder, secondRequest) + if secondRecorder.Code != http.StatusTooManyRequests { + t.Fatalf("status = %d, want %d: %s", secondRecorder.Code, http.StatusTooManyRequests, secondRecorder.Body.String()) + } + if firstRecorder := <-firstDone; firstRecorder.Code != http.StatusOK { + t.Fatalf("first status = %d, want %d: %s", firstRecorder.Code, http.StatusOK, firstRecorder.Body.String()) + } + if activeRequests.Load() != 0 { + t.Fatalf("active requests = %d, want 0", activeRequests.Load()) + } +} + +func TestChatCompletionsSupportsJSONAndSSE(t *testing.T) { + recorder := postJSON(t, "/v1/chat/completions", `{ + "model":"test-model", + "messages":[{"role":"user","content":"hello"}] + }`) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body := responseMap(t, recorder) + if got := body["object"]; got != "chat.completion" { + t.Fatalf("object = %#v, want chat.completion", got) + } + choices := body["choices"].([]any) + message := choices[0].(map[string]any)["message"].(map[string]any) + if got := message["content"]; got != defaultChunk { + t.Fatalf("content = %#v, want %s", got, defaultChunk) + } + + recorder = postJSONWithHeaders(t, "/v1/chat/completions", `{ + "model":"test-model", + "messages":[{"role":"user","content":"hello"}], + "stream":true, + "stream_options":{"include_usage":true} + }`, map[string]string{ + headerOutputChunks: "2", + }) + if recorder.Code != http.StatusOK { + t.Fatalf("stream status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + stream := recorder.Body.String() + for _, expected := range []string{`"role":"assistant"`, `"content":"xxxx"`, `"finish_reason":"stop"`, `"usage":`, "data: [DONE]"} { + if !strings.Contains(stream, expected) { + t.Fatalf("stream missing %q: %s", expected, stream) + } + } +} + +func TestLegacyStreamsFlushOncePerFrame(t *testing.T) { + for _, test := range []struct { + name string + path string + body string + }{ + { + name: "chat completions", + path: "/v1/chat/completions", + body: `{"model":"test-model","messages":[],"stream":true}`, + }, + { + name: "completions", + path: "/v1/completions", + body: `{"model":"test-model","prompt":"hello","stream":true}`, + }, + } { + t.Run(test.name, func(t *testing.T) { + recorder := newFlushCountingRecorder() + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, test.path, strings.NewReader(test.body)) + request.Header.Set("Content-Type", "application/json") + request.Header.Set(headerOutputChunks, "2") + newRouter().ServeHTTP(recorder, request) + + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if frames := strings.Count(recorder.Body.String(), "\n\n"); recorder.flushes != frames { + t.Fatalf("flushes = %d, want %d frames", recorder.flushes, frames) + } + }) + } +} + +func TestCompletionsAndModels(t *testing.T) { + recorder := postJSON(t, "/v1/completions", `{"model":"test-model","prompt":"hello"}`) + if recorder.Code != http.StatusOK { + t.Fatalf("completion status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body := responseMap(t, recorder) + if got := body["object"]; got != "text_completion" { + t.Fatalf("object = %#v, want text_completion", got) + } + + request := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/v1/models", nil) + recorder = httptest.NewRecorder() + newRouter().ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("models status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + data := responseMap(t, recorder)["data"].([]any) + if got := data[0].(map[string]any)["id"]; got != defaultModel { + t.Fatalf("model id = %#v, want %s", got, defaultModel) + } +} + +func TestEmbeddingsSupportFloatAndBase64(t *testing.T) { + recorder := postJSON(t, "/v1/embeddings", `{ + "model":"test-model", + "input":["one","two"], + "encoding_format":"float" + }`) + if recorder.Code != http.StatusOK { + t.Fatalf("float status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body := responseMap(t, recorder) + data := body["data"].([]any) + if len(data) != 2 { + t.Fatalf("embedding count = %d, want 2", len(data)) + } + vector := data[0].(map[string]any)["embedding"].([]any) + if len(vector) != 3 { + t.Fatalf("vector length = %d, want 3", len(vector)) + } + + recorder = postJSON(t, "/v1/embeddings", `{ + "model":"test-model", + "input":"one", + "encoding_format":"base64" + }`) + if recorder.Code != http.StatusOK { + t.Fatalf("base64 status = %d, want %d: %s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + body = responseMap(t, recorder) + encoded := body["data"].([]any)[0].(map[string]any)["embedding"].(string) + decoded, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + t.Fatalf("base64 embedding decode failed: %v", err) + } + if len(decoded) != 12 { + t.Fatalf("base64 embedding length = %d, want 12", len(decoded)) + } +} + +func TestEmbeddingsRejectEmptyInput(t *testing.T) { + recorder := postJSON(t, "/v1/embeddings", `{"model":"test-model","input":[]}`) + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d: %s", recorder.Code, http.StatusBadRequest, recorder.Body.String()) + } +} + +func postJSON(t *testing.T, path, body string) *httptest.ResponseRecorder { + return postJSONWithHeaders(t, path, body, nil) +} + +func postJSONWithHeaders(t *testing.T, path, body string, headers map[string]string) *httptest.ResponseRecorder { + t.Helper() + request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, path, strings.NewReader(body)) + request.Header.Set("Content-Type", "application/json") + for name, value := range headers { + request.Header.Set(name, value) + } + recorder := httptest.NewRecorder() + newRouter().ServeHTTP(recorder, request) + return recorder +} + +func responseMap(t *testing.T, recorder *httptest.ResponseRecorder) map[string]any { + t.Helper() + var response map[string]any + if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { + t.Fatalf("response JSON unmarshal failed: %v: %s", err, recorder.Body.String()) + } + return response +} + +type sseEvent struct { + name string + data map[string]any +} + +type flushCountingRecorder struct { + *httptest.ResponseRecorder + flushes int + onWrite func([]byte) +} + +func newFlushCountingRecorder() *flushCountingRecorder { + return &flushCountingRecorder{ResponseRecorder: httptest.NewRecorder()} +} + +func (recorder *flushCountingRecorder) Write(frame []byte) (int, error) { + if recorder.onWrite != nil { + recorder.onWrite(frame) + } + return recorder.ResponseRecorder.Write(frame) +} + +func (recorder *flushCountingRecorder) Flush() { + recorder.flushes++ +} + +type delayedDeltaRecorder struct { + *flushCountingRecorder + delay time.Duration + deltaWrites []time.Time +} + +func (recorder *delayedDeltaRecorder) Write(frame []byte) (int, error) { + isDelta := bytes.Contains(frame, []byte("event: response.output_text.delta\n")) + if isDelta && len(recorder.deltaWrites) == 1 { + time.Sleep(recorder.delay) + } + written, err := recorder.flushCountingRecorder.Write(frame) + if isDelta { + recorder.deltaWrites = append(recorder.deltaWrites, time.Now()) + } + return written, err +} + +type partialResponseWriter struct { + header http.Header + body bytes.Buffer + maxWrite int + flushes int +} + +func (writer *partialResponseWriter) Header() http.Header { + return writer.header +} + +func (writer *partialResponseWriter) Write(frame []byte) (int, error) { + written := len(frame) + if writer.maxWrite > 0 && written > writer.maxWrite { + written = writer.maxWrite + } + _, _ = writer.body.Write(frame[:written]) + return written, nil +} + +func (writer *partialResponseWriter) WriteHeader(int) {} + +func (writer *partialResponseWriter) Flush() { + writer.flushes++ +} + +type discardResponseWriter struct { + header http.Header +} + +func (writer *discardResponseWriter) Header() http.Header { + return writer.header +} + +func (writer *discardResponseWriter) Write(frame []byte) (int, error) { + return len(frame), nil +} + +func (writer *discardResponseWriter) WriteHeader(int) {} + +func (writer *discardResponseWriter) Flush() {} + +func parseResponsesSSEEvents(t *testing.T, stream string) []sseEvent { + t.Helper() + frames := strings.Split(strings.TrimSuffix(stream, "\n\n"), "\n\n") + events := make([]sseEvent, 0, len(frames)) + for _, frame := range frames { + var event sseEvent + for _, line := range strings.Split(frame, "\n") { + switch { + case strings.HasPrefix(line, "event: "): + event.name = strings.TrimPrefix(line, "event: ") + case strings.HasPrefix(line, "data: "): + if err := json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &event.data); err != nil { + t.Fatalf("SSE data JSON unmarshal failed: %v: %s", err, line) + } + } + } + if event.name == "" || event.data == nil { + t.Fatalf("invalid SSE frame: %q", frame) + } + events = append(events, event) + } + return events +} + +func sseSequence(t *testing.T, event sseEvent) int { + t.Helper() + sequence, ok := event.data["sequence_number"].(float64) + if !ok { + t.Fatalf("sequence_number = %#v, want number", event.data["sequence_number"]) + } + return int(sequence) +} + +func eventText(t *testing.T, event sseEvent, field string) string { + t.Helper() + text, ok := event.data[field].(string) + if !ok { + t.Fatalf("%s = %#v, want string", field, event.data[field]) + } + return text +} + +func nestedEventText(t *testing.T, event sseEvent, field string) string { + t.Helper() + part, ok := event.data[field].(map[string]any) + if !ok { + t.Fatalf("%s = %#v, want object", field, event.data[field]) + } + text, ok := part["text"].(string) + if !ok { + t.Fatalf("%s.text = %#v, want string", field, part["text"]) + } + return text +} + +func nestedResponseText(t *testing.T, response map[string]any) string { + t.Helper() + output, ok := response["output"].([]any) + if !ok || len(output) != 1 { + t.Fatalf("output = %#v, want one item", response["output"]) + } + item, ok := output[0].(map[string]any) + if !ok { + t.Fatalf("output item = %#v, want object", output[0]) + } + content, ok := item["content"].([]any) + if !ok || len(content) != 1 { + t.Fatalf("content = %#v, want one item", item["content"]) + } + part, ok := content[0].(map[string]any) + if !ok { + t.Fatalf("content part = %#v, want object", content[0]) + } + text, ok := part["text"].(string) + if !ok { + t.Fatalf("content text = %#v, want string", part["text"]) + } + return text +} diff --git a/examples/function-samples/openai-compatible-sample/http-server/openai_client_check.py b/examples/function-samples/openai-compatible-sample/http-server/openai_client_check.py new file mode 100644 index 0000000000..ef0d714a03 --- /dev/null +++ b/examples/function-samples/openai-compatible-sample/http-server/openai_client_check.py @@ -0,0 +1,118 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Check the sample through the public OpenAI Python client.""" + +import os +from urllib.parse import urlparse + +import httpx +from openai import OpenAI + + +BASE_URL = os.environ.get("OPENAI_BASE_URL", "http://127.0.0.1:8000/v1") +API_KEY = os.environ.get("OPENAI_API_KEY", "not-needed") +if API_KEY != "not-needed" and urlparse(BASE_URL).scheme != "https": + raise ValueError("OPENAI_BASE_URL must use HTTPS when OPENAI_API_KEY is set") + +CLIENT = OpenAI( + api_key=API_KEY, + base_url=BASE_URL, + http_client=httpx.Client(follow_redirects=False), + _strict_response_validation=True, +) + + +def main() -> None: + response = CLIENT.responses.create( + model="test-model", + input="hello", + extra_headers={ + "X-Load-Tester-Chunk": "token", + "X-Load-Tester-Output-Chunks": "3", + }, + ) + if response.output_text != "tokentokentoken": + raise RuntimeError("unexpected Responses output") + + response_events = list( + CLIENT.responses.create( + model="test-model", + input="hello", + stream=True, + extra_headers={ + "X-Load-Tester-Chunk": "stream", + "X-Load-Tester-Output-Chunks": "2", + }, + ) + ) + response_text = "".join( + event.delta + for event in response_events + if event.type == "response.output_text.delta" + ) + if response_text != "streamstream": + raise RuntimeError("unexpected streamed Responses output") + + chat = CLIENT.chat.completions.create( + model="test-model", + messages=[{"role": "user", "content": "hello"}], + ) + if chat.choices[0].message.content != "xxxx": + raise RuntimeError("unexpected chat completion output") + + chat_chunks = list( + CLIENT.chat.completions.create( + model="test-model", + messages=[{"role": "user", "content": "hello"}], + stream=True, + extra_headers={ + "X-Load-Tester-Chunk": "chat", + "X-Load-Tester-Output-Chunks": "2", + }, + ) + ) + chat_text = "".join( + choice.delta.content or "" + for chunk in chat_chunks + for choice in chunk.choices + ) + if chat_text != "chatchat": + raise RuntimeError("unexpected streamed chat completion output") + + completion = CLIENT.completions.create(model="test-model", prompt="hello") + if completion.choices[0].text != "xxxx": + raise RuntimeError("unexpected completion output") + + embedding = CLIENT.embeddings.create( + model="test-model", + input=["one", "two"], + encoding_format="float", + ) + if len(embedding.data) != 2: + raise RuntimeError("unexpected embedding count") + + models = CLIENT.models.list() + if not any(model.id == "test-model" for model in models.data): + raise RuntimeError("test-model is missing from the model list") + if CLIENT.models.retrieve("test-model").id != "test-model": + raise RuntimeError("unexpected retrieved model") + + print("OpenAI client compatibility check passed") + + +if __name__ == "__main__": + main() diff --git a/examples/load-tests/README.md b/examples/load-tests/README.md index 6c298f2a3b..0520f009cd 100644 --- a/examples/load-tests/README.md +++ b/examples/load-tests/README.md @@ -31,6 +31,7 @@ tasks/ NVCT task load tests | `supreme_large_response_test.js` | Large payload responses. | | `oai_compatible_llm_stream_load_test.js` | Streaming OpenAI-compatible LLM completions. | | `oai_compatible_llm_load_test.js` | Non-streaming OpenAI-compatible LLM completions. | +| `oai_compatible_responses_sse_load_test.js` | Streaming OpenAI Responses API benchmark with TTFT, ITL, and throughput metrics. | | `oai_list_models_load_test.js` | OpenAI-compatible model listing endpoint. | | `sdxl_load_test.js` | Stable Diffusion XL image generation. | | `nvcf_health_load_test.js` | NVCF health endpoint. | @@ -106,7 +107,26 @@ These helpers use the `ngc` CLI today and target cloud NVCF. Porting them to sel | Variable | Description | |----------|-------------| | `OAI_COMPAT_URL` | OpenAI-compatible API endpoint. | +| `TOKEN` | Optional Bearer token. Non-loopback endpoints must use HTTPS. | | `LLM_MODEL_NAME` | Model identifier. | +| `OPENAI_RESPONSES_PROFILE` | `calibration` (default) or `load` for the Responses SSE benchmark. | +| `OPENAI_RESPONSES_VUS` | Virtual users for the Responses SSE benchmark. Defaults to 1 for calibration and 10 for load. | +| `OPENAI_RESPONSES_ITERATIONS` | Per-VU iterations for calibration. Defaults to 10. | +| `OPENAI_RESPONSES_MAX_DURATION` | Maximum duration for calibration. Defaults to `10m`. | +| `OPENAI_RESPONSES_DURATION` | Test duration for load. Defaults to `30s`. | +| `OPENAI_RESPONSES_TOKENS_PER_CHUNK` | Declared synthetic tokens in each output chunk. Defaults to 1. | +| `OPENAI_RESPONSES_EXPECTED_DELTAS` | Required text-delta count. Defaults to the configured output chunks for calibration and is disabled for load. | +| `OPENAI_RESPONSES_CALIBRATION_TOLERANCE_MS` | Allowed early-observation tolerance for calibration timing checks. Defaults to 10 ms. | +| `OPENAI_RESPONSES_INPUT` | Responses API input string. Defaults to `benchmark`. | +| `LOAD_TESTER_QUEUE_DELAY_MS` | Maps to `X-Load-Tester-Queue-Delay-Ms`. Sent by default only in calibration. | +| `LOAD_TESTER_TTFT_MS` | Maps to `X-Load-Tester-TTFT-Ms`. Sent by default only in calibration. | +| `LOAD_TESTER_TTFT_JITTER_MS` | Maps to `X-Load-Tester-TTFT-Jitter-Ms`. Sent by default only in calibration. | +| `LOAD_TESTER_ITL_MS` | Maps to `X-Load-Tester-ITL-Ms`. Sent by default only in calibration. | +| `LOAD_TESTER_ITL_JITTER_MS` | Maps to `X-Load-Tester-ITL-Jitter-Ms`. Sent by default only in calibration. | +| `LOAD_TESTER_OUTPUT_CHUNKS` | Maps to `X-Load-Tester-Output-Chunks`. Defaults to 8 and must not exceed the sample's startup chunk limit. | +| `LOAD_TESTER_CHUNK` | Maps to `X-Load-Tester-Chunk`. Defaults to `xxxx`. | +| `LOAD_TESTER_STREAM_ERROR_AFTER_CHUNKS` | Maps to `X-Load-Tester-Stream-Error-After-Chunks` for failure-path validation. | +| `LOAD_TESTER_STREAM_TRUNCATE_AFTER_CHUNKS` | Maps to `X-Load-Tester-Stream-Truncate-After-Chunks` for truncated-stream validation. | ### Multi-Endpoint Tests @@ -169,6 +189,79 @@ Then run with the local binary: -e TOKEN=$TOKEN -e HTTP_SUPREME_NVCF_URL=$HTTP_SUPREME_NVCF_URL ``` +### OpenAI Responses SSE Benchmark + +`oai_compatible_responses_sse_load_test.js` measures two start latencies: the +first SSE event and the first `response.output_text.delta`. It records one ITL +sample between each pair of text-delta events, then reports output chunks per +second and declared tokens per second. The output chunk and declared token +counters also provide aggregate rates for streams whose delta timestamps share +the same millisecond. A stream succeeds only after HTTP 200, +`response.completed`, no transport or protocol error, and an optional expected +delta count. For this script, `OAI_COMPAT_URL` must be the full +`/v1/responses` endpoint URL. + +Metrics: + +- `openai_responses_first_sse_event_ms`, `openai_responses_ttft_ms`, and `openai_responses_itl_ms` +- `openai_responses_output_chunks_per_second` and `openai_responses_declared_tokens_per_second` +- `openai_responses_stream_duration_ms`, `openai_responses_stream_success`, and stream outcome counters + +The default calibration profile sends eight `xxxx` chunks with 200 ms TTFT and +50 ms ITL. It validates that those delays are not observed materially early: + +```bash +./k6 run functions/oai_compatible_responses_sse_load_test.js \ + -e OAI_COMPAT_URL=http://127.0.0.1:8000/v1/responses +``` + +The load profile leaves queue, TTFT, and ITL delays unset unless their +`LOAD_TESTER_*` variables are supplied: + +```bash +./k6 run functions/oai_compatible_responses_sse_load_test.js \ + -e OAI_COMPAT_URL=$OAI_COMPAT_URL \ + -e OPENAI_RESPONSES_PROFILE=load \ + -e OPENAI_RESPONSES_VUS=10 \ + -e OPENAI_RESPONSES_DURATION=30s +``` + +For a 60-second per-connection capacity run at 5 ms ITL, start the sample with +`LOAD_TESTER_MAX_OUTPUT_CHUNKS=12000`, then use calibration mode with 12000 +chunks and two iterations per VU. Raise the generator file-descriptor limit +before a high-concurrency run. Set `OAI_COMPAT_URL` to the full +`/v1/responses` endpoint. + +```bash +ulimit -n 65536 + +./k6 run functions/oai_compatible_responses_sse_load_test.js \ + --summary-export responses-sse-summary.json \ + -e OAI_COMPAT_URL=$OAI_COMPAT_URL \ + -e OPENAI_RESPONSES_PROFILE=calibration \ + -e OPENAI_RESPONSES_VUS=1024 \ + -e OPENAI_RESPONSES_ITERATIONS=2 \ + -e OPENAI_RESPONSES_MAX_DURATION=5m \ + -e OPENAI_RESPONSES_EXPECTED_DELTAS=12000 \ + -e OPENAI_RESPONSES_CALIBRATION_TOLERANCE_MS=1 \ + -e LOAD_TESTER_QUEUE_DELAY_MS=0 \ + -e LOAD_TESTER_TTFT_MS=1 \ + -e LOAD_TESTER_TTFT_JITTER_MS=0 \ + -e LOAD_TESTER_ITL_MS=5 \ + -e LOAD_TESTER_ITL_JITTER_MS=0 \ + -e LOAD_TESTER_CHUNK=xxxx \ + -e LOAD_TESTER_OUTPUT_CHUNKS=12000 +``` + +The high-concurrency profile opens a new connection for each iteration. Check +generator file descriptors, sockets, CPU, and network saturation before +interpreting a failure as target capacity. + +`openai_responses_declared_tokens_per_second` is synthetic. It multiplies the +observed chunk rate by `OPENAI_RESPONSES_TOKENS_PER_CHUNK`; `xxxx` is not a +tokenizer-derived token. Use `openai_responses_output_chunks_per_second` when +the chunk-to-token mapping is unknown. + ## Resources - [k6 Documentation](https://grafana.com/docs/k6/latest/) diff --git a/examples/load-tests/functions/oai_compatible_responses_sse_load_test.js b/examples/load-tests/functions/oai_compatible_responses_sse_load_test.js new file mode 100644 index 0000000000..5d9bc072ef --- /dev/null +++ b/examples/load-tests/functions/oai_compatible_responses_sse_load_test.js @@ -0,0 +1,361 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +import { check } from 'k6' +import sse from 'k6/x/sse' +import { Counter, Rate, Trend } from 'k6/metrics' + +const PROFILE_CALIBRATION = 'calibration' +const PROFILE_LOAD = 'load' + +const firstSSEEventMs = new Trend('openai_responses_first_sse_event_ms', true) +const ttftMs = new Trend('openai_responses_ttft_ms', true) +const itlMs = new Trend('openai_responses_itl_ms', true) +const outputChunksPerSecond = new Trend('openai_responses_output_chunks_per_second') +const declaredTokensPerSecond = new Trend('openai_responses_declared_tokens_per_second') +const streamDurationMs = new Trend('openai_responses_stream_duration_ms', true) +const emittedOutputChunks = new Counter('openai_responses_output_chunks') +const declaredOutputTokens = new Counter('openai_responses_declared_output_tokens') +const streamsCompleted = new Counter('openai_responses_streams_completed') +const streamsFailed = new Counter('openai_responses_streams_failed') +const streamsTruncated = new Counter('openai_responses_streams_truncated') +const protocolErrors = new Counter('openai_responses_protocol_errors') +const expectedDeltaMismatches = new Counter('openai_responses_expected_delta_mismatches') +const streamSuccess = new Rate('openai_responses_stream_success') + +const config = buildConfig() + +export const options = buildOptions() + +export function setup() { + if (config.url === '') { + throw new Error('OAI_COMPAT_URL must be the full /v1/responses endpoint URL') + } +} + +export default function () { + const requestStartedAt = Date.now() + let firstEventAt = null + let firstDeltaAt = null + let lastDeltaAt = null + let deltaCount = 0 + let completed = false + let protocolError = false + let transportError = false + let earlyITL = false + + const response = sse.open(config.url, requestParams(), function (client) { + client.on('event', function (event) { + const eventAt = Date.now() + if (firstEventAt === null) { + firstEventAt = eventAt + } + + if (event.name === 'error') { + protocolError = true + client.close() + return + } + + if (event.name === 'response.output_text.delta') { + const data = parseEvent(event) + if (data === null || data.type !== event.name || typeof data.delta !== 'string') { + protocolError = true + return + } + + if (firstDeltaAt === null) { + firstDeltaAt = eventAt + } else { + const interval = eventAt - lastDeltaAt + itlMs.add(interval) + if (config.profile === PROFILE_CALIBRATION && interval < calibrationITLMinimum()) { + earlyITL = true + } + } + + lastDeltaAt = eventAt + deltaCount += 1 + return + } + + if (event.name === 'response.completed') { + const data = parseEvent(event) + if (data === null || data.type !== event.name) { + protocolError = true + } else { + completed = true + } + client.close() + } + }) + + client.on('error', function () { + transportError = true + client.close() + }) + }) + + const streamFinishedAt = Date.now() + const status = response && response.status ? response.status : 0 + const responseError = response && response.error ? response.error : null + const hasTransportError = transportError || responseError !== null + const firstEventDuration = firstEventAt === null ? null : firstEventAt - requestStartedAt + const ttftDuration = firstDeltaAt === null ? null : firstDeltaAt - requestStartedAt + const outputDuration = firstDeltaAt === null || lastDeltaAt === null ? 0 : lastDeltaAt - firstDeltaAt + const expectedDeltaCount = config.expectedDeltas === 0 || deltaCount === config.expectedDeltas + const success = status === 200 && completed && !protocolError && !hasTransportError && expectedDeltaCount + + streamDurationMs.add(streamFinishedAt - requestStartedAt) + if (firstEventDuration !== null) { + firstSSEEventMs.add(firstEventDuration) + } + if (ttftDuration !== null) { + ttftMs.add(ttftDuration) + } + if (success) { + emittedOutputChunks.add(deltaCount) + declaredOutputTokens.add(deltaCount * config.tokensPerChunk) + } + if (success && deltaCount > 1 && outputDuration > 0) { + const chunksPerSecond = (deltaCount - 1) * 1000 / outputDuration + outputChunksPerSecond.add(chunksPerSecond) + declaredTokensPerSecond.add(chunksPerSecond * config.tokensPerChunk) + } + + if (completed) { + streamsCompleted.add(1) + } + if (protocolError) { + protocolErrors.add(1) + } + if (!completed && !protocolError && !hasTransportError && status === 200) { + streamsTruncated.add(1) + } + if (!expectedDeltaCount) { + expectedDeltaMismatches.add(1) + } + if (!success) { + streamsFailed.add(1) + } + streamSuccess.add(success) + + const checks = { + 'responses stream returns HTTP 200': function (result) { + return result.status === 200 + }, + 'responses stream completes': function (result) { + return result.completed + }, + 'responses stream has no protocol error': function (result) { + return !result.protocolError + }, + 'responses stream has no transport error': function (result) { + return !result.transportError + }, + } + if (config.expectedDeltas !== 0) { + checks['responses stream emits the expected delta count'] = function (result) { + return result.expectedDeltaCount + } + } + if (config.profile === PROFILE_CALIBRATION) { + checks['calibration first SSE event is not early'] = function (result) { + return result.firstEventDuration !== null && result.firstEventDuration >= calibrationStartMinimum() + } + checks['calibration first token is not early'] = function (result) { + return result.ttftDuration !== null && result.ttftDuration >= calibrationStartMinimum() + } + checks['calibration inter-token latency is not early'] = function (result) { + return result.deltaCount > 1 && !result.earlyITL + } + } + + check({ + status: status, + completed: completed, + protocolError: protocolError, + transportError: hasTransportError, + expectedDeltaCount: expectedDeltaCount, + firstEventDuration: firstEventDuration, + ttftDuration: ttftDuration, + deltaCount: deltaCount, + earlyITL: earlyITL, + }, checks) +} + +function buildConfig() { + const profile = stringEnv('OPENAI_RESPONSES_PROFILE', PROFILE_CALIBRATION).toLowerCase() + if (profile !== PROFILE_CALIBRATION && profile !== PROFILE_LOAD) { + throw new Error('OPENAI_RESPONSES_PROFILE must be calibration or load') + } + + const isCalibration = profile === PROFILE_CALIBRATION + const outputChunks = integerEnv('LOAD_TESTER_OUTPUT_CHUNKS', 8, 1) + const tokensPerChunk = integerEnv('OPENAI_RESPONSES_TOKENS_PER_CHUNK', 1, 1) + const chunk = stringEnv('LOAD_TESTER_CHUNK', 'xxxx') + if (chunk === '') { + throw new Error('LOAD_TESTER_CHUNK must not be empty') + } + + return { + profile: profile, + url: stringEnv('OAI_COMPAT_URL', ''), + token: stringEnv('TOKEN', ''), + model: stringEnv('LLM_MODEL_NAME', 'test-model'), + input: stringEnv('OPENAI_RESPONSES_INPUT', 'benchmark'), + vus: integerEnv('OPENAI_RESPONSES_VUS', isCalibration ? 1 : 10, 1), + iterations: integerEnv('OPENAI_RESPONSES_ITERATIONS', 10, 1), + calibrationMaxDuration: stringEnv('OPENAI_RESPONSES_MAX_DURATION', '10m'), + duration: stringEnv('OPENAI_RESPONSES_DURATION', '30s'), + expectedDeltas: integerEnv('OPENAI_RESPONSES_EXPECTED_DELTAS', isCalibration ? outputChunks : 0, 0), + tokensPerChunk: tokensPerChunk, + outputChunks: outputChunks, + chunk: chunk, + queueDelayMs: optionalIntegerEnv('LOAD_TESTER_QUEUE_DELAY_MS', 0, isCalibration), + ttftDelayMs: optionalIntegerEnv('LOAD_TESTER_TTFT_MS', 200, isCalibration), + ttftJitterMs: optionalIntegerEnv('LOAD_TESTER_TTFT_JITTER_MS', 0, isCalibration), + itlDelayMs: optionalIntegerEnv('LOAD_TESTER_ITL_MS', 50, isCalibration), + itlJitterMs: optionalIntegerEnv('LOAD_TESTER_ITL_JITTER_MS', 0, isCalibration), + streamErrorAfterChunks: optionalIntegerEnv('LOAD_TESTER_STREAM_ERROR_AFTER_CHUNKS', 0, false), + streamTruncateAfterChunks: optionalIntegerEnv('LOAD_TESTER_STREAM_TRUNCATE_AFTER_CHUNKS', 0, false), + calibrationToleranceMs: integerEnv('OPENAI_RESPONSES_CALIBRATION_TOLERANCE_MS', 10, 0), + } +} + +function buildOptions() { + const scenarios = {} + if (config.profile === PROFILE_CALIBRATION) { + scenarios.calibration = { + executor: 'per-vu-iterations', + vus: config.vus, + iterations: config.iterations, + maxDuration: config.calibrationMaxDuration, + } + } else { + scenarios.load = { + executor: 'constant-vus', + vus: config.vus, + duration: config.duration, + } + } + + return { + scenarios: scenarios, + thresholds: { + checks: ['rate==1'], + openai_responses_stream_success: ['rate==1'], + }, + } +} + +function requestParams() { + const headers = { + 'Accept': 'text/event-stream', + 'Content-Type': 'application/json', + 'X-Load-Tester-Chunk': config.chunk, + 'X-Load-Tester-Output-Chunks': String(config.outputChunks), + } + if (config.token !== '') { + validateCredentialURL(config.url) + headers.Authorization = `Bearer ${config.token}` + } + addHeader(headers, 'X-Load-Tester-Queue-Delay-Ms', config.queueDelayMs) + addHeader(headers, 'X-Load-Tester-TTFT-Ms', config.ttftDelayMs) + addHeader(headers, 'X-Load-Tester-TTFT-Jitter-Ms', config.ttftJitterMs) + addHeader(headers, 'X-Load-Tester-ITL-Ms', config.itlDelayMs) + addHeader(headers, 'X-Load-Tester-ITL-Jitter-Ms', config.itlJitterMs) + addHeader(headers, 'X-Load-Tester-Stream-Error-After-Chunks', config.streamErrorAfterChunks) + addHeader(headers, 'X-Load-Tester-Stream-Truncate-After-Chunks', config.streamTruncateAfterChunks) + + return { + method: 'POST', + body: JSON.stringify({ + model: config.model, + input: config.input, + stream: true, + }), + headers: headers, + tags: { + name: 'OpenAIResponsesSSE', + profile: config.profile, + }, + } +} + +function validateCredentialURL(value) { + const parsed = /^([a-z][a-z0-9+.-]*):\/\/([^/?#:]+|\[[^\]]+\])(?::\d+)?(?:[/?#]|$)/i.exec(value) + if (parsed === null) { + throw new Error('OAI_COMPAT_URL must be an absolute URL when TOKEN is set') + } + + const scheme = parsed[1].toLowerCase() + const host = parsed[2].replace(/^\[|\]$/g, '').toLowerCase() + if (scheme === 'https' || isLoopbackHost(host)) { + return + } + throw new Error('OAI_COMPAT_URL must use HTTPS when TOKEN is set, except for a loopback host') +} + +function isLoopbackHost(host) { + return host === 'localhost' || host === '::1' || /^127(?:\.\d{1,3}){3}$/.test(host) +} + +function addHeader(headers, name, value) { + if (value !== undefined) { + headers[name] = String(value) + } +} + +function parseEvent(event) { + try { + return JSON.parse(event.data) + } catch (error) { + return null + } +} + +function calibrationStartMinimum() { + return Math.max(0, (config.queueDelayMs || 0) + (config.ttftDelayMs || 0) - config.calibrationToleranceMs) +} + +function calibrationITLMinimum() { + return Math.max(0, (config.itlDelayMs || 0) - config.calibrationToleranceMs) +} + +function stringEnv(name, defaultValue) { + const value = __ENV[name] + return value === undefined || value === '' ? defaultValue : value +} + +function integerEnv(name, defaultValue, minimum) { + const value = optionalIntegerEnv(name, defaultValue, true) + if (value < minimum) { + throw new Error(`${name} must be at least ${minimum}`) + } + return value +} + +function optionalIntegerEnv(name, defaultValue, useDefault) { + const raw = __ENV[name] + if (raw === undefined || raw === '') { + return useDefault ? defaultValue : undefined + } + const value = Number(raw) + if (!isFinite(value) || Math.floor(value) !== value || value < 0) { + throw new Error(`${name} must be a non-negative integer`) + } + return value +}