Skip to content
Open
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions src/benchflow/agents/providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -275,6 +275,30 @@ def all_endpoints(self) -> dict[str, str]:
},
],
),
"zai-coding": ProviderConfig(
name="zai-coding",
base_url="https://api.z.ai/api/coding/paas/v4",
api_protocol="openai-completions",
auth_type="api_key",
auth_env="ZAI_API_KEY",
endpoints={
"openai-completions": "https://api.z.ai/api/coding/paas/v4",
"openai-responses": "https://api.z.ai/api/coding/paas/v4",
"anthropic-messages": "https://api.z.ai/api/anthropic",
},
models=[
{
"id": model,
"name": model.upper(),
"reasoning": True,
"input": ["text"],
"cost": {"input": 0, "output": 0, "cacheRead": 0, "cacheWrite": 0},
"contextWindow": 200000,
"maxTokens": 131072,
}
for model in ("glm-5", "glm-5.1", "glm-5.2", "glm-5-turbo")
],
),
"kimi": ProviderConfig(
name="kimi",
base_url="{base_url}",
Expand Down
60 changes: 56 additions & 4 deletions src/benchflow/providers/litellm_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import hashlib
import json
import math
import re
from dataclasses import dataclass
from urllib.parse import urlparse
Expand Down Expand Up @@ -177,6 +178,38 @@ def _registered_api_key_ref(cfg: ProviderConfig) -> str | None:
return None


def _apply_model_generation_params(
params: dict[str, str | int | float | bool | list[str]], env: dict[str, str]
) -> None:
for env_name, param, upper_bound in (
("BENCHFLOW_MODEL_TEMPERATURE", "temperature", None),
("BENCHFLOW_MODEL_TOP_P", "top_p", 1.0),
):
value = (env.get(env_name) or "").strip()
try:
parsed = float(value)
except ValueError:
continue
if (
math.isfinite(parsed)
and parsed >= 0
and (upper_bound is None or parsed <= upper_bound)
):
params[param] = parsed

try:
max_tokens = int((env.get("BENCHFLOW_MODEL_MAX_TOKENS") or "").strip())
except ValueError:
return
if max_tokens > 0:
token_param = (
"max_output_tokens"
if env.get("BENCHFLOW_PROVIDER_PROTOCOL") == "openai-responses"
else "max_tokens"
)
params[token_param] = max_tokens


def _provider_reasoning_effort(env: dict[str, str]) -> str | None:
"""Return an explicitly requested gateway-side reasoning effort."""
raw = (
Expand Down Expand Up @@ -291,7 +324,24 @@ def _route_registered_provider(
)
explicit_api_base = (env.get("BENCHFLOW_PROVIDER_BASE_URL") or "").strip()
explicit_api_key = (env.get("BENCHFLOW_PROVIDER_API_KEY") or "").strip()
if explicit_api_base and explicit_api_key:
registry_api_bases = set()
for endpoint_protocol in provider_cfg.all_endpoints:
try:
endpoint = resolve_base_url(
provider_cfg, env, protocol=endpoint_protocol
).rstrip("/")
except KeyError:
continue
registry_api_bases.add(endpoint)
native_api_key = (
(env.get(provider_cfg.auth_env) or "").strip() if provider_cfg.auth_env else ""
)
registry_api_base = (
explicit_api_base.rstrip("/") in registry_api_bases
and bool(native_api_key)
and explicit_api_key == native_api_key
)
Comment thread
kywch marked this conversation as resolved.
Outdated
if explicit_api_base and explicit_api_key and not registry_api_base:
api_base = explicit_api_base
else:
try:
Expand Down Expand Up @@ -335,7 +385,7 @@ def _route_registered_provider(
params["api_base"] = api_base
api_key_ref = (
_env_ref("BENCHFLOW_PROVIDER_API_KEY")
if explicit_api_base and explicit_api_key
if explicit_api_base and explicit_api_key and not registry_api_base
else _registered_api_key_ref(provider_cfg)
)
if api_key_ref:
Expand All @@ -344,7 +394,6 @@ def _route_registered_provider(
required_env.append("BENCHFLOW_PROVIDER_API_KEY")
elif provider_cfg.auth_env:
required_env.append(provider_cfg.auth_env)

return LiteLLMRoute(
requested_model=model,
model_alias=safe_model_alias(model),
Expand All @@ -360,12 +409,14 @@ def resolve_litellm_route(model: str, env: dict[str, str]) -> LiteLLMRoute:
provider = find_provider(model)
if provider is not None:
provider_name, provider_cfg = provider
return _route_registered_provider(
route = _route_registered_provider(
model=model,
provider_name=provider_name,
provider_cfg=provider_cfg,
env=env,
)
_apply_model_generation_params(route.litellm_params, env)
return route

lower = model.lower()
bare = strip_provider_prefix(model)
Expand Down Expand Up @@ -404,6 +455,7 @@ def resolve_litellm_route(model: str, env: dict[str, str]) -> LiteLLMRoute:
key = required[0] if required else None
if key and "api_key" not in params:
params["api_key"] = _env_ref(key)
_apply_model_generation_params(params, env)
return LiteLLMRoute(
requested_model=model,
model_alias=safe_model_alias(model),
Expand Down
107 changes: 107 additions & 0 deletions tests/test_litellm_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,113 @@ def test_registered_provider_route_honors_explicit_generic_proxy_env():
assert route.required_env == ("BENCHFLOW_PROVIDER_API_KEY",)


@pytest.mark.parametrize(
("model", "key", "derived_base", "expected_base", "expected_upstream"),
[
(
"zai-coding/glm-5.9",
"ZAI_API_KEY",
"https://api.z.ai/api/anthropic/",
"https://api.z.ai/api/coding/paas/v4",
"openai/glm-5.9",
),
(
"openrouter/qwen/qwen3.5-397b-a17b",
"OPENROUTER_API_KEY",
"https://openrouter.ai/api/v1",
"https://openrouter.ai/api/v1",
"openai/qwen/qwen3.5-397b-a17b",
),
],
)
def test_registered_provider_route_ignores_registry_derived_generic_proxy_env(
model, key, derived_base, expected_base, expected_upstream
):
route = resolve_litellm_route(
model,
{
key: "native-key",
"BENCHFLOW_PROVIDER_BASE_URL": derived_base,
"BENCHFLOW_PROVIDER_API_KEY": "native-key",
},
)

assert route.litellm_params["api_base"] == expected_base
assert route.litellm_params["api_key"] == f"os.environ/{key}"
assert route.required_env == (key,)
assert route.upstream_model == expected_upstream


def test_registered_endpoint_with_generic_key_remains_explicit_override():
route = resolve_litellm_route(
"openrouter/qwen/qwen3.5-397b-a17b",
{
"OPENROUTER_API_KEY": "native-key",
"BENCHFLOW_PROVIDER_BASE_URL": "https://openrouter.ai/api/v1",
"BENCHFLOW_PROVIDER_API_KEY": "generic-key",
},
)

assert route.litellm_params["api_base"] == "https://openrouter.ai/api/v1"
assert route.litellm_params["api_key"] == "os.environ/BENCHFLOW_PROVIDER_API_KEY"
assert route.required_env == ("BENCHFLOW_PROVIDER_API_KEY",)


@pytest.mark.parametrize(
("model", "key", "protocol", "token_param"),
[
("zai-coding/glm-5.2", "ZAI_API_KEY", "openai-responses", "max_output_tokens"),
("gemini-3.5-flash", "GEMINI_API_KEY", "openai-completions", "max_tokens"),
],
)
def test_litellm_route_generation_overrides(model, key, protocol, token_param):
params = resolve_litellm_route(
model,
{
key: "key",
"BENCHFLOW_PROVIDER_PROTOCOL": protocol,
"BENCHFLOW_MODEL_TEMPERATURE": "1.0",
"BENCHFLOW_MODEL_TOP_P": "0.95",
"BENCHFLOW_MODEL_MAX_TOKENS": "131072",
},
).litellm_params
expected = {"temperature": 1.0, "top_p": 0.95, token_param: 131072}
assert {name: params[name] for name in expected} == expected
assert ({"max_tokens", "max_output_tokens"} - {token_param}).isdisjoint(params)


@pytest.mark.parametrize(
("env_name", "param", "value"),
[
("BENCHFLOW_MODEL_TEMPERATURE", "temperature", "nan"),
("BENCHFLOW_MODEL_TEMPERATURE", "temperature", "inf"),
("BENCHFLOW_MODEL_TEMPERATURE", "temperature", "-0.1"),
("BENCHFLOW_MODEL_TOP_P", "top_p", "1.1"),
("BENCHFLOW_MODEL_TOP_P", "top_p", "-0.1"),
("BENCHFLOW_MODEL_MAX_TOKENS", "max_tokens", "0"),
("BENCHFLOW_MODEL_MAX_TOKENS", "max_tokens", "-1"),
("BENCHFLOW_MODEL_MAX_TOKENS", "max_tokens", "1.5"),
],
)
def test_litellm_route_rejects_invalid_generation_overrides(env_name, param, value):
params = resolve_litellm_route(
"zai-coding/glm-5.2",
{
"ZAI_API_KEY": "key",
env_name: value,
},
).litellm_params
assert param not in params


def test_special_registered_provider_generation_overrides():
route = resolve_litellm_route(
"aws-bedrock/us.anthropic.claude-opus-4-8",
{"BENCHFLOW_MODEL_MAX_TOKENS": "4096"},
)
assert route.litellm_params["max_tokens"] == 4096


@pytest.mark.parametrize("model", ["gemini/gemini-2.5-flash", "gemini-2.5-flash"])
def test_gemini_native_route_honors_explicit_base_url(model):
"""Guards the fix from PR #881 for issue #672."""
Expand Down
11 changes: 11 additions & 0 deletions tests/test_providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,17 @@ def test_protocol_selects_endpoint(self):
== "https://api.z.ai/api/paas/v4"
)

@pytest.mark.parametrize(
("protocol", "expected"),
[
("openai-completions", "https://api.z.ai/api/coding/paas/v4"),
("openai-responses", "https://api.z.ai/api/coding/paas/v4"),
("anthropic-messages", "https://api.z.ai/api/anthropic"),
],
)
def test_zai_coding_protocol_selects_endpoint(self, protocol, expected):
assert resolve_base_url(PROVIDERS["zai-coding"], {}, protocol) == expected

def test_protocol_fallback_to_base_url(self):
"""Unknown protocol falls back to primary base_url."""
p = PROVIDERS["zai"]
Expand Down
1 change: 1 addition & 0 deletions tests/test_registry_invariants.py
Original file line number Diff line number Diff line change
Expand Up @@ -486,6 +486,7 @@ def test_provider_model_prefixes_unique_and_resolvable():
("aws-bedrock/openai.gpt-oss-20b-1:0", "aws-bedrock"),
("github-models/openai/gpt-4.1-mini", "github-models"),
("zai/glm-5", "zai"),
("zai-coding/glm-5.2", "zai-coding"),
("vllm/local-model", "vllm"),
("kimi/kimi-k2.6", "kimi"),
("qwen-dashscope/qwen3.6-max-preview", "qwen-dashscope"),
Expand Down
Loading