mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-27 06:56:39 +03:00
refactor(llm_clients): build the client keyword arguments in llm_clients
- build_llm_kwargs(config) replaces the graph's _get_provider_kwargs, with the retry and token coercion beside it
This commit is contained in:
@@ -11,7 +11,7 @@ import importlib
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config_module
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph, _coerce_max_retries
|
||||
from tradingagents.llm_clients.factory import _coerce_max_retries, build_llm_kwargs
|
||||
|
||||
# --- coercion / validation -------------------------------------------------
|
||||
|
||||
@@ -44,36 +44,31 @@ def test_coerce_rejects_non_integers(bad):
|
||||
|
||||
# --- forwarding into provider kwargs --------------------------------------
|
||||
|
||||
def _bare_graph(config):
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
g.config = config
|
||||
return g
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_not_forwarded_when_unset():
|
||||
kwargs = _bare_graph({"llm_provider": "openai", "llm_max_retries": None})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": "openai", "llm_max_retries": None})
|
||||
assert "max_retries" not in kwargs
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider", ["openai", "anthropic", "google"])
|
||||
def test_forwarded_across_providers(provider):
|
||||
kwargs = _bare_graph({"llm_provider": provider, "llm_max_retries": 6})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": provider, "llm_max_retries": 6})
|
||||
assert kwargs["max_retries"] == 6
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_forwarded_env_string_is_coerced():
|
||||
# env vars arrive as strings; the consumer coerces (like temperature)
|
||||
kwargs = _bare_graph({"llm_provider": "openai", "llm_max_retries": "4"})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": "openai", "llm_max_retries": "4"})
|
||||
assert kwargs["max_retries"] == 4
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_invalid_config_value_fails_loudly():
|
||||
with pytest.raises(ValueError):
|
||||
_bare_graph({"llm_provider": "openai", "llm_max_retries": -1})._get_provider_kwargs()
|
||||
build_llm_kwargs({"llm_provider": "openai", "llm_max_retries": -1})
|
||||
|
||||
|
||||
# --- env overlay -----------------------------------------------------------
|
||||
|
||||
@@ -13,7 +13,7 @@ import importlib
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config_module
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph, _coerce_max_tokens
|
||||
from tradingagents.llm_clients.factory import _coerce_max_tokens, build_llm_kwargs
|
||||
|
||||
# --- coercion / validation -------------------------------------------------
|
||||
|
||||
@@ -46,15 +46,10 @@ def test_coerce_rejects_non_integers(bad):
|
||||
|
||||
# --- forwarding into provider kwargs (right key per provider) --------------
|
||||
|
||||
def _bare_graph(config):
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
g.config = config
|
||||
return g
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_not_forwarded_when_unset():
|
||||
kwargs = _bare_graph({"llm_provider": "openai", "max_tokens": None})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": "openai", "max_tokens": None})
|
||||
assert "max_tokens" not in kwargs
|
||||
assert "max_output_tokens" not in kwargs
|
||||
|
||||
@@ -62,7 +57,7 @@ def test_not_forwarded_when_unset():
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider", ["openai", "anthropic", "deepseek", "openai_compatible"])
|
||||
def test_forwarded_as_max_tokens_for_non_google(provider):
|
||||
kwargs = _bare_graph({"llm_provider": provider, "max_tokens": 8192})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": provider, "max_tokens": 8192})
|
||||
assert kwargs["max_tokens"] == 8192
|
||||
assert "max_output_tokens" not in kwargs
|
||||
|
||||
@@ -70,21 +65,21 @@ def test_forwarded_as_max_tokens_for_non_google(provider):
|
||||
@pytest.mark.unit
|
||||
def test_forwarded_as_max_output_tokens_for_google():
|
||||
# Gemini's kwarg name differs; forwarding plain max_tokens would be rejected.
|
||||
kwargs = _bare_graph({"llm_provider": "google", "max_tokens": 8192})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": "google", "max_tokens": 8192})
|
||||
assert kwargs["max_output_tokens"] == 8192
|
||||
assert "max_tokens" not in kwargs
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_env_string_is_coerced():
|
||||
kwargs = _bare_graph({"llm_provider": "openai", "max_tokens": "4096"})._get_provider_kwargs()
|
||||
kwargs = build_llm_kwargs({"llm_provider": "openai", "max_tokens": "4096"})
|
||||
assert kwargs["max_tokens"] == 4096
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_invalid_value_fails_loudly():
|
||||
with pytest.raises(ValueError):
|
||||
_bare_graph({"llm_provider": "openai", "max_tokens": 0})._get_provider_kwargs()
|
||||
build_llm_kwargs({"llm_provider": "openai", "max_tokens": 0})
|
||||
|
||||
|
||||
# --- client-side allowlists carry the kwarg --------------------------------
|
||||
|
||||
@@ -61,14 +61,11 @@ class TestTemperatureEnvOverlay:
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestProviderKwargsTemperature:
|
||||
"""_get_provider_kwargs float-coerces and forwards temperature, or omits it."""
|
||||
"""build_llm_kwargs float-coerces and forwards temperature, or omits it."""
|
||||
|
||||
def _kwargs_for(self, temperature):
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
# Call the method without constructing the full graph.
|
||||
graph = TradingAgentsGraph.__new__(TradingAgentsGraph)
|
||||
graph.config = {"llm_provider": "openai", "temperature": temperature}
|
||||
return TradingAgentsGraph._get_provider_kwargs(graph)
|
||||
from tradingagents.llm_clients import build_llm_kwargs
|
||||
return build_llm_kwargs({"llm_provider": "openai", "temperature": temperature})
|
||||
|
||||
def test_float_string_coerced(self):
|
||||
assert self._kwargs_for("0.3")["temperature"] == 0.3
|
||||
|
||||
Reference in New Issue
Block a user