Studio: expose Anthropic / OpenAI sampling knobs per provider
Adds the missing sampling parameters that the upstream APIs accept and that Studio's chat UI previously hid. Each knob is gated per provider so the picker never offers a field the upstream would 400 on, and the per-provider stream functions translate / drop fields to match each API's naming. New `InferenceParams` fields (round-trip through PersistedInferenceParams and the chat-settings server store automatically): - frequencyPenalty (-2..2): OpenAI Chat Completions only. - seed (int | null): OpenAI Chat + OpenAI-compat local backends. - stop (string[]): all OpenAI Chat + Anthropic Messages. Backend truncates to 4 entries on OpenAI Chat per docs and renames to `stop_sequences` on Anthropic. - serviceTier (auto|default|flex|priority|scale|standard_only): per-provider enum sets resolved by getServiceTierOptions. - parallelToolCalls (bool, default true): forwarded as `parallel_tool_calls` on both OpenAI APIs and inverted into `disable_parallel_tool_use` on Anthropic. OpenAI Responses (gpt-5.x / o3) explicitly drops frequencyPenalty / seed / stop alongside the existing temperature / top_p drop, since the upstream 400s on all of them. service_tier on Responses accepts a subset (no `scale`) which the dispatch already enforces. UI rows land in the existing Sampling section of the chat settings sheet using ParamSlider (frequency penalty), a numeric Input (seed), a new chips editor `StopSequencesInput` (stop), Select (service tier), and Switch (parallel tool calls). Each row's visibility follows the new ProviderCapabilities flag. Tests pin the gating contract: stop_sequences renamed on Anthropic, 4-entry truncation on OpenAI Chat, every Responses-rejected field dropped, schema-level validation for the service_tier Literal and frequency_penalty range. Plan: plans/hashed-riding-porcupine.md
This commit is contained in:
parent
83b20976f7
commit
91d04741ff
11 changed files with 1002 additions and 4 deletions
|
|
@ -11,7 +11,7 @@ Anthropic uses native Messages API with translation in this client.
|
|||
import json as _json
|
||||
import re
|
||||
import time
|
||||
from typing import Any, AsyncGenerator, Literal, NamedTuple, Optional
|
||||
from typing import Any, AsyncGenerator, Literal, NamedTuple, Optional, Union
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
|
@ -348,6 +348,11 @@ class ExternalProviderClient:
|
|||
anthropic_code_exec_container_id: Optional[str] = None,
|
||||
prompt_cache_ttl: Optional[str] = None,
|
||||
compaction_threshold: Optional[int] = None,
|
||||
frequency_penalty: Optional[float] = None,
|
||||
seed: Optional[int] = None,
|
||||
stop: Optional[Union[str, list[str]]] = None,
|
||||
service_tier: Optional[str] = None,
|
||||
parallel_tool_calls: Optional[bool] = None,
|
||||
stream: bool = True,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
|
|
@ -360,6 +365,12 @@ class ExternalProviderClient:
|
|||
supplies a value the provider accepts — the frontend's
|
||||
provider-capability map already filters these per provider, so we
|
||||
treat them as opt-in here.
|
||||
|
||||
``frequency_penalty``, ``seed``, ``stop``, ``service_tier``,
|
||||
``parallel_tool_calls`` follow the same rule: the per-provider
|
||||
stream helpers silently drop fields the upstream API does not
|
||||
accept (e.g. Responses rejects all of seed / frequency / stop;
|
||||
Anthropic does not implement seed / frequency / logprobs).
|
||||
"""
|
||||
if not self._is_openai_compatible():
|
||||
async for line in self._stream_anthropic(
|
||||
|
|
@ -376,6 +387,9 @@ class ExternalProviderClient:
|
|||
anthropic_code_exec_container_id,
|
||||
prompt_cache_ttl,
|
||||
compaction_threshold,
|
||||
stop = stop,
|
||||
service_tier = service_tier,
|
||||
parallel_tool_calls = parallel_tool_calls,
|
||||
):
|
||||
yield line
|
||||
return
|
||||
|
|
@ -398,6 +412,8 @@ class ExternalProviderClient:
|
|||
enable_prompt_caching,
|
||||
openai_code_exec_container_id,
|
||||
compaction_threshold,
|
||||
service_tier = service_tier,
|
||||
parallel_tool_calls = parallel_tool_calls,
|
||||
):
|
||||
yield line
|
||||
return
|
||||
|
|
@ -438,6 +454,38 @@ class ExternalProviderClient:
|
|||
else:
|
||||
body["max_tokens"] = max_tokens
|
||||
|
||||
# Optional sampling extensions (added in #5XXX). Only forwarded
|
||||
# when the caller passed a value. Each upstream provider that
|
||||
# 400s on the field appears in `body_omit` (see providers.py)
|
||||
# so the registry-driven drop loop below removes them before
|
||||
# the request hits the wire. The Responses path
|
||||
# (_stream_openai_responses) drops these explicitly because it
|
||||
# never reaches this body construction.
|
||||
if frequency_penalty is not None:
|
||||
body["frequency_penalty"] = frequency_penalty
|
||||
if seed is not None:
|
||||
body["seed"] = seed
|
||||
if stop is not None:
|
||||
# OpenAI Chat caps the list at 4 entries. Truncate with a
|
||||
# warning rather than letting the upstream 400; users editing
|
||||
# chips in the picker shouldn't get a cryptic API error.
|
||||
if isinstance(stop, list) and len(stop) > 4:
|
||||
logger.warning(
|
||||
"stop sequences truncated to 4 entries "
|
||||
"(received %d, OpenAI's hard cap is 4)",
|
||||
len(stop),
|
||||
)
|
||||
body["stop"] = stop[:4]
|
||||
elif isinstance(stop, list) and len(stop) == 0:
|
||||
# Empty list = unset; don't ship `stop: []`.
|
||||
pass
|
||||
else:
|
||||
body["stop"] = stop
|
||||
if service_tier is not None:
|
||||
body["service_tier"] = service_tier
|
||||
if parallel_tool_calls is not None:
|
||||
body["parallel_tool_calls"] = parallel_tool_calls
|
||||
|
||||
# Strip body fields a provider's registry entry declares unusable —
|
||||
# reasoning-class models that lock these to fixed defaults (e.g.
|
||||
# Kimi k2.5/k2.6 only accept temperature=1, top_p=1) 400 otherwise.
|
||||
|
|
@ -1186,6 +1234,10 @@ class ExternalProviderClient:
|
|||
anthropic_code_exec_container_id: Optional[str] = None,
|
||||
prompt_cache_ttl: Optional[str] = None,
|
||||
compaction_threshold: Optional[int] = None,
|
||||
*,
|
||||
stop: Optional[Union[str, list[str]]] = None,
|
||||
service_tier: Optional[str] = None,
|
||||
parallel_tool_calls: Optional[bool] = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Call the Anthropic Messages API and translate its SSE to OpenAI format.
|
||||
|
|
@ -1347,6 +1399,29 @@ class ExternalProviderClient:
|
|||
body["temperature"] = temperature
|
||||
if top_k is not None and top_k > 0 and not sampling_removed:
|
||||
body["top_k"] = top_k
|
||||
|
||||
# Optional sampling extensions. Anthropic has no
|
||||
# frequency_penalty / seed / logprobs equivalents, so those are
|
||||
# silently dropped by virtue of not being forwarded from
|
||||
# stream_chat_completion. The three knobs Anthropic does accept
|
||||
# land here:
|
||||
# stop → stop_sequences (renamed)
|
||||
# service_tier → service_tier (auto|standard_only only)
|
||||
# parallel_tool_calls → disable_parallel_tool_use (inverted)
|
||||
if stop is not None:
|
||||
sequences: list[str]
|
||||
if isinstance(stop, str):
|
||||
sequences = [stop] if stop else []
|
||||
else:
|
||||
sequences = [s for s in stop if isinstance(s, str) and s]
|
||||
if sequences:
|
||||
body["stop_sequences"] = sequences
|
||||
if service_tier in ("auto", "standard_only"):
|
||||
body["service_tier"] = service_tier
|
||||
if parallel_tool_calls is False:
|
||||
# Default upstream behavior is parallel-allowed; only
|
||||
# forward when the user explicitly disabled it.
|
||||
body["disable_parallel_tool_use"] = True
|
||||
# Anthropic only caches a prefix when at least one cache_control
|
||||
# marker is attached to it — the frontend defaults
|
||||
# enable_prompt_caching to True for Anthropic, so treat `None` the
|
||||
|
|
@ -2558,6 +2633,9 @@ class ExternalProviderClient:
|
|||
enable_prompt_caching: Optional[bool] = None,
|
||||
openai_code_exec_container_id: Optional[str] = None,
|
||||
compaction_threshold: Optional[int] = None,
|
||||
*,
|
||||
service_tier: Optional[str] = None,
|
||||
parallel_tool_calls: Optional[bool] = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Call OpenAI's /v1/responses endpoint and translate its SSE stream back
|
||||
|
|
@ -2665,6 +2743,15 @@ class ExternalProviderClient:
|
|||
"input": input_items,
|
||||
"stream": True,
|
||||
}
|
||||
# Responses accepts service_tier on the same enum set as Chat
|
||||
# Completions minus `scale`. parallel_tool_calls follows the
|
||||
# same shape (default true). The frontend capability gate
|
||||
# (provider-capabilities.ts) already filters the option lists
|
||||
# per provider, so we just forward what we got.
|
||||
if service_tier in ("auto", "default", "flex", "priority"):
|
||||
body["service_tier"] = service_tier
|
||||
if parallel_tool_calls is not None:
|
||||
body["parallel_tool_calls"] = bool(parallel_tool_calls)
|
||||
# `summary: "auto"` is what makes /v1/responses emit reasoning
|
||||
# summary events — without it OpenAI returns no thinking text on
|
||||
# most reasoning models, the SSE handler has no <think>…</think>
|
||||
|
|
|
|||
|
|
@ -786,6 +786,48 @@ class ChatCompletionRequest(BaseModel):
|
|||
"to auto-create."
|
||||
),
|
||||
)
|
||||
frequency_penalty: Optional[float] = Field(
|
||||
None,
|
||||
ge = -2.0,
|
||||
le = 2.0,
|
||||
description = (
|
||||
"OpenAI Chat Completions frequency penalty (-2.0 to 2.0). "
|
||||
"Forwarded only on providers that accept it; Anthropic "
|
||||
"Messages and the OpenAI Responses family silently drop it."
|
||||
),
|
||||
)
|
||||
seed: Optional[int] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Best-effort determinism seed. Forwarded to OpenAI Chat "
|
||||
"Completions and OpenAI-compatible local backends. The "
|
||||
"Responses family rejects it server-side and Anthropic does "
|
||||
"not implement it, so it is silently dropped on those routes."
|
||||
),
|
||||
)
|
||||
service_tier: Optional[
|
||||
Literal["auto", "default", "flex", "priority", "scale", "standard_only"]
|
||||
] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Provider service tier. Anthropic accepts only `auto` and "
|
||||
"`standard_only`; OpenAI Chat accepts "
|
||||
"`auto|default|flex|priority|scale`; OpenAI Responses accepts "
|
||||
"`auto|default|flex|priority`. Unsupported values per provider "
|
||||
"are dropped in the per-provider stream helper instead of "
|
||||
"422'ing here so a stale frontend never breaks a request."
|
||||
),
|
||||
)
|
||||
parallel_tool_calls: Optional[bool] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Whether the provider may dispatch tool calls in parallel. "
|
||||
"OpenAI: forwarded as `parallel_tool_calls`. Anthropic: "
|
||||
"inverted into `disable_parallel_tool_use` on the Messages "
|
||||
"body. Default `None` preserves each provider's upstream "
|
||||
"default (which is `true` everywhere today)."
|
||||
),
|
||||
)
|
||||
|
||||
@model_validator(mode = "after")
|
||||
def _resolve_missing_tool_call_ids(self) -> "ChatCompletionRequest":
|
||||
|
|
|
|||
|
|
@ -1868,6 +1868,11 @@ async def _proxy_to_external_provider(
|
|||
anthropic_code_exec_container_id = payload.anthropic_code_exec_container_id,
|
||||
prompt_cache_ttl = payload.prompt_cache_ttl,
|
||||
compaction_threshold = payload.compaction_threshold,
|
||||
frequency_penalty = payload.frequency_penalty,
|
||||
seed = payload.seed,
|
||||
stop = payload.stop,
|
||||
service_tier = payload.service_tier,
|
||||
parallel_tool_calls = payload.parallel_tool_calls,
|
||||
stream = payload.stream,
|
||||
)
|
||||
try:
|
||||
|
|
|
|||
385
studio/backend/tests/test_sampling_params_routing.py
Normal file
385
studio/backend/tests/test_sampling_params_routing.py
Normal file
|
|
@ -0,0 +1,385 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""End-to-end routing tests for the new sampling parameters.
|
||||
|
||||
Pins the per-provider gating contract added by the
|
||||
expose-sampling-params PR: each of `frequency_penalty`, `seed`, `stop`
|
||||
/ `stop_sequences`, `service_tier`, `parallel_tool_calls` only appears
|
||||
on the outbound body when the upstream provider actually accepts it.
|
||||
|
||||
The provider matrix is captured per docs:
|
||||
- Anthropic Messages: accepts stop_sequences, service_tier
|
||||
(auto|standard_only), disable_parallel_tool_use (inverted). REJECTS
|
||||
frequency_penalty, seed, logprobs (silently dropped client-side).
|
||||
- OpenAI Chat Completions (default OAI-compat branch): accepts every
|
||||
field; OpenAI cloud uses `max_completion_tokens` rather than
|
||||
`max_tokens`.
|
||||
- OpenAI Responses (gpt-5.x / o3): rejects temperature, top_p,
|
||||
frequency_penalty, seed, stop, logprobs. Accepts service_tier
|
||||
(auto|default|flex|priority) and parallel_tool_calls.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from core.inference import external_provider as ep_mod
|
||||
from core.inference.external_provider import ExternalProviderClient
|
||||
|
||||
|
||||
def _drive(coro):
|
||||
return asyncio.new_event_loop().run_until_complete(coro)
|
||||
|
||||
|
||||
def _install_mock(monkeypatch, *, sse_payload: bytes | None = None) -> dict:
|
||||
captured: dict = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
try:
|
||||
captured["body"] = json.loads(request.content.decode("utf-8"))
|
||||
except json.JSONDecodeError:
|
||||
captured["body"] = None
|
||||
captured["url"] = str(request.url)
|
||||
captured["headers"] = dict(request.headers)
|
||||
return httpx.Response(
|
||||
200,
|
||||
content = sse_payload
|
||||
or (b"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"),
|
||||
headers = {"content-type": "text/event-stream"},
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
ep_mod,
|
||||
"_http_client",
|
||||
httpx.AsyncClient(transport = httpx.MockTransport(handler)),
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
# ── Anthropic ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _drive_anthropic(captured, **kwargs) -> dict:
|
||||
async def run():
|
||||
client = ExternalProviderClient(
|
||||
provider_type = "anthropic",
|
||||
base_url = "https://api.anthropic.com/v1",
|
||||
api_key = "sk-ant-test",
|
||||
)
|
||||
async for _ in client.stream_chat_completion(
|
||||
messages = [{"role": "user", "content": "hi"}],
|
||||
model = "claude-opus-4-7",
|
||||
temperature = 0.7,
|
||||
top_p = 0.95,
|
||||
max_tokens = 64,
|
||||
**kwargs,
|
||||
):
|
||||
pass
|
||||
await client.close()
|
||||
|
||||
_drive(run())
|
||||
return captured["body"]
|
||||
|
||||
|
||||
def test_anthropic_stop_sequences_forwarded_as_renamed_field(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, stop = ["END", "DONE"])
|
||||
assert body.get("stop_sequences") == ["END", "DONE"], body
|
||||
# Anthropic does not have a `stop` field; the unrenamed key must not appear.
|
||||
assert "stop" not in body, body
|
||||
|
||||
|
||||
def test_anthropic_single_string_stop_is_wrapped(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, stop = "STOPHERE")
|
||||
assert body.get("stop_sequences") == ["STOPHERE"], body
|
||||
|
||||
|
||||
def test_anthropic_empty_stop_omitted(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, stop = [])
|
||||
assert "stop_sequences" not in body, body
|
||||
assert "stop" not in body, body
|
||||
|
||||
|
||||
def test_anthropic_service_tier_forwarded_when_valid(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, service_tier = "standard_only")
|
||||
assert body.get("service_tier") == "standard_only", body
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bogus", ["flex", "priority", "scale", "default", "", "auto-foo"])
|
||||
def test_anthropic_service_tier_unsupported_values_dropped(monkeypatch, bogus):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, service_tier = bogus)
|
||||
assert "service_tier" not in body, body
|
||||
|
||||
|
||||
def test_anthropic_disable_parallel_tool_use_only_when_false(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, parallel_tool_calls = False)
|
||||
assert body.get("disable_parallel_tool_use") is True, body
|
||||
# Anthropic has no `parallel_tool_calls` field.
|
||||
assert "parallel_tool_calls" not in body, body
|
||||
|
||||
|
||||
def test_anthropic_parallel_tool_calls_default_not_sent(monkeypatch):
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(captured, parallel_tool_calls = True)
|
||||
# True is the upstream default; do not surface
|
||||
# `disable_parallel_tool_use: false` which would over-specify the request.
|
||||
assert "disable_parallel_tool_use" not in body, body
|
||||
assert "parallel_tool_calls" not in body, body
|
||||
|
||||
|
||||
def test_anthropic_rejects_openai_only_knobs(monkeypatch):
|
||||
"""frequency_penalty / seed are dropped at the dispatch layer.
|
||||
|
||||
Anthropic has no equivalent; the keyword args are not even forwarded
|
||||
from stream_chat_completion to _stream_anthropic. This test pins
|
||||
that no such field reaches the Messages body.
|
||||
"""
|
||||
captured = _install_mock(monkeypatch)
|
||||
body = _drive_anthropic(
|
||||
captured,
|
||||
frequency_penalty = 1.5,
|
||||
seed = 42,
|
||||
)
|
||||
assert "frequency_penalty" not in body, body
|
||||
assert "seed" not in body, body
|
||||
|
||||
|
||||
# ── OpenAI Chat Completions (default OAI-compat) ─────────────────────────
|
||||
|
||||
|
||||
def _drive_openai_compat(captured, **kwargs) -> dict:
|
||||
"""Send through the default OAI-compat branch (NOT /v1/responses).
|
||||
|
||||
Use a non-OpenAI provider_type so the dispatcher takes the default
|
||||
branch at the bottom of stream_chat_completion rather than the
|
||||
Responses translator path that routes provider_type=="openai".
|
||||
"""
|
||||
|
||||
async def run():
|
||||
client = ExternalProviderClient(
|
||||
provider_type = "mistral",
|
||||
base_url = "https://api.mistral.ai/v1",
|
||||
api_key = "test-key",
|
||||
)
|
||||
# mistral's OpenAI-compat /v1/chat/completions returns OpenAI
|
||||
# SSE; a single DONE frame is enough to drain the stream.
|
||||
async for _ in client.stream_chat_completion(
|
||||
messages = [{"role": "user", "content": "hi"}],
|
||||
model = "mistral-small-latest",
|
||||
temperature = 0.5,
|
||||
top_p = 0.9,
|
||||
max_tokens = 64,
|
||||
**kwargs,
|
||||
):
|
||||
pass
|
||||
await client.close()
|
||||
|
||||
_drive(run())
|
||||
return captured["body"]
|
||||
|
||||
|
||||
def _oai_done_payload() -> bytes:
|
||||
return b"data: [DONE]\n\n"
|
||||
|
||||
|
||||
def test_openai_compat_forwards_frequency_penalty(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, frequency_penalty = 1.25)
|
||||
assert body.get("frequency_penalty") == 1.25, body
|
||||
|
||||
|
||||
def test_openai_compat_forwards_seed(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, seed = 12345)
|
||||
assert body.get("seed") == 12345, body
|
||||
|
||||
|
||||
def test_openai_compat_forwards_stop_array(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, stop = ["END", "DONE"])
|
||||
assert body.get("stop") == ["END", "DONE"], body
|
||||
# The default OAI-compat branch does not rename to stop_sequences.
|
||||
assert "stop_sequences" not in body, body
|
||||
|
||||
|
||||
def test_openai_compat_truncates_stop_to_four(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(
|
||||
captured, stop = ["a", "b", "c", "d", "e", "f"]
|
||||
)
|
||||
assert body.get("stop") == ["a", "b", "c", "d"], body
|
||||
|
||||
|
||||
def test_openai_compat_empty_stop_omitted(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, stop = [])
|
||||
assert "stop" not in body, body
|
||||
|
||||
|
||||
def test_openai_compat_forwards_service_tier(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, service_tier = "flex")
|
||||
assert body.get("service_tier") == "flex", body
|
||||
|
||||
|
||||
def test_openai_compat_forwards_parallel_tool_calls(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured, parallel_tool_calls = False)
|
||||
assert body.get("parallel_tool_calls") is False, body
|
||||
|
||||
|
||||
def test_openai_compat_omits_unset_optionals(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _oai_done_payload())
|
||||
body = _drive_openai_compat(captured)
|
||||
# Optional knobs default to None / unset -> never appear.
|
||||
assert "frequency_penalty" not in body, body
|
||||
assert "seed" not in body, body
|
||||
assert "stop" not in body, body
|
||||
assert "service_tier" not in body, body
|
||||
assert "parallel_tool_calls" not in body, body
|
||||
|
||||
|
||||
# ── OpenAI Responses (gpt-5.x via /v1/responses) ─────────────────────────
|
||||
|
||||
|
||||
def _responses_done_payload() -> bytes:
|
||||
return (
|
||||
b"event: response.completed\n"
|
||||
b"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{}}}\n\n"
|
||||
)
|
||||
|
||||
|
||||
def _drive_openai_responses(captured, **kwargs) -> dict:
|
||||
async def run():
|
||||
client = ExternalProviderClient(
|
||||
provider_type = "openai",
|
||||
base_url = "https://api.openai.com/v1",
|
||||
api_key = "sk-test",
|
||||
)
|
||||
async for _ in client.stream_chat_completion(
|
||||
messages = [{"role": "user", "content": "hi"}],
|
||||
model = "gpt-5.5",
|
||||
temperature = 1.0,
|
||||
top_p = 1.0,
|
||||
max_tokens = 64,
|
||||
**kwargs,
|
||||
):
|
||||
pass
|
||||
await client.close()
|
||||
|
||||
_drive(run())
|
||||
return captured["body"]
|
||||
|
||||
|
||||
def test_openai_responses_drops_temperature_top_p(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(captured)
|
||||
assert "temperature" not in body, body
|
||||
assert "top_p" not in body, body
|
||||
|
||||
|
||||
def test_openai_responses_drops_frequency_penalty_seed_stop(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(
|
||||
captured,
|
||||
frequency_penalty = 1.5,
|
||||
seed = 99,
|
||||
stop = ["END"],
|
||||
)
|
||||
# Responses 400s on any of these; the dispatch must drop them
|
||||
# before they hit the wire.
|
||||
assert "frequency_penalty" not in body, body
|
||||
assert "seed" not in body, body
|
||||
assert "stop" not in body, body
|
||||
assert "stop_sequences" not in body, body
|
||||
|
||||
|
||||
def test_openai_responses_forwards_service_tier(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(captured, service_tier = "priority")
|
||||
assert body.get("service_tier") == "priority", body
|
||||
|
||||
|
||||
def test_openai_responses_rejects_chat_only_service_tier(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(captured, service_tier = "scale")
|
||||
# Responses only accepts auto|default|flex|priority -- `scale` is
|
||||
# silently dropped so a stale frontend cannot 400 the request.
|
||||
assert "service_tier" not in body, body
|
||||
|
||||
|
||||
def test_openai_responses_forwards_parallel_tool_calls(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(captured, parallel_tool_calls = False)
|
||||
assert body.get("parallel_tool_calls") is False, body
|
||||
|
||||
|
||||
def test_openai_responses_omits_unset_optionals(monkeypatch):
|
||||
captured = _install_mock(monkeypatch, sse_payload = _responses_done_payload())
|
||||
body = _drive_openai_responses(captured)
|
||||
assert "service_tier" not in body, body
|
||||
assert "parallel_tool_calls" not in body, body
|
||||
|
||||
|
||||
# ── Schema-level smoke tests ─────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_chat_completion_request_accepts_new_sampling_fields():
|
||||
from models.inference import ChatCompletionRequest
|
||||
|
||||
payload = ChatCompletionRequest.model_validate(
|
||||
{
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"frequency_penalty": -1.0,
|
||||
"seed": 0,
|
||||
"stop": ["END"],
|
||||
"service_tier": "auto",
|
||||
"parallel_tool_calls": True,
|
||||
}
|
||||
)
|
||||
assert payload.frequency_penalty == -1.0
|
||||
assert payload.seed == 0
|
||||
assert payload.stop == ["END"]
|
||||
assert payload.service_tier == "auto"
|
||||
assert payload.parallel_tool_calls is True
|
||||
|
||||
|
||||
def test_chat_completion_request_rejects_bad_service_tier():
|
||||
import pydantic
|
||||
from models.inference import ChatCompletionRequest
|
||||
|
||||
with pytest.raises(pydantic.ValidationError):
|
||||
ChatCompletionRequest.model_validate(
|
||||
{
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"service_tier": "bogus",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_chat_completion_request_clamps_frequency_penalty_range():
|
||||
import pydantic
|
||||
from models.inference import ChatCompletionRequest
|
||||
|
||||
with pytest.raises(pydantic.ValidationError):
|
||||
ChatCompletionRequest.model_validate(
|
||||
{
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"frequency_penalty": 3.0,
|
||||
}
|
||||
)
|
||||
with pytest.raises(pydantic.ValidationError):
|
||||
ChatCompletionRequest.model_validate(
|
||||
{
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"frequency_penalty": -3.0,
|
||||
}
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue