diff --git a/src/benchflow/agents/env.py b/src/benchflow/agents/env.py index ebf83e15c..0e1c5a6b3 100644 --- a/src/benchflow/agents/env.py +++ b/src/benchflow/agents/env.py @@ -469,12 +469,17 @@ def resolve_provider_env( ) -> None: """Detect provider for model, inject BENCHFLOW_PROVIDER_* and env_mapping.""" from benchflow.agents.providers import ( + PROVIDER_ENV_SOURCE_ENV, find_provider, find_provider_for_bare_model, resolve_base_url, strip_provider_prefix, ) + provider_route_missing = not any( + key in agent_env + for key in ("BENCHFLOW_PROVIDER_BASE_URL", "BENCHFLOW_PROVIDER_API_KEY") + ) agent_env.setdefault("BENCHFLOW_PROVIDER_MODEL", strip_provider_prefix(model)) agent_cfg = AGENTS.get(agent) # Agent-declared protocol takes precedence over provider's primary so @@ -532,6 +537,11 @@ def resolve_provider_env( "BENCHFLOW_PROVIDER_API_KEY", _BEDROCK_PROVIDER_PLACEHOLDER_API_KEY, ) + if provider_route_missing and all( + agent_env.get(key) + for key in ("BENCHFLOW_PROVIDER_BASE_URL", "BENCHFLOW_PROVIDER_API_KEY") + ): + agent_env[PROVIDER_ENV_SOURCE_ENV] = "registry" else: # No registered provider prefix — bridge the model's well-known API key # to BENCHFLOW_PROVIDER_API_KEY so env_mapping can translate it to diff --git a/src/benchflow/agents/providers.py b/src/benchflow/agents/providers.py index 9a576cbbc..4277b2c9e 100644 --- a/src/benchflow/agents/providers.py +++ b/src/benchflow/agents/providers.py @@ -61,6 +61,8 @@ from dataclasses import dataclass, field +PROVIDER_ENV_SOURCE_ENV = "_BENCHFLOW_PROVIDER_ENV_SOURCE" + @dataclass class ProviderConfig: @@ -275,6 +277,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}", diff --git a/src/benchflow/providers/litellm_config.py b/src/benchflow/providers/litellm_config.py index f703b32d2..db6c1696b 100644 --- a/src/benchflow/providers/litellm_config.py +++ b/src/benchflow/providers/litellm_config.py @@ -4,11 +4,13 @@ import hashlib import json +import math import re from dataclasses import dataclass from urllib.parse import urlparse from benchflow.agents.providers import ( + PROVIDER_ENV_SOURCE_ENV, ProviderConfig, find_provider, resolve_base_url, @@ -177,6 +179,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 = ( @@ -291,7 +325,8 @@ 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: + provider_env_is_registry = env.get(PROVIDER_ENV_SOURCE_ENV) == "registry" + if explicit_api_base and explicit_api_key and not provider_env_is_registry: api_base = explicit_api_base else: try: @@ -335,7 +370,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 provider_env_is_registry else _registered_api_key_ref(provider_cfg) ) if api_key_ref: @@ -344,7 +379,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), @@ -360,12 +394,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) @@ -404,6 +440,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), diff --git a/tests/test_litellm_config.py b/tests/test_litellm_config.py index b829c4021..57e5d3714 100644 --- a/tests/test_litellm_config.py +++ b/tests/test_litellm_config.py @@ -2,6 +2,7 @@ import pytest +from benchflow.agents.env import resolve_provider_env from benchflow.providers.litellm_config import ( litellm_proxy_config, resolve_litellm_route, @@ -120,6 +121,125 @@ def test_registered_provider_route_honors_explicit_generic_proxy_env(): assert route.required_env == ("BENCHFLOW_PROVIDER_API_KEY",) +@pytest.mark.parametrize( + ("model", "key", "agent", "expected_base", "expected_upstream"), + [ + ( + "zai-coding/glm-5.9", + "ZAI_API_KEY", + "claude-agent-acp", + "https://api.z.ai/api/coding/paas/v4", + "openai/glm-5.9", + ), + ( + "openrouter/qwen/qwen3.5-397b-a17b", + "OPENROUTER_API_KEY", + "pi-acp", + "https://openrouter.ai/api/v1", + "openai/qwen/qwen3.5-397b-a17b", + ), + ], +) +def test_registered_provider_route_ignores_registry_derived_generic_proxy_env( + model, key, agent, expected_base, expected_upstream +): + env = {key: "native-key"} + resolve_provider_env(env, model, agent) + route = resolve_litellm_route(model, env) + + 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 + + +@pytest.mark.parametrize("resolve_first", [False, True], ids=["direct", "resolved"]) +def test_registered_provider_route_preserves_explicit_endpoint_with_native_key( + resolve_first, +): + env = { + "ZAI_API_KEY": "same-key", + "BENCHFLOW_PROVIDER_BASE_URL": "https://api.z.ai/api/anthropic", + "BENCHFLOW_PROVIDER_API_KEY": "same-key", + } + if resolve_first: + resolve_provider_env(env, "zai-coding/glm-5.2", "claude-agent-acp") + route = resolve_litellm_route("zai-coding/glm-5.2", env) + + assert route.litellm_params["api_base"] == "https://api.z.ai/api/anthropic" + assert route.required_env == ("BENCHFLOW_PROVIDER_API_KEY",) + + +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.""" diff --git a/tests/test_providers.py b/tests/test_providers.py index 3d6c486e7..4c7ad1dcd 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -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"] diff --git a/tests/test_registry_invariants.py b/tests/test_registry_invariants.py index a6dc7d03f..5ff505b73 100644 --- a/tests/test_registry_invariants.py +++ b/tests/test_registry_invariants.py @@ -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"),