mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-25 05:52:35 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
90be56b9b2 | ||
|
|
25c751346e | ||
|
|
6caf0db866 | ||
|
|
f692949f14 | ||
|
|
f58a585390 | ||
|
|
e1a5f20480 | ||
|
|
bbd4d0413a | ||
|
|
b176c9374b | ||
|
|
4187716842 | ||
|
|
24c38602e5 | ||
|
|
c924f84114 | ||
|
|
f281feb483 | ||
|
|
8efb702871 | ||
|
|
91c4319147 | ||
|
|
1e9ad314d8 | ||
|
|
825e6321ae | ||
|
|
4a30cb1c0a | ||
|
|
56bd98f690 | ||
|
|
969861e8be | ||
|
|
c9695673bf | ||
|
|
e7354d8af4 | ||
|
|
dcbff5ea42 | ||
|
|
852ffead43 | ||
|
|
6097b582d9 | ||
|
|
c42a2f2c61 | ||
|
|
a58aa613fc | ||
|
|
f142e79660 | ||
|
|
9b14233a24 | ||
|
|
41fc25ac0d | ||
|
|
4e1faf6465 | ||
|
|
a6e92a6b5e | ||
|
|
15b8276b92 | ||
|
|
2938e1c7e9 | ||
|
|
f197e09dcc | ||
|
|
b4479b0c70 | ||
|
|
4a71dc708d | ||
|
|
1e439352c0 | ||
|
|
632fb51d18 | ||
|
|
5c2dded82c | ||
|
|
1f7a2aa865 | ||
|
|
25a86a8890 | ||
|
|
2340fe4396 | ||
|
|
96daaf1152 | ||
|
|
a9cc3be731 | ||
|
|
c78fa86500 | ||
|
|
2d17df8da1 | ||
|
|
7fe2252244 | ||
|
|
76a93d615a | ||
|
|
c039e39056 | ||
|
|
4acdc16513 | ||
|
|
d5ba41bac3 | ||
|
|
10cc070fa3 | ||
|
|
ed6eae44b6 | ||
|
|
aef4af90e3 | ||
|
|
bbcd6661af | ||
|
|
f8042efdde | ||
|
|
d0881f081e | ||
|
|
8bde10cf44 | ||
|
|
9683194793 | ||
|
|
04b691804c | ||
|
|
c3bb991974 | ||
|
|
f0a1cf6290 | ||
|
|
de7e43fc4a | ||
|
|
13b35e8aeb | ||
|
|
008ac655a9 | ||
|
|
2ddfe4ceb5 | ||
|
|
85d9137437 | ||
|
|
2ca59cc795 | ||
|
|
63989515c9 | ||
|
|
486dec1710 | ||
|
|
8d30fee06b | ||
|
|
3244a568ed | ||
|
|
b6dad747e9 | ||
|
|
f881c4a111 | ||
|
|
8ac4371387 | ||
|
|
60dcf64723 | ||
|
|
8d64416436 | ||
|
|
c4152d6946 | ||
|
|
375af054b7 | ||
|
|
d8eceb6571 | ||
|
|
63738c8f2c | ||
|
|
7bf2117491 | ||
|
|
8721b92f65 | ||
|
|
9794a90734 | ||
|
|
6436d1ff30 | ||
|
|
dffff22951 | ||
|
|
ef52d4452b | ||
|
|
2942655f70 | ||
|
|
fadc698e20 | ||
|
|
29e331a9af | ||
|
|
99313bfdbc | ||
|
|
d04693a208 | ||
|
|
4a9f196e92 | ||
|
|
2c1ba388d6 | ||
|
|
b9666f6e09 | ||
|
|
34899bd320 | ||
|
|
b20c8e60a4 | ||
|
|
241638da68 | ||
|
|
2e38b47dca | ||
|
|
673abb3c68 | ||
|
|
9b4c741d33 | ||
|
|
62d3479217 | ||
|
|
be952b8ecc | ||
|
|
821848bb82 | ||
|
|
d6ca23aee5 | ||
|
|
ef383df8f4 | ||
|
|
d58b838081 | ||
|
|
ffd5d9a180 | ||
|
|
16f7fd613c | ||
|
|
260c899c72 | ||
|
|
94113c8d11 | ||
|
|
7cc478ad07 | ||
|
|
1c44dd1ffc | ||
|
|
96111aa368 | ||
|
|
9dee508c44 | ||
|
|
5a26ae17a1 | ||
|
|
a4acd8a174 | ||
|
|
2322dd9baa | ||
|
|
70b58c21dc | ||
|
|
2448d0a125 | ||
|
|
c95f83dfaf | ||
|
|
ecbe3e3a21 | ||
|
|
e93c5c53c2 | ||
|
|
45c1744b86 | ||
|
|
63be7fe7f1 | ||
|
|
30d42abd5d | ||
|
|
a2f51da917 | ||
|
|
b43bc31479 | ||
|
|
8db41f6bca | ||
|
|
51a245dbe1 | ||
|
|
43fc275b36 | ||
|
|
0ef56e6a33 | ||
|
|
539eae8fd6 | ||
|
|
9b98f09613 | ||
|
|
8b7ece8a3e | ||
|
|
a33fd4c0f1 | ||
|
|
7bbe33ab1d | ||
|
|
030b434585 | ||
|
|
3f6c082695 | ||
|
|
d78c698d0e | ||
|
|
40774ca042 | ||
|
|
01477f9afb | ||
|
|
0f70af2f31 | ||
|
|
43bd32befa | ||
|
|
a102afa090 | ||
|
|
daf1da9c35 | ||
|
|
b47a828a4f | ||
|
|
622f99d28a | ||
|
|
3570f2e1e6 | ||
|
|
85946c2f60 | ||
|
|
cbd17ac3e0 | ||
|
|
8ab24f30af | ||
|
|
2b2d685df6 | ||
|
|
a0120e1805 | ||
|
|
0b61effd6c | ||
|
|
ec3974b84e | ||
|
|
0405168f20 | ||
|
|
709fe2b646 | ||
|
|
517eeaf4b9 | ||
|
|
9ad98c55c5 | ||
|
|
ee1ece3347 | ||
|
|
7bb16c5daa | ||
|
|
a420ad0f3b | ||
|
|
c15200dc28 | ||
|
|
7aef10acbd | ||
|
|
03600f3121 | ||
|
|
6b6177ebf7 | ||
|
|
e3bc872982 | ||
|
|
cbc5f67d42 | ||
|
|
3cddf1e331 | ||
|
|
308757c999 | ||
|
|
eeb84aa63b | ||
|
|
9fd54f8368 | ||
|
|
7df18fc912 | ||
|
|
db059034a2 | ||
|
|
ddfb840ecf | ||
|
|
895ed130f9 | ||
|
|
295e84cd54 | ||
|
|
20d3b0782f | ||
|
|
4e7821d574 | ||
|
|
0c1231a405 | ||
|
|
e4be7cc5a3 | ||
|
|
a597063747 | ||
|
|
dab07688fb | ||
|
|
65608831f8 | ||
|
|
76add9048f | ||
|
|
7c8fe2fe9f | ||
|
|
2a58c2208f | ||
|
|
04f434e86d | ||
|
|
2e67782f20 | ||
|
|
1ff3f07a73 | ||
|
|
2f85be624e | ||
|
|
c93b92c7a4 | ||
|
|
d6762d6095 | ||
|
|
8694bd070d | ||
|
|
2c9f1bfe65 | ||
|
|
8a22594607 | ||
|
|
47cbb321fe | ||
|
|
e80636fc0e | ||
|
|
a66aa8fb94 | ||
|
|
3543e5397e | ||
|
|
d7b40a2a5c | ||
|
|
61522e103e | ||
|
|
e848b5e812 | ||
|
|
3e5e99b368 | ||
|
|
a2e7ac1599 | ||
|
|
b16fe53efe | ||
|
|
a5cb7cbd61 | ||
|
|
78d063dc5c | ||
|
|
819e813a14 | ||
|
|
800862405d | ||
|
|
f10daa2824 | ||
|
|
249caba06f | ||
|
|
a2f343bb54 | ||
|
|
879e2bb5da | ||
|
|
9f7abfcbd5 | ||
|
|
d13e9b7946 | ||
|
|
6b384f74f9 | ||
|
|
384fe1a3d2 | ||
|
|
0fcf13624e | ||
|
|
d0dd0420ad | ||
|
|
faaeebac70 | ||
|
|
0011b5ebf5 | ||
|
|
4f057e290c | ||
|
|
9e00c8117f | ||
|
|
78fe77f4e6 | ||
|
|
e1316686f8 | ||
|
|
9482cae188 | ||
|
|
19d22b54a9 | ||
|
|
704b7627f2 | ||
|
|
22bb91bd83 | ||
|
|
afdc6d4ec1 | ||
|
|
e2c850eb17 | ||
|
|
c405867bde | ||
|
|
db7e0a67e2 | ||
|
|
5bae826749 | ||
|
|
7e9e7b83c7 | ||
|
|
2c97bad45c | ||
|
|
7c37249f80 | ||
|
|
4016fd4efa | ||
|
|
bba147798f | ||
|
|
0fda24515f | ||
|
|
4cbd4b086f | ||
|
|
ebd2e12e67 | ||
|
|
f85f5d9f5d | ||
|
|
8e7654f0df | ||
|
|
872b063e69 | ||
|
|
6abc768c1d | ||
|
|
99ec63f966 | ||
|
|
e7ec980021 | ||
|
|
8536ccacdd | ||
|
|
fa4d01c23a | ||
|
|
b0f6058299 | ||
|
|
59d6b2152d | ||
|
|
10c136f49c | ||
|
|
4f965bf46a | ||
|
|
bdb9c29d44 | ||
|
|
bdc5fc62d3 | ||
|
|
78fb66aed1 | ||
|
|
7269f877c1 | ||
|
|
28d5cc661f | ||
|
|
7004dfe554 | ||
|
|
f4519bcb84 | ||
|
|
4300b68f19 | ||
|
|
2d2c9e6d66 | ||
|
|
4641c03340 | ||
|
|
e75d17bc51 | ||
|
|
6cddd26d6e | ||
|
|
c61242a28c | ||
|
|
58e99421bd | ||
|
|
46e1b600b8 | ||
|
|
ae8c8aebe8 | ||
|
|
f3f58bdbdc | ||
|
|
e1113880a1 | ||
|
|
bd6a5b75b5 | ||
|
|
8793336dad | ||
|
|
047b38971c | ||
|
|
f5026009f9 | ||
|
|
589b351f2a | ||
|
|
6c9c9ce1fd | ||
|
|
b8b2825783 | ||
|
|
318adda0c6 | ||
|
|
c3ba3bf428 | ||
|
|
7cca9c924e | ||
|
|
bd9b1e5efa | ||
|
|
77755f0431 | ||
|
|
0b13145dc0 | ||
|
|
3ff28f3559 | ||
|
|
7d200d834a | ||
|
|
08bfe70a69 | ||
|
|
f362a160c3 | ||
|
|
64f07671b9 | ||
|
|
b19c5c18fb | ||
|
|
551fd7f074 | ||
|
|
b0f9d180f9 | ||
|
|
9cc283ac22 | ||
|
|
fe9c8d5d31 | ||
|
|
eec6ca4b53 | ||
|
|
3642f5917c | ||
|
|
907bc8022a | ||
|
|
8a60662070 | ||
|
|
f047f26df0 | ||
|
|
35856ff33e | ||
|
|
5fec171a1e | ||
|
|
50c82a25b5 | ||
|
|
8b3068d091 | ||
|
|
66a02b3193 | ||
|
|
e9470b69c4 | ||
|
|
b4b133eb2d | ||
|
|
80aab35119 | ||
|
|
393d4c6a1b | ||
|
|
aba1880c8c | ||
|
|
6cd35179fa | ||
|
|
102b026d23 | ||
|
|
224941d8c2 | ||
|
|
93b87d5119 | ||
|
|
54cdb146d0 | ||
|
|
b06936f420 | ||
|
|
b75940e901 | ||
|
|
3d040f8da4 | ||
|
|
50961b2477 | ||
|
|
a3761bdd66 | ||
|
|
d4dadb82fc | ||
|
|
79051580b8 | ||
|
|
13b826a31d | ||
|
|
b2ef960da7 | ||
|
|
a5dcc7da45 | ||
|
|
7bb2941b07 | ||
|
|
32be17c606 | ||
|
|
c07dcf026b | ||
|
|
d23fb539e9 | ||
|
|
b01051b9f4 | ||
|
|
8fdbbcca3d | ||
|
|
86bc0e793f | ||
|
|
7fc9c28a94 | ||
|
|
7bcc2cbd8a | ||
|
|
6211b1132a | ||
|
|
8b04ec307f | ||
|
|
0ab323c2c6 | ||
|
|
a6734d71bc | ||
|
|
a438acdbbd | ||
|
|
c73e374e7c | ||
|
|
f704828f89 | ||
|
|
fda4f664e8 | ||
|
|
718df34932 | ||
|
|
43aa9c5d09 | ||
|
|
26c5ba5a78 | ||
|
|
78ea029a0b | ||
|
|
ee3d499894 | ||
|
|
7abff0f354 | ||
|
|
b575bd0941 | ||
|
|
b8f712b170 | ||
|
|
52284ce13c | ||
|
|
11804f88ff | ||
|
|
1e86e74314 | ||
|
|
c2f897fc67 | ||
|
|
ed32081f57 | ||
|
|
2af7ef3d79 | ||
|
|
383deb72aa | ||
|
|
7eaf4d995f | ||
|
|
da84ef43aa | ||
|
|
90b23e72f5 | ||
|
|
417b09712c | ||
|
|
570644d939 | ||
|
|
9647359246 | ||
|
|
99789f9cd1 | ||
|
|
a879868396 | ||
|
|
0013415378 | ||
|
|
0fdfd35867 | ||
|
|
e994e56c23 |
@@ -0,0 +1,15 @@
|
||||
.git
|
||||
.venv
|
||||
.env
|
||||
.claude
|
||||
.idea
|
||||
.vscode
|
||||
.DS_Store
|
||||
__pycache__
|
||||
*.egg-info
|
||||
build
|
||||
dist
|
||||
results
|
||||
eval_results
|
||||
Dockerfile
|
||||
docker-compose.yml
|
||||
@@ -0,0 +1,5 @@
|
||||
# Azure OpenAI
|
||||
AZURE_OPENAI_API_KEY=
|
||||
AZURE_OPENAI_ENDPOINT=https://your-resource-name.openai.azure.com/
|
||||
AZURE_OPENAI_DEPLOYMENT_NAME=
|
||||
# OPENAI_API_VERSION=2024-10-21 # optional, required for non-v1 API
|
||||
@@ -0,0 +1,68 @@
|
||||
# LLM Providers (set the one you use)
|
||||
OPENAI_API_KEY=
|
||||
GOOGLE_API_KEY=
|
||||
ANTHROPIC_API_KEY=
|
||||
XAI_API_KEY=
|
||||
DEEPSEEK_API_KEY=
|
||||
DASHSCOPE_API_KEY=
|
||||
DASHSCOPE_CN_API_KEY=
|
||||
ZHIPU_API_KEY=
|
||||
ZHIPU_CN_API_KEY=
|
||||
MINIMAX_API_KEY=
|
||||
MINIMAX_CN_API_KEY=
|
||||
OPENROUTER_API_KEY=
|
||||
MISTRAL_API_KEY=
|
||||
MOONSHOT_API_KEY=
|
||||
GROQ_API_KEY=
|
||||
NVIDIA_API_KEY=
|
||||
|
||||
# SEC EDGAR (US company filings, point-in-time). No key; a contact address SEC can reach you at.
|
||||
#SEC_EDGAR_USER_AGENT=Your Name your@email.com
|
||||
|
||||
# FRED (Federal Reserve macro data). Free key: https://fred.stlouisfed.org/docs/api/api_key.html
|
||||
#FRED_API_KEY=
|
||||
|
||||
# TypeSafe Jev: screens StockTwits and Reddit posts for the Sentiment Analyst. Optional.
|
||||
#TYPESAFE_API_KEY=
|
||||
#TYPESAFE_DEFAULT_MODEL=jev-latest
|
||||
|
||||
# Custom OpenAI-compatible endpoint (vLLM, LM Studio, llama.cpp). Local servers need no key.
|
||||
#OPENAI_COMPATIBLE_API_KEY=
|
||||
|
||||
# AWS Bedrock (pip install ".[bedrock]"). Bearer token, or the AWS credential chain; set the region either way.
|
||||
#AWS_BEARER_TOKEN_BEDROCK=
|
||||
#AWS_DEFAULT_REGION=us-west-2
|
||||
#AWS_PROFILE=
|
||||
|
||||
# Remote Ollama server. Unset uses http://localhost:11434/v1.
|
||||
#OLLAMA_BASE_URL=http://your-ollama-host:11434/v1
|
||||
|
||||
# Override these DEFAULT_CONFIG keys. Provider, models, language and round counts also skip their CLI prompt.
|
||||
#TRADINGAGENTS_LLM_PROVIDER=openai
|
||||
#TRADINGAGENTS_DEEP_THINK_LLM=gpt-6-sol
|
||||
#TRADINGAGENTS_QUICK_THINK_LLM=gpt-6-luna
|
||||
#TRADINGAGENTS_LLM_BACKEND_URL=
|
||||
#TRADINGAGENTS_OUTPUT_LANGUAGE=English
|
||||
#TRADINGAGENTS_MAX_DEBATE_ROUNDS=1
|
||||
#TRADINGAGENTS_MAX_RISK_ROUNDS=1
|
||||
#TRADINGAGENTS_CHECKPOINT_ENABLED=false
|
||||
|
||||
# Paths and alpha benchmark. Unset uses ~/.tradingagents and the regional index.
|
||||
#TRADINGAGENTS_RESULTS_DIR=
|
||||
#TRADINGAGENTS_CACHE_DIR=
|
||||
#TRADINGAGENTS_MEMORY_LOG_PATH=
|
||||
#TRADINGAGENTS_BENCHMARK_TICKER=SPY
|
||||
|
||||
# Lower temperature means less run-to-run variation on models that honor it.
|
||||
#TRADINGAGENTS_TEMPERATURE=0.0
|
||||
|
||||
# Retry budget for every LLM SDK. Raise it to ride out 429 throttling.
|
||||
#TRADINGAGENTS_LLM_MAX_RETRIES=6
|
||||
|
||||
# Cap output tokens to bound a model that runs long and trips a timeout.
|
||||
#TRADINGAGENTS_MAX_TOKENS=8192
|
||||
|
||||
# Reasoning depth per provider; setting one skips its prompt.
|
||||
#TRADINGAGENTS_OPENAI_REASONING_EFFORT=medium
|
||||
#TRADINGAGENTS_GOOGLE_THINKING_LEVEL=high
|
||||
#TRADINGAGENTS_ANTHROPIC_EFFORT=high
|
||||
@@ -0,0 +1,58 @@
|
||||
name: Bug report
|
||||
description: Something fails, crashes, or behaves differently from the docs.
|
||||
labels: ["bug"]
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Search [existing issues](https://github.com/TauricResearch/TradingAgents/issues?q=is%3Aissue) first, and check that the bug reproduces on the latest release.
|
||||
Remove API keys and other secrets from anything you paste.
|
||||
- type: input
|
||||
id: version
|
||||
attributes:
|
||||
label: Version
|
||||
description: Output of `pip show tradingagents | grep Version`, or `git rev-parse --short HEAD` for a source checkout.
|
||||
placeholder: "0.5.1"
|
||||
validations:
|
||||
required: true
|
||||
- type: dropdown
|
||||
id: entry
|
||||
attributes:
|
||||
label: How you run it
|
||||
options:
|
||||
- CLI (tradingagents)
|
||||
- Python package (TradingAgentsGraph)
|
||||
- Backtest (run_backtest or tradingagents backtest)
|
||||
- Docker
|
||||
validations:
|
||||
required: true
|
||||
- type: input
|
||||
id: models
|
||||
attributes:
|
||||
label: LLM provider and models
|
||||
placeholder: "openai, gpt-6-sol / gpt-6-luna"
|
||||
validations:
|
||||
required: true
|
||||
- type: input
|
||||
id: run
|
||||
attributes:
|
||||
label: Ticker and analysis date
|
||||
placeholder: "NVDA, 2026-09-23"
|
||||
- type: textarea
|
||||
id: what
|
||||
attributes:
|
||||
label: What happened, and what you expected
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: repro
|
||||
attributes:
|
||||
label: Steps or code to reproduce
|
||||
description: The smallest steps or snippet that shows the problem.
|
||||
render: python
|
||||
- type: textarea
|
||||
id: logs
|
||||
attributes:
|
||||
label: Error output
|
||||
description: The full traceback or log lines, as text rather than a screenshot.
|
||||
render: text
|
||||
@@ -0,0 +1,8 @@
|
||||
blank_issues_enabled: false
|
||||
contact_links:
|
||||
- name: Questions and showcases
|
||||
url: https://github.com/TauricResearch/TradingAgents/discussions
|
||||
about: Usage questions, ideas, and projects built on TradingAgents
|
||||
- name: Security vulnerability
|
||||
url: https://github.com/TauricResearch/TradingAgents/security/advisories/new
|
||||
about: Report privately; please do not open a public issue
|
||||
@@ -0,0 +1,52 @@
|
||||
name: Wrong or missing data
|
||||
description: A report shows a wrong price, figure, date or company, or data is missing for a symbol.
|
||||
labels: ["bug", "data"]
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Data is served as of the analysis date: a historical run does not see later prices, filings or news, and some sources (news, StockTwits, Reddit) only serve recent items. A report that says a source is unavailable for a past window is expected.
|
||||
- type: input
|
||||
id: symbol
|
||||
attributes:
|
||||
label: Ticker and market
|
||||
placeholder: "0700.HK (Hong Kong)"
|
||||
validations:
|
||||
required: true
|
||||
- type: input
|
||||
id: date
|
||||
attributes:
|
||||
label: Analysis date
|
||||
placeholder: "2026-09-23"
|
||||
validations:
|
||||
required: true
|
||||
- type: dropdown
|
||||
id: kind
|
||||
attributes:
|
||||
label: Which data
|
||||
multiple: true
|
||||
options:
|
||||
- Prices or indicators
|
||||
- Fundamentals or statements
|
||||
- News
|
||||
- Social (StockTwits, Reddit)
|
||||
- Macro (FRED)
|
||||
- Insider transactions
|
||||
- Prediction markets
|
||||
- Company identity (wrong company or name)
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: shown
|
||||
attributes:
|
||||
label: What the report shows, and what is correct
|
||||
description: Quote the report line, and give the correct value with a public source (exchange page, filing, Yahoo Finance link).
|
||||
validations:
|
||||
required: true
|
||||
- type: input
|
||||
id: version
|
||||
attributes:
|
||||
label: Version
|
||||
placeholder: "0.5.1"
|
||||
validations:
|
||||
required: true
|
||||
@@ -0,0 +1,25 @@
|
||||
name: Feature request
|
||||
description: A capability TradingAgents lacks, or a change to how it works.
|
||||
labels: ["enhancement"]
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Any OpenAI-compatible endpoint already works through the `openai_compatible` provider, and OpenRouter serves many hosted models, so a new LLM provider is often not needed.
|
||||
- type: textarea
|
||||
id: problem
|
||||
attributes:
|
||||
label: The problem
|
||||
description: What you are trying to do, and where TradingAgents falls short today.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: proposal
|
||||
attributes:
|
||||
label: Proposed change
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: alternatives
|
||||
attributes:
|
||||
label: Alternatives considered
|
||||
@@ -0,0 +1,67 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main, 'v[0-9]+.[0-9]+.[0-9]+']
|
||||
pull_request:
|
||||
|
||||
concurrency:
|
||||
group: ci-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
test:
|
||||
name: tests (py${{ matrix.python-version }}${{ matrix.tz && format(', {0}', matrix.tz) || '' }})
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: ["3.10", "3.11", "3.12", "3.13"]
|
||||
# One run away from UTC, so a test that only passes in UTC fails here.
|
||||
include:
|
||||
- python-version: "3.12"
|
||||
tz: America/New_York
|
||||
env:
|
||||
TZ: ${{ matrix.tz }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install (with dev extras)
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -e ".[dev]"
|
||||
- name: Run test suite
|
||||
run: pytest -q
|
||||
|
||||
smoke-install:
|
||||
name: clean-install smoke
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- name: Fresh install (no dev extras) and import
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install .
|
||||
# Catches undeclared runtime deps (e.g. #994 python-dotenv): a bare
|
||||
# install must import the package and the CLI module.
|
||||
python -c "import tradingagents, cli.main; print('clean-install import OK')"
|
||||
|
||||
lint:
|
||||
name: ruff (strict, full repo)
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- name: Install ruff
|
||||
run: pip install "ruff>=0.15"
|
||||
- name: Lint the repository
|
||||
# The repo is fully clean under the strict select, so we lint everything
|
||||
# (generated results/ is excluded via pyproject extend-exclude).
|
||||
run: ruff check .
|
||||
+221
-6
@@ -1,8 +1,223 @@
|
||||
env/
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
.DS_Store
|
||||
*.csv
|
||||
src/
|
||||
eval_results/
|
||||
eval_data/
|
||||
*.py[codz]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
share/python-wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.nox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
*.py.cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
cover/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
.pybuilder/
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# IPython
|
||||
profile_default/
|
||||
ipython_config.py
|
||||
|
||||
# pyenv
|
||||
# For a library or package, you might want to ignore these files since the code is
|
||||
# intended to run in multiple environments; otherwise, check them in:
|
||||
# .python-version
|
||||
|
||||
# pipenv
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||
# install all needed dependencies.
|
||||
# Pipfile.lock
|
||||
|
||||
# UV
|
||||
# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# uv.lock
|
||||
|
||||
# poetry
|
||||
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||
# poetry.lock
|
||||
# poetry.toml
|
||||
|
||||
# pdm
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
||||
# pdm recommends including project-wide configuration in pdm.toml, but excluding .pdm-python.
|
||||
# https://pdm-project.org/en/latest/usage/project/#working-with-version-control
|
||||
# pdm.lock
|
||||
# pdm.toml
|
||||
.pdm-python
|
||||
.pdm-build/
|
||||
|
||||
# pixi
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pixi.lock in version control.
|
||||
# pixi.lock
|
||||
# Pixi creates a virtual environment in the .pixi directory, just like venv module creates one
|
||||
# in the .venv directory. It is recommended not to include this directory in version control.
|
||||
.pixi
|
||||
|
||||
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
||||
__pypackages__/
|
||||
|
||||
# Celery stuff
|
||||
celerybeat-schedule
|
||||
celerybeat.pid
|
||||
|
||||
# Redis
|
||||
*.rdb
|
||||
*.aof
|
||||
*.pid
|
||||
|
||||
# RabbitMQ
|
||||
mnesia/
|
||||
rabbitmq/
|
||||
rabbitmq-data/
|
||||
|
||||
# ActiveMQ
|
||||
activemq-data/
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.envrc
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
.dmypy.json
|
||||
dmypy.json
|
||||
|
||||
# Pyre type checker
|
||||
.pyre/
|
||||
|
||||
# pytype static type analyzer
|
||||
.pytype/
|
||||
|
||||
# Cython debug symbols
|
||||
cython_debug/
|
||||
|
||||
# PyCharm
|
||||
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
||||
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||
# .idea/
|
||||
|
||||
# Abstra
|
||||
# Abstra is an AI-powered process automation framework.
|
||||
# Ignore directories containing user credentials, local state, and settings.
|
||||
# Learn more at https://abstra.io/docs
|
||||
.abstra/
|
||||
|
||||
# Visual Studio Code
|
||||
# Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore
|
||||
# that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. However, if you prefer,
|
||||
# you could uncomment the following to ignore the entire vscode folder
|
||||
# .vscode/
|
||||
|
||||
# Ruff stuff:
|
||||
.ruff_cache/
|
||||
|
||||
# PyPI configuration file
|
||||
.pypirc
|
||||
|
||||
# Marimo
|
||||
marimo/_static/
|
||||
marimo/_lsp/
|
||||
__marimo__/
|
||||
|
||||
# Streamlit
|
||||
.streamlit/secrets.toml
|
||||
|
||||
# Cache
|
||||
**/data_cache/
|
||||
|
||||
# Enterprise env file (secrets) and generated run reports
|
||||
.env.enterprise
|
||||
reports/
|
||||
|
||||
+611
@@ -0,0 +1,611 @@
|
||||
# Changelog
|
||||
|
||||
All notable changes to TradingAgents are documented here.
|
||||
|
||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
|
||||
and this project follows [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||
Changes that need action when upgrading are listed first in their release.
|
||||
|
||||
## [0.5.1] — 2026-09-24
|
||||
|
||||
A package layout organised by what each module holds, social posts screened by
|
||||
TypeSafe's Jev when a key is set, GPT-6 Sol and Luna as the default models, and
|
||||
fixes to run isolation, SEC EDGAR statements and historical runs.
|
||||
|
||||
### Upgrading from 0.5.0
|
||||
|
||||
Some modules moved, and the old import paths are gone. Update imports as follows:
|
||||
|
||||
- `tradingagents.dataflows.interface` is `tradingagents.dataflows.router`, and `dataflows.symbol_utils` is `dataflows.symbols`. `dataflows.utils` is gone: `get_current_date` is in `dataflows.date_window`, `safe_ticker_component` in `dataflows.symbols`.
|
||||
- Vendor modules live under `tradingagents.dataflows.vendors`: `yahoo` (`ohlcv`, `market`, `fundamentals`, `news`, `snapshot`, from the former `stockstats_utils`, `y_finance`, `yfinance_news` and `market_data_validator`), `alpha_vantage` (a package, from the `alpha_vantage_*` modules), and `sec_edgar`, `fred`, `polymarket`, `reddit`, `stocktwits`.
|
||||
- `tradingagents.agents.utils` is gone: the agent tools are in `agents.tools`, and `agent_utils`, `agent_states`, `rating` and `structured` are `agents.context`, `agents.state`, `agents.rating` and `agents.structured`.
|
||||
- The decision log is `tradingagents.decision_log` (was `agents.utils.memory`), and `cli.utils` is `cli.prompts`.
|
||||
- `backtest.summarize` takes a `run_backtest` result or the path of a decision log, in place of a `TradingMemoryLog`.
|
||||
- Removed: `SignalProcessor` (the rating is parsed by `process_signal`), the `create_social_media_analyst` alias (use `create_sentiment_analyst`), the unused `project_dir` config key, and the graph attributes `curr_state`, `ticker` and `log_states_dict`, which held the previous run's state.
|
||||
|
||||
### Added
|
||||
|
||||
- **Jev post screening.** With `TYPESAFE_API_KEY` set, TypeSafe's Jev reads each StockTwits and Reddit post the Sentiment Analyst fetches: posts that are not about the company are dropped, and each source opens with a count of the rest by stance. Without the key nothing changes. (#1376)
|
||||
- B3 tickers (`.SA`) are benchmarked against the Ibovespa. (#1366)
|
||||
|
||||
### Models
|
||||
|
||||
- GPT-6 Sol and GPT-6 Luna are the default deep and quick models, and Claude Opus 5.5 replaces Opus 5 in the picker. Opus 5 and GPT-5.4 Mini remain valid model IDs.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Several graphs in one process each read their own data vendors. (#1369)
|
||||
- SEC EDGAR cash flow statements find capital expenditure for filers that moved it to purchases of productive assets (NVIDIA since fiscal 2022, Amazon since 2016), whose recent periods read as empty. (#1370)
|
||||
- SEC EDGAR annual statements list fiscal years only.
|
||||
- A historical run is no longer told today's date, or a date after the run, in coverage notices and the instrument context.
|
||||
- The Fundamentals Analyst can call the insider transactions tool.
|
||||
- An unreachable Yahoo on insider transactions is reported as unavailable, not as a symbol without data.
|
||||
- A graph reused across runs, as in a backtest, no longer keeps every run's full state.
|
||||
- A checkpointed CLI run says whether it resumed or started fresh.
|
||||
- A blank path variable (`TRADINGAGENTS_RESULTS_DIR` and the like) keeps the default path.
|
||||
- StockTwits messages reach the prompt as plain text rather than HTML-escaped.
|
||||
- The test suite is independent of the developer's `.env`, time zone and network. (#1368, #1372)
|
||||
|
||||
### Contributors
|
||||
|
||||
Thanks to everyone who reported these or sent a fix:
|
||||
|
||||
[@codify88](https://github.com/codify88), [@davidalmeida90](https://github.com/davidalmeida90), [@duongylinh](https://github.com/duongylinh), [@jccl2](https://github.com/jccl2), [@yuina368](https://github.com/yuina368).
|
||||
|
||||
## [0.5.0] — 2026-09-18
|
||||
|
||||
Point-in-time integrity across every dated path, decisions that are recorded as
|
||||
they were made, backtesting over a grid of tickers and dates, the caller's
|
||||
portfolio as run input, and SEC EDGAR fundamentals served as filed.
|
||||
|
||||
### Highlights
|
||||
|
||||
- **Fundamentals as filed.** SEC EDGAR serves US company statements as they stood on the run's date: a period that has ended but has not been filed is not served, and a figure restated later still reads as first reported. Keyless, opt-in via the vendor chain.
|
||||
- **Backtesting.** `run_backtest` runs the pipeline over a ticker and date grid into its own decision log, and `summarize` scores the settled cells; `tradingagents backtest` does the same from the CLI.
|
||||
- **Portfolio context.** `propagate(..., portfolio=...)` and `--portfolio` let the trader, risk and portfolio agents size against real holdings. A run without one is never treated as a flat book.
|
||||
- **Decisions are recorded as made.** An unreadable decision is flagged for review everywhere instead of becoming a tradeable Hold, and a rating argued against is no longer read as the call.
|
||||
|
||||
### Point-in-time and honest attribution
|
||||
|
||||
- Dated tools take the run's date from graph state, so an omitted or later date cannot reach a vendor. (#1331, #1319, #1118)
|
||||
- Insider filings and prediction-market odds are bounded by the run date; insider rows state that a trade becomes public when its Form 4 is filed.
|
||||
- A feed that never observed a window reports it as unavailable rather than as an absence, across news, Reddit and StockTwits.
|
||||
- The resolved company identity says when it describes today rather than the run date.
|
||||
- The verification snapshot quotes the prices the vendor reported, never a gap-filled value.
|
||||
- A vendor failure is a vendor failure: yfinance raises instead of returning its errors as text, an outage is not reported as a company with no data, and a chain where every vendor is unavailable says so instead of ending the run.
|
||||
- The macro vintage pin is clamped to the vendor's own clock, so a run dated today cannot ask for a vintage it does not have.
|
||||
- A historical run is not served a present-day company profile by either fundamentals vendor. (#1300)
|
||||
|
||||
### Decisions and evaluation
|
||||
|
||||
- The labelled rating decides, whatever separates it, and prose naming several ratings is reviewed rather than guessed.
|
||||
- Decision prompts state the shape of their answer, so a provider without structured output still returns a readable decision.
|
||||
- A report that was not produced says so, instead of appearing as an empty section.
|
||||
- Backtest scoring reads the direction each rating claimed: a Sell that fell is a hit, and Hold reports no hit rate.
|
||||
- The outcome window is configurable (`holding_period_days`), and reflection states the window it judges.
|
||||
- A settled decision is not logged twice, and a failed reflection no longer stops the next run. (#645)
|
||||
- The trader states entry and stop levels as prices, so a percentage no longer fails the whole proposal. (#1288)
|
||||
|
||||
### CLI
|
||||
|
||||
- `tradingagents backtest`, with `--run-id` to continue an interrupted sweep. (#1234)
|
||||
- The previous run's selections come back as prompt defaults. (#1236, #920)
|
||||
- A run with no readable rating says so; the live view no longer scrolls; messages that read like Python values are shown. (#649, #784)
|
||||
- The state log keeps non-ASCII readable. (#1081)
|
||||
|
||||
### Data sources
|
||||
|
||||
- SEC EDGAR fundamentals vendor (US filers, keyless).
|
||||
- Hong Kong and Shanghai tickers resolve to the symbols Yahoo serves. (#1342, #957, #1260)
|
||||
- Reddit is fetched as one combined request per run. (#1286)
|
||||
- One OHLCV cache file per symbol. (#1330)
|
||||
|
||||
### Models
|
||||
|
||||
- Current lineups for every provider: GPT-6 Astra and the GPT-5.6 family, Gemini 3.8 Flash, Claude Opus 5 and Fable 5.1, Grok 4.6, DeepSeek Flash, Qwen 3.8, GLM-5.3, MiniMax M3, Kimi K3 and the current Mistral snapshots.
|
||||
- Every provider accepts a model ID the picker does not list.
|
||||
- GLM traffic goes to the platform its key belongs to, and Ollama structured output no longer sends a tool_choice it rejects. (#1062)
|
||||
|
||||
### Changed
|
||||
|
||||
- The memory log records `REVIEW` for a decision with no readable rating, where it previously recorded `Hold`.
|
||||
- Optional fields the model did not provide are named as such rather than omitted.
|
||||
- Removed dependencies nothing imports: backtrader, redis, setuptools, langchain-experimental, parsel, tqdm. (#1353, #1070)
|
||||
|
||||
### Contributors
|
||||
|
||||
Thanks to everyone who reported these or sent a fix:
|
||||
|
||||
[@akashkpfreelancer](https://github.com/akashkpfreelancer), [@angziii](https://github.com/angziii), [@anupamme](https://github.com/anupamme), [@AyushKar2005](https://github.com/AyushKar2005), [@bulkypanda](https://github.com/bulkypanda), [@CadeYu](https://github.com/CadeYu), [@chiang21fcb](https://github.com/chiang21fcb), [@dajiaohuang](https://github.com/dajiaohuang), [@dewrama](https://github.com/dewrama), [@DogInfantry](https://github.com/DogInfantry), [@emitov](https://github.com/emitov), [@farukerdem34](https://github.com/farukerdem34), [@flydragon2018](https://github.com/flydragon2018), [@fusshell](https://github.com/fusshell), [@Ganesh1729-ui](https://github.com/Ganesh1729-ui), [@gyx09212214-prog](https://github.com/gyx09212214-prog), [@hamzabudeir](https://github.com/hamzabudeir), [@ihsieh31](https://github.com/ihsieh31), [@jaylew20250206](https://github.com/jaylew20250206), [@kaushik-yadav](https://github.com/kaushik-yadav), [@kbnnf](https://github.com/kbnnf), [@kevinkda](https://github.com/kevinkda), [@LudwigJMarx](https://github.com/LudwigJMarx), [@lx7720](https://github.com/lx7720), [@malandrindev](https://github.com/malandrindev), [@mhd325ic-hash](https://github.com/mhd325ic-hash), [@minhdn90](https://github.com/minhdn90), [@miznan](https://github.com/miznan), [@mmssix](https://github.com/mmssix), [@mrbob-git](https://github.com/mrbob-git), [@newnewself](https://github.com/newnewself), [@prithvirajrh](https://github.com/prithvirajrh), [@PyriteResearch](https://github.com/PyriteResearch), [@Rajatendu1](https://github.com/Rajatendu1), [@Recnelis0](https://github.com/Recnelis0), [@Rodvask](https://github.com/Rodvask), [@samhoooo](https://github.com/samhoooo), [@sheiun-xu](https://github.com/sheiun-xu), [@shivsin25](https://github.com/shivsin25), [@SmileShaun](https://github.com/SmileShaun), [@SonnyRajagopalan](https://github.com/SonnyRajagopalan), [@taro0915](https://github.com/taro0915), [@wupengbo125](https://github.com/wupengbo125), [@wxggzz](https://github.com/wxggzz), [@Yixiang-Wu](https://github.com/Yixiang-Wu), [@ZahirBodrike](https://github.com/ZahirBodrike), [@ZHUYAWEI](https://github.com/ZHUYAWEI), [@zkwang616](https://github.com/zkwang616).
|
||||
|
||||
## [0.4.0] — 2026-08-31
|
||||
|
||||
Look-ahead and point-in-time fixes across the data and memory layers, clearer
|
||||
decision signals, working CLI checkpoint resume, and the GPT-5.6 / GLM-5.3 models.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **FRED macro look-ahead.** Historical macro requests were served from today's
|
||||
data vintage, leaking later revisions into a backtest; both the observations
|
||||
and metadata requests now pin the vintage to the as-of date. (#1275)
|
||||
- **Social sentiment look-ahead.** StockTwits and Reddit were fetched with no
|
||||
date, so a historical run showed today's chatter as if it were from the as-of
|
||||
date; the social path is now trimmed to the analysis window, via one shared
|
||||
UTC half-open window rule (`dataflows/date_window`) used by news too. (#1220)
|
||||
- **Memory point-in-time guard.** `get_past_context` returned every resolved
|
||||
lesson regardless of the run date; each resolved entry now records the date
|
||||
its outcome became known, and a historical run only sees lessons resolved by
|
||||
the trade date. (#1251)
|
||||
- **Premature reflection.** A decision was settled on a partial return if a rerun
|
||||
happened before its holding window fully traded; resolution now waits for the
|
||||
full window. (#1169)
|
||||
- **Latest OHLCV bar dropped.** The newest bar with a NaN close was silently
|
||||
dropped before the date cutoff, making the previous trading day look like the
|
||||
latest; dates are normalized per element (DST- and non-US-market safe) and a
|
||||
missing latest close raises rather than falling back. (#1201)
|
||||
- **Debate opening fabrication.** The first speaker in each debate round rebutted
|
||||
an empty opponent response, fabricating the other side; all five debators now
|
||||
open with their own case when no opponent has spoken. (#1176)
|
||||
- **Silent Hold.** An unparseable Portfolio Manager rating (including a fullwidth
|
||||
colon) was coerced to a tradeable Hold; it now surfaces a `REVIEW` sentinel,
|
||||
with `parse_rating` keeping its silent default for compatibility callers. (#1170)
|
||||
- **`--checkpoint` was a no-op on the CLI.** Checkpoint setup lived only in
|
||||
`propagate()`; the CLI streamed the checkpointer-less graph. The lifecycle is
|
||||
now shared, and a resume feeds `None` so LangGraph continues the interrupted
|
||||
run instead of duplicating messages. (#1249)
|
||||
- **DeepSeek via OpenRouter.** `deepseek/<id>` fell through to default
|
||||
capabilities and had object-form `tool_choice` forced on it; the official
|
||||
namespace is stripped so it reuses the native DeepSeek quirks. (#1199)
|
||||
- **Trader price grounding.** The Trader saw only the digested plan; it now also
|
||||
receives the technical market report so entry/stop levels anchor to real price
|
||||
structure. (#1167)
|
||||
|
||||
### Added
|
||||
|
||||
- **Configurable output-token cap.** `max_tokens` / `TRADINGAGENTS_MAX_TOKENS`,
|
||||
forwarded to every provider (Gemini as `max_output_tokens`), so a model that
|
||||
emits unbounded reasoning can be bounded instead of hanging. (#1204)
|
||||
- **Latest models.** Added the GPT-5.6 family (`gpt-5.6` / `gpt-5.6-terra` /
|
||||
`gpt-5.6-luna`) and GLM-5.3 (`glm-5.3`, `glm-5.3-flash`). The default models
|
||||
are now `gpt-5.6` (deep) and `gpt-5.6-luna` (quick).
|
||||
|
||||
### Contributors
|
||||
|
||||
Thanks to everyone who reported these or sent a fix:
|
||||
|
||||
[@PyriteResearch](https://github.com/PyriteResearch), [@yiran1268](https://github.com/yiran1268), [@fabiolenine](https://github.com/fabiolenine), [@lx7720](https://github.com/lx7720), [@taro0915](https://github.com/taro0915), [@Jaswanth-Sriram-Veturi](https://github.com/Jaswanth-Sriram-Veturi), [@ariesy](https://github.com/ariesy), [@liangzj1999](https://github.com/liangzj1999), [@zkwang616](https://github.com/zkwang616), [@aniketshukla1](https://github.com/aniketshukla1), [@loulanyue](https://github.com/loulanyue), [@hudsonwa](https://github.com/hudsonwa), [@daleselaji-dev](https://github.com/daleselaji-dev), [@wolfoswald777-crypto](https://github.com/wolfoswald777-crypto).
|
||||
|
||||
## [0.3.1] — 2026-07-05
|
||||
|
||||
Correctness and stability patch: data look-ahead, graph-router crash-safety,
|
||||
checkpoint identity, crypto sentiment sources, and configurable resilience.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Alpha Vantage look-ahead filter now runs.** The fundamentals payload is a
|
||||
JSON string, so the dict-only guard skipped filtering and future-dated reports
|
||||
leaked into historical runs; parse before filtering. (#1115, @zachthebird)
|
||||
- **News analyst prompt matches the tool.** The prompt advertised
|
||||
`get_news(query, ...)` but the tool takes a ticker; aligned to stop
|
||||
hallucinated free-text query calls. (#1116, @shcheuk)
|
||||
- **Shared debate/risk routers can't crash mid-run.** Both routers return more
|
||||
targets than any one edge mapped; every edge now shares the complete path map,
|
||||
so a fall-through under prompt/i18n/refactor drift stays routable.
|
||||
(#1088, @Fr3ya, @sa7an7, @Sushanth012)
|
||||
- **Checkpoint resume respects graph shape.** The thread id folds in selected
|
||||
analysts, debate/risk depth, and asset mode, so a resume under different
|
||||
choices no longer continues the wrong graph. (#1089, @bossjoker1, @Ghraven)
|
||||
- **Crypto sentiment sources resolve.** StockTwits lists crypto as `<BASE>.X`
|
||||
(Yahoo's `BTC-USD` 404s) and Reddit needs the base symbol to match; the social
|
||||
path now maps crypto correctly for both. (#1113, @suremadoreai)
|
||||
|
||||
### Added
|
||||
|
||||
- **Configurable LLM retry budget.** `llm_max_retries` /
|
||||
`TRADINGAGENTS_LLM_MAX_RETRIES` is forwarded to every provider, so a transient
|
||||
429 burst no longer aborts a run. (#1091, @yanggaome)
|
||||
- **Bedrock API-key auth.** `AWS_BEARER_TOKEN_BEDROCK` authenticates Amazon
|
||||
Bedrock without AWS access keys and takes precedence over an ambient
|
||||
`AWS_PROFILE`. (#1103, @praxstack)
|
||||
- **Latest Claude models.** Added Claude Sonnet 5 (`claude-sonnet-5`) and
|
||||
Fable 5 (`claude-fable-5`); effort control now covers the Claude 5 line.
|
||||
|
||||
## [0.3.0] — 2026-06-22
|
||||
|
||||
Stabilization and extensibility release: a CI gate, a unified verified
|
||||
data-access contract, a provider and data-vendor registry, and a maintenance
|
||||
sweep that hardened config precedence, the model catalog, data resilience, and
|
||||
structured output.
|
||||
|
||||
### Added
|
||||
|
||||
- **CI gate.** GitHub Actions runs the pytest suite across Python 3.10-3.13,
|
||||
strict `ruff`, and a clean-install smoke that imports the package and CLI to
|
||||
catch undeclared dependencies. (#994, #197)
|
||||
- **Provider registry.** OpenAI-compatible providers register as a single spec,
|
||||
and a generic `openai_compatible` endpoint covers vLLM, LM Studio, and relays.
|
||||
Adds NVIDIA NIM, Kimi, Groq, Mistral, and a native Amazon Bedrock client.
|
||||
- **Macro and prediction-market vendors.** FRED macro indicators and Polymarket
|
||||
event probabilities, surfaced to the news and macro analysts.
|
||||
- **Programmatic report output.** `TradingAgentsGraph.save_reports()` writes the
|
||||
same report tree the CLI produces, for headless and API runs. (#1037)
|
||||
- **Env-configurable reasoning depth** via `TRADINGAGENTS_OPENAI_REASONING_EFFORT`,
|
||||
`TRADINGAGENTS_GOOGLE_THINKING_LEVEL`, and `TRADINGAGENTS_ANTHROPIC_EFFORT`,
|
||||
each gated to the models that accept it.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Verified data-access contract.** Symbol normalization on every vendor path
|
||||
(identity, returns, CLI, news); the configured vendor list is the exact
|
||||
resolution chain with no silent fallback to unselected vendors; a typed
|
||||
`VendorError` taxonomy; look-ahead-safe news windows; stale-OHLCV rejection;
|
||||
inclusive yfinance date ranges.
|
||||
- **Config precedence.** An explicit `TRADINGAGENTS_*` value or CLI flag now wins
|
||||
over interactive defaults for debate and risk round counts,
|
||||
`--checkpoint / --no-checkpoint`, and the Docker provider profile; invalid
|
||||
boolean env values fail loudly. (#975, #976, #977)
|
||||
- **Current-generation model catalog.** Refreshed provider lineups; retired
|
||||
`gpt-4.1`, Claude Sonnet 4.5, and the Gemini 2.5 line.
|
||||
- **Optional vendors degrade** instead of aborting a run: a failed macro or
|
||||
prediction-market lookup returns a no-data sentinel.
|
||||
- **Analyst prompts lead with the current date** so tool-call date ranges anchor
|
||||
to the run date rather than the model's training cutoff. (#836)
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Instrument identity.** Deterministic ticker-to-company resolution prevents
|
||||
wrong-company hallucination, and a verified market-data snapshot grounds price
|
||||
and indicator claims. (#814, #830)
|
||||
- **Social and market data sources.** Reddit RSS-first with 429 backoff,
|
||||
StockTwits transport hardening, and Alpha Vantage timeout plus
|
||||
key-versus-rate-limit handling.
|
||||
- **Structured output.** Local OpenAI-compatible servers no longer reject
|
||||
object-form `tool_choice`; a thinking model that returns no parsed result falls
|
||||
back to free text; null-ish strings in optional price fields coerce to `None`.
|
||||
(#1038, #1051, #1057)
|
||||
|
||||
### Removed
|
||||
|
||||
- The no-op `analyst_concurrency_limit` config knob; parallel analyst execution
|
||||
is planned for a later release. (#979)
|
||||
- The unused committed `uv.lock`. (#1030)
|
||||
|
||||
### Contributors
|
||||
|
||||
Thanks to everyone who shaped this release through code, design, and reports:
|
||||
|
||||
[@CadeYu](https://github.com/CadeYu), [@Zavianx](https://github.com/Zavianx), [@weijianz-opc](https://github.com/weijianz-opc), [@naltun](https://github.com/naltun), [@brahmasky](https://github.com/brahmasky), [@nik2208](https://github.com/nik2208), [@thieucong98](https://github.com/thieucong98), [@Derekko-web](https://github.com/Derekko-web), [@LukiPrince](https://github.com/LukiPrince), [@Eddieargenal](https://github.com/Eddieargenal), [@Ghraven](https://github.com/Ghraven), [@ms32035](https://github.com/ms32035), [@yting27](https://github.com/yting27), [@nyxst4ck](https://github.com/nyxst4ck), [@KenCheung-AIxFinance](https://github.com/KenCheung-AIxFinance), [@yangyusheng2n](https://github.com/yangyusheng2n), [@fareloj](https://github.com/fareloj), [@haosenwang1018](https://github.com/haosenwang1018), [@octo-patch](https://github.com/octo-patch), [@seifenk](https://github.com/seifenk), [@CaoYuhaoCarl](https://github.com/CaoYuhaoCarl), [@mihailnica10](https://github.com/mihailnica10), [@Dado-hash](https://github.com/Dado-hash), [@Handsomemikezzz](https://github.com/Handsomemikezzz), [@ydhawesome](https://github.com/ydhawesome), [@macd2](https://github.com/macd2), [@AyushKar2005](https://github.com/AyushKar2005), [@wildhuman](https://github.com/wildhuman), [@robert23kim](https://github.com/robert23kim), [@bngness](https://github.com/bngness), [@tedix-rodrigo](https://github.com/tedix-rodrigo), [@malaccan](https://github.com/malaccan), [@rfalken78](https://github.com/rfalken78), [@dengli1971-droid](https://github.com/dengli1971-droid), [@proofconcept39](https://github.com/proofconcept39), [@prasta1](https://github.com/prasta1), [@liximin](https://github.com/liximin), [@jeffhuen](https://github.com/jeffhuen), [@mazar](https://github.com/mazar), [@soyangelromero](https://github.com/soyangelromero), [@CNQQC](https://github.com/CNQQC), [@dovetaill](https://github.com/dovetaill), [@fperdigon](https://github.com/fperdigon), [@gyx09212214-prog](https://github.com/gyx09212214-prog), [@RSXLX](https://github.com/RSXLX).
|
||||
|
||||
## [0.2.5] — 2026-05-11
|
||||
|
||||
### Added
|
||||
|
||||
- **Grounded Sentiment Analyst.** The renamed `sentiment_analyst` now reads
|
||||
real Yahoo News, StockTwits, and Reddit data before generating its report,
|
||||
replacing the prior flow that could fabricate social posts under prompt
|
||||
pressure. (#557, #607)
|
||||
- **MiniMax provider** with the full M2.x catalog (M2.7 / M2.5 / M2.1 / M2
|
||||
plus highspeed variants, 204K context). Dual-region: Global
|
||||
(`MINIMAX_API_KEY`) and China (`MINIMAX_CN_API_KEY`).
|
||||
- **Dual-region Qwen and GLM** with separate keys per region — international
|
||||
(`DASHSCOPE_API_KEY`, `ZHIPU_API_KEY`) and China (`DASHSCOPE_CN_API_KEY`,
|
||||
`ZHIPU_CN_API_KEY`), selectable via a secondary region prompt. (#758)
|
||||
- **`TRADINGAGENTS_*` env-var configurability for `DEFAULT_CONFIG`.** Override
|
||||
`llm_provider`, deep/quick model IDs, `backend_url`, `output_language`,
|
||||
debate-round counts, checkpoint flag, and benchmark ticker via `.env` with
|
||||
type-aware coercion (string / int / bool). (#602)
|
||||
- **Interactive API-key detection in the CLI.** When the selected provider's
|
||||
key is missing, the CLI prompts for it and persists the value to `.env`
|
||||
so the analysis run continues without restart.
|
||||
- **Remote Ollama support.** `OLLAMA_BASE_URL` points the CLI and the
|
||||
programmatic client at a remote `ollama-serve`. The CLI surfaces the
|
||||
resolved endpoint and warns on common malformed inputs. Adds a
|
||||
`"Custom model ID"` option for models pulled via `ollama pull`. (#648, #768)
|
||||
- **Configurable news-fetch parameters** in `DEFAULT_CONFIG` — per-ticker
|
||||
article limit, macro headline limit, lookback window, and macro search
|
||||
queries. (#606, #683)
|
||||
- **Configurable alpha benchmark** for non-US tickers. Replaces hardcoded
|
||||
SPY with regional indices for `.NS` (^NSEI), `.T` (^N225), `.HK` (^HSI),
|
||||
`.L` (^FTSE), `.TO` (^GSPTSE), `.AX` (^AXJO), `.BO` (^BSESN); explicit
|
||||
`benchmark_ticker` override available. Eliminates FX drift dominating
|
||||
alpha for non-USD listings. (#628, #684)
|
||||
- **Multi-language output covers every user-facing agent** — researchers,
|
||||
risk debators, research manager, and trader, ending the previous
|
||||
partial-localization reports. (#575)
|
||||
- **Model catalog refresh.** OpenAI GPT-5.5 frontier, Anthropic Claude Opus
|
||||
4.7, Gemini 3.1 Flash-Lite GA, xAI Grok 4.20, Qwen 3.6 line. Versioned IDs
|
||||
only; auto-shifting aliases moved to the `"Custom model ID"` option.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Sentiment Analyst** is now consistently named across the CLI dropdown,
|
||||
status panel, and final reports (previously the backend was renamed but
|
||||
the CLI still said "Social Analyst"). The `AnalystType.SOCIAL = "social"`
|
||||
wire value is kept for saved-config back-compat.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Structured output works on DeepSeek V4 / reasoner and MiniMax M2.x.**
|
||||
Those providers reject `tool_choice` per their tool-calling docs; the
|
||||
binding flow now skips it automatically via a capability table.
|
||||
- **`pip install .` installations pick up the project `.env`** when running
|
||||
the CLI as a console script. (#747)
|
||||
- **Reports save end-to-end** — streamed chunks were previously dropped from
|
||||
`complete_report.md`. (#719, #736)
|
||||
- **Ticker prompt preserves exchange suffixes** (`.SH`, `.SZ`, `.SS`, `.HK`,
|
||||
`.T`, etc.) for A-share, HK, Tokyo, and other non-US flows. (#770)
|
||||
- **Docker permission errors** no longer block first-run write to
|
||||
`~/.tradingagents/`. (#519, #627, #672, #771)
|
||||
- **Config state no longer leaks between runs** when sub-dicts are mutated;
|
||||
`set_config` partial updates preserve sibling defaults. (#788)
|
||||
- **`max_recur_limit` config actually applies** — previously read but not
|
||||
forwarded to the propagator. (#764)
|
||||
- **Missing-API-key error** names the exact env var to set. (#680)
|
||||
- **Quieter startup** — suppressed the noisy upstream
|
||||
`LangChainPendingDeprecationWarning` from langgraph-checkpoint; will be
|
||||
removed once that package ships its fix.
|
||||
|
||||
### Security
|
||||
|
||||
- **Ticker path-traversal validation** at every filesystem-path site (cache,
|
||||
checkpoint database, results) so a malicious ticker cannot escape its
|
||||
intended directory. (#618)
|
||||
|
||||
## [0.2.4] — 2026-04-25
|
||||
|
||||
### Added
|
||||
|
||||
- **Structured-output decision agents.** Research Manager, Trader, and Portfolio
|
||||
Manager now use `llm.with_structured_output(Schema)` on their primary call
|
||||
and return typed Pydantic instances. Each provider's native structured-output
|
||||
mode is used (`json_schema` for OpenAI / xAI, `response_schema` for Gemini,
|
||||
tool-use for Anthropic, function-calling for OpenAI-compatible providers).
|
||||
Render helpers preserve the existing markdown shape so memory log, CLI
|
||||
display, and saved reports keep working unchanged. (#434)
|
||||
- **LangGraph checkpoint resume** — opt-in via `--checkpoint`. State is saved
|
||||
after each node so crashed or interrupted runs resume from the last
|
||||
successful step. Per-ticker SQLite databases under
|
||||
`~/.tradingagents/cache/checkpoints/`. `--clear-checkpoints` resets them. (#594)
|
||||
- **Persistent decision log** replacing the per-agent BM25 memory. Decisions
|
||||
are stored automatically at the end of `propagate()`; the next same-ticker
|
||||
run resolves prior pending entries with realised return, alpha vs SPY, and
|
||||
a one-paragraph reflection. Override path with `TRADINGAGENTS_MEMORY_LOG_PATH`.
|
||||
Optional `memory_log_max_entries` config caps resolved entries; pending
|
||||
entries are never pruned. (#578, #563, #564, #579)
|
||||
- **DeepSeek, Qwen (Alibaba DashScope), GLM (Zhipu), and Azure OpenAI**
|
||||
providers, plus dynamic OpenRouter model selection.
|
||||
- **Docker support** — multi-stage build with separate dev and runtime images.
|
||||
- **`scripts/smoke_structured_output.py`** — diagnostic that exercises the
|
||||
three structured-output agents against any provider so contributors can
|
||||
verify their setup with one command.
|
||||
- **5-tier rating scale** (Buy / Overweight / Hold / Underweight / Sell) used
|
||||
consistently by Research Manager, Portfolio Manager, signal processor, and
|
||||
the memory log; Trader keeps 3-tier (Buy / Hold / Sell) since transaction
|
||||
direction is naturally ternary.
|
||||
- **Pytest fixtures** — lazy LLM client imports plus placeholder API keys so
|
||||
the test suite runs cleanly without credentials. (#588)
|
||||
|
||||
### Changed
|
||||
|
||||
- **`backend_url` default is now `None`** rather than the OpenAI URL. Each
|
||||
provider client falls back to its native default. The previous default
|
||||
leaked the OpenAI URL into non-OpenAI clients (e.g. Gemini), producing
|
||||
malformed request URLs for Python users who switched providers without
|
||||
overriding `backend_url`. The CLI flow is unaffected.
|
||||
- All file I/O passes explicit `encoding="utf-8"` so Windows users no longer
|
||||
hit `UnicodeEncodeError` with the cp1252 default. (#543, #550, #576)
|
||||
- Cache and log directories moved to `~/.tradingagents/` to resolve Docker
|
||||
permission issues. (#519)
|
||||
- `SignalProcessor` reads the rating from the Portfolio Manager's rendered
|
||||
markdown via a deterministic heuristic — no extra LLM call.
|
||||
- OpenAI structured-output calls default to `method="function_calling"` to
|
||||
avoid noisy `PydanticSerializationUnexpectedValue` warnings emitted by
|
||||
langchain-openai's Responses-API parse path. Same typed result, no warnings.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Empty memory no longer triggers fabricated past-lessons in agent prompts;
|
||||
the memory-log redesign makes this structurally impossible since only the
|
||||
Portfolio Manager consults memory and only when entries exist. (#572)
|
||||
- Tool-call logging processes every chunk message, not just the last one, and
|
||||
memory score normalization handles empty score arrays. (#534, #531)
|
||||
|
||||
### Removed
|
||||
|
||||
- `FinancialSituationMemory` (the per-agent BM25 system) and the dead
|
||||
`reflect_and_remember()` plumbing; subsumed by the persistent decision log.
|
||||
- Hardcoded Google endpoint that caused 404 when `langchain-google-genai`
|
||||
changed its API path. (#493, #496)
|
||||
|
||||
### Contributors
|
||||
|
||||
Thanks to everyone who shaped this release through code, design, and reports:
|
||||
|
||||
- [@claytonbrown](https://github.com/claytonbrown) — checkpoint resume (#594), test fixtures (#588), design feedback on cost tracking (#582) and structured validation (#583)
|
||||
- [@Bcardo](https://github.com/Bcardo) — memory-log redesign (#579), empty-memory hallucination report (#572), encoding fix proposal (#570)
|
||||
- [@voidborne-d](https://github.com/voidborne-d) — memory persistence design (#564), portfolio manager state fix (#503)
|
||||
- [@mannubaveja007](https://github.com/mannubaveja007) — structured-output feature request (#434)
|
||||
- [@kelder66](https://github.com/kelder66) — RAM-only memory issue (#563)
|
||||
- [@Gujiassh](https://github.com/Gujiassh) — tool-call logging fix (#534), test stub PR (#533)
|
||||
- [@iuyup](https://github.com/iuyup) — memory score normalization fix (#531)
|
||||
- [@kaihg](https://github.com/kaihg) — Google base_url fix (#496)
|
||||
- [@32ryh98yfe](https://github.com/32ryh98yfe) — Gemini 404 report (#493)
|
||||
- [@uppb](https://github.com/uppb) — OpenRouter dynamic model selection (#482)
|
||||
- [@guoz14](https://github.com/guoz14) — OpenRouter limited-model report (#337)
|
||||
- [@samchenku](https://github.com/samchenku) — indicator name normalization (#490)
|
||||
- [@JasonOA888](https://github.com/JasonOA888) — y_finance pandas import fix (#488)
|
||||
- [@tiffanychum](https://github.com/tiffanychum) — stale import cleanup (#499)
|
||||
- [@zaizou](https://github.com/zaizou) — Docker permission issue (#519)
|
||||
- [@Stosman123](https://github.com/Stosman123), [@mauropuga](https://github.com/mauropuga), [@hotwind2015](https://github.com/hotwind2015) — Windows encoding bug reports (#543, #550, #576)
|
||||
- [@nnishad](https://github.com/nnishad), [@atharvajoshi01](https://github.com/atharvajoshi01) — encoding fix proposals (#568, #549)
|
||||
|
||||
## [0.2.3] — 2026-03-29
|
||||
|
||||
### Added
|
||||
|
||||
- **Multi-language output** for analyst reports and final decisions, with a
|
||||
CLI selector. Internal agent debate stays in English for reasoning quality. (#472)
|
||||
- **GPT-5.4 family models** in the default catalog, with deep/quick model split.
|
||||
- **Unified model catalog** as a single source of truth for CLI options and
|
||||
provider validation.
|
||||
|
||||
### Changed
|
||||
|
||||
- `base_url` is forwarded to Google and Anthropic clients so corporate proxies
|
||||
work consistently across providers. (#427)
|
||||
- Standardised the Google `api_key` parameter to the unified `api_key` form.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Backtesting fetchers no longer leak look-ahead data when `curr_date` is in
|
||||
the middle of a fetched window. (#475)
|
||||
- Invalid indicator names from the LLM are caught at the tool boundary instead
|
||||
of crashing the run. (#429)
|
||||
- yfinance news fetchers respect the same exponential-backoff retry as price
|
||||
fetchers. (#445)
|
||||
|
||||
### Contributors
|
||||
|
||||
- [@ahmedk20](https://github.com/ahmedk20) — multi-language output (#472)
|
||||
- [@CadeYu](https://github.com/CadeYu) — model catalog typing (#464)
|
||||
- [@javierdejesusda](https://github.com/javierdejesusda) — unified Google API key parameter (#453)
|
||||
- [@voidborne-d](https://github.com/voidborne-d) — yfinance news retry (#445)
|
||||
- [@kostakost2](https://github.com/kostakost2) — look-ahead bias report (#475)
|
||||
- [@lu-zhengda](https://github.com/lu-zhengda) — proxy/base_url support request (#427)
|
||||
- [@VamsiKrishna2021](https://github.com/VamsiKrishna2021) — invalid indicator crash report (#429)
|
||||
|
||||
## [0.2.2] — 2026-03-22
|
||||
|
||||
### Added
|
||||
|
||||
- **Five-tier rating scale** (Buy / Overweight / Hold / Underweight / Sell)
|
||||
introduced for the Portfolio Manager.
|
||||
- **Anthropic effort level** support for Claude models.
|
||||
- **OpenAI Responses API** path for native OpenAI models.
|
||||
|
||||
### Changed
|
||||
|
||||
- `risk_manager` renamed to `portfolio_manager` to match the role description
|
||||
shown in the CLI display.
|
||||
- Exchange-qualified tickers (e.g. `7203.T`, `BRK.B`) preserved across all
|
||||
agent prompts and tool calls.
|
||||
- Process-level UTF-8 default attempted for cross-platform consistency
|
||||
(note: this approach did not actually take effect; replaced in v0.2.4 with
|
||||
explicit per-call `encoding="utf-8"` arguments).
|
||||
|
||||
### Fixed
|
||||
|
||||
- yfinance rate-limit errors are retried with exponential backoff. (#426)
|
||||
- HTTP client SSL customisation is supported for environments that need
|
||||
custom certificate bundles. (#379)
|
||||
- Report-section writes handle list-of-string content gracefully.
|
||||
|
||||
### Contributors
|
||||
|
||||
- [@CadeYu](https://github.com/CadeYu) — exchange-qualified ticker preservation (#413)
|
||||
- [@yang1002378395-cmyk](https://github.com/yang1002378395-cmyk) — HTTP client SSL customisation (#379)
|
||||
|
||||
## [0.2.1] — 2026-03-15
|
||||
|
||||
### Security
|
||||
|
||||
- Patched `langchain-core` vulnerability (LangGrinch). (#335)
|
||||
- Removed `chainlit` dependency affected by CVE-2026-22218.
|
||||
|
||||
### Added
|
||||
|
||||
- `pyproject.toml` build-system configuration; the project now installs via
|
||||
modern packaging tooling.
|
||||
|
||||
### Removed
|
||||
|
||||
- `setup.py` — dependencies consolidated to `pyproject.toml`.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Risk manager reads the correct fundamental report source. (#341)
|
||||
- All `open()` calls receive an explicit UTF-8 encoding (initial pass).
|
||||
- `get_indicators` tool handles comma-separated indicator names from the LLM. (#368)
|
||||
- `Propagation` initialises every debate-state field so risk debaters never
|
||||
see missing keys.
|
||||
- Stock data parsing tolerates malformed CSVs and NaN values.
|
||||
- Conditional debate logic respects the configured round count. (#361)
|
||||
|
||||
### Contributors
|
||||
|
||||
- [@RinZ27](https://github.com/RinZ27) — `langchain-core` security patch (#335)
|
||||
- [@Ljx-007](https://github.com/Ljx-007) — risk manager fundamental-report fix (#341)
|
||||
- [@makk9](https://github.com/makk9) — debate-rounds config issue (#361)
|
||||
|
||||
## [0.2.0] — 2026-02-04
|
||||
|
||||
This is the largest release since the initial public version. The framework
|
||||
moved from single-provider to a multi-provider architecture and grew several
|
||||
production-ready surfaces.
|
||||
|
||||
### Added
|
||||
|
||||
- **Multi-provider LLM support** (OpenAI, Google, Anthropic, xAI, OpenRouter,
|
||||
Ollama) via a factory pattern, with provider-specific thinking configurations.
|
||||
- **Alpha Vantage** integration as a configurable primary data provider, with
|
||||
yfinance as a community-stability fallback.
|
||||
- **Footer statistics** in the CLI: real-time tracking of LLM calls, tool
|
||||
calls, and token usage via LangChain callbacks.
|
||||
- **Post-analysis report saving** — the framework writes per-section markdown
|
||||
files (analyst reports, debate transcripts, final decision) when a run
|
||||
completes.
|
||||
- **Announcements panel** — fetches updates from `api.tauric.ai/v1/announcements`
|
||||
for the CLI welcome screen.
|
||||
- **Tool fallbacks** so a single vendor outage does not stop the pipeline.
|
||||
|
||||
### Changed
|
||||
|
||||
- Risky / Safe risk debaters renamed to **Aggressive / Conservative** for
|
||||
consistency with the displayed agent labels.
|
||||
- Default data vendor switched to balance reliability and quota across
|
||||
community deployments.
|
||||
- Ollama and OpenRouter model lists updated; default endpoints clarified.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Analyst status tracking and message deduplication in the live display.
|
||||
- Infinite-loop guard in the agent loop; reflection and logging hardened.
|
||||
- Various data-vendor implementation bugs and tool-signature mismatches.
|
||||
|
||||
### Contributors
|
||||
|
||||
This release is the first with substantial outside contributions; many community
|
||||
PRs from late 2025 also landed here.
|
||||
|
||||
- [@luohy15](https://github.com/luohy15) — Alpha Vantage data-vendor integration (#235)
|
||||
- [@EdwardoSunny](https://github.com/EdwardoSunny) — yfinance fetching optimisations (#245)
|
||||
- [@Mirza-Samad-Ahmed-Baig](https://github.com/Mirza-Samad-Ahmed-Baig) — infinite-loop guard, reflection, and logging fixes (#89)
|
||||
- [@ZeroAct](https://github.com/ZeroAct) — saved results path support (#29)
|
||||
- [@Zhongyi-Lu](https://github.com/Zhongyi-Lu) — `.env` gitignore (#49)
|
||||
- [@csoboy](https://github.com/csoboy) — local Ollama setup (#53)
|
||||
- [@chauhang](https://github.com/chauhang) — initial Docker support attempt (#47, later reverted; the merged Docker support shipped in v0.2.4)
|
||||
|
||||
## [0.1.1] — 2025-06-07
|
||||
|
||||
### Removed
|
||||
|
||||
- Static site assets that had been bundled with v0.1.0; the public site now
|
||||
lives separately.
|
||||
|
||||
## [0.1.0] — 2025-06-05
|
||||
|
||||
### Added
|
||||
|
||||
- **Initial public release** of the TradingAgents multi-agent trading
|
||||
framework: market / sentiment / news / fundamentals analysts; bull and bear
|
||||
researchers; trader; aggressive, conservative, and neutral risk debaters;
|
||||
portfolio manager. LangGraph orchestration, yfinance data, per-agent
|
||||
BM25 memory, single-provider OpenAI integration, interactive CLI.
|
||||
|
||||
[0.2.4]: https://github.com/TauricResearch/TradingAgents/compare/v0.2.3...v0.2.4
|
||||
[0.2.3]: https://github.com/TauricResearch/TradingAgents/compare/v0.2.2...v0.2.3
|
||||
[0.2.2]: https://github.com/TauricResearch/TradingAgents/compare/v0.2.1...v0.2.2
|
||||
[0.2.1]: https://github.com/TauricResearch/TradingAgents/compare/v0.2.0...v0.2.1
|
||||
[0.2.0]: https://github.com/TauricResearch/TradingAgents/compare/v0.1.1...v0.2.0
|
||||
[0.1.1]: https://github.com/TauricResearch/TradingAgents/compare/v0.1.0...v0.1.1
|
||||
[0.1.0]: https://github.com/TauricResearch/TradingAgents/releases/tag/v0.1.0
|
||||
+28
@@ -0,0 +1,28 @@
|
||||
FROM python:3.12-slim AS builder
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
PIP_DISABLE_PIP_VERSION_CHECK=1
|
||||
|
||||
RUN python -m venv /opt/venv
|
||||
ENV PATH="/opt/venv/bin:$PATH"
|
||||
|
||||
WORKDIR /build
|
||||
COPY . .
|
||||
RUN pip install --no-cache-dir .
|
||||
|
||||
FROM python:3.12-slim
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1
|
||||
|
||||
COPY --from=builder /opt/venv /opt/venv
|
||||
ENV PATH="/opt/venv/bin:$PATH"
|
||||
|
||||
RUN useradd --create-home appuser \
|
||||
&& install -d -m 0755 -o appuser -g appuser /home/appuser/.tradingagents
|
||||
USER appuser
|
||||
WORKDIR /home/appuser/app
|
||||
|
||||
COPY --from=builder --chown=appuser:appuser /build .
|
||||
|
||||
ENTRYPOINT ["tradingagents"]
|
||||
@@ -5,19 +5,53 @@
|
||||
<div align="center" style="line-height: 1;">
|
||||
<a href="https://arxiv.org/abs/2412.20138" target="_blank"><img alt="arXiv" src="https://img.shields.io/badge/arXiv-2412.20138-B31B1B?logo=arxiv"/></a>
|
||||
<a href="https://discord.com/invite/hk9PGKShPK" target="_blank"><img alt="Discord" src="https://img.shields.io/badge/Discord-TradingResearch-7289da?logo=discord&logoColor=white&color=7289da"/></a>
|
||||
<a href="./assets/wechat.png" target="_blank"><img alt="WeChat" src="https://img.shields.io/badge/WeChat-TauricResearch-brightgreen?logo=wechat&logoColor=white"/></a>
|
||||
<a href="https://x.com/TauricResearch" target="_blank"><img alt="X Follow" src="https://img.shields.io/badge/X-TauricResearch-white?logo=x&logoColor=white"/></a>
|
||||
<br>
|
||||
<a href="https://github.com/TauricResearch/" target="_blank"><img alt="Community" src="https://img.shields.io/badge/Join_GitHub_Community-TauricResearch-14C290?logo=discourse"/></a>
|
||||
<a href="https://github.com/TauricResearch/" target="_blank"><img alt="Community" src="https://img.shields.io/badge/GitHub_Community-TauricResearch-14C290?logo=discourse"/></a>
|
||||
</div>
|
||||
<br>
|
||||
<div align="center">
|
||||
<a href="https://github.com/TauricResearch" target="_blank"><img alt="TradingAgents #1 Repository of the Day" src="https://trendshift.io/api/badge/repositories/16192" width="250" height="55"/></a>
|
||||
</div>
|
||||
<br>
|
||||
<div align="center">
|
||||
<!-- Keep these links. Translations will automatically update with the README. -->
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=de">Deutsch</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=es">Español</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=fr">français</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=ja">日本語</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=ko">한국어</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=pt">Português</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=ru">Русский</a> |
|
||||
<a href="https://www.readme-i18n.com/TauricResearch/TradingAgents?lang=zh">中文</a>
|
||||
</div>
|
||||
|
||||
---
|
||||
|
||||
# TradingAgents: Multi-Agents LLM Financial Trading Framework
|
||||
# TradingAgents: Multi-Agents LLM Financial Trading Framework
|
||||
|
||||
> 🎉 **TradingAgents** officially released! We have received numerous inquiries about the work, and we would like to express our thanks for the enthusiasm in our community.
|
||||
>
|
||||
> So we decided to fully open-source the framework. Looking forward to building impactful projects with you!
|
||||
## News
|
||||
|
||||
<!-- news:start -->
|
||||
- [2026-09] **TradingAgents v0.5.1** released with a package layout organised by what each module holds (import paths moved), optional Jev screening of social posts, GPT-6 Sol and Luna as the default models, and fixes to run isolation and SEC EDGAR statements.
|
||||
- [2026-09] **TradingAgents v0.5.0** released with point-in-time integrity across every dated path, SEC EDGAR fundamentals served as filed, backtesting over a ticker and date grid, portfolio-aware runs, and current model lineups across every provider.
|
||||
- [2026-08] **TradingAgents v0.4.0** released with look-ahead / point-in-time fixes across FRED macro, social sentiment, and the decision-log memory; clearer decision signals; working CLI checkpoint resume; Trader price grounding; and the GPT-5.6 and GLM-5.3 models.
|
||||
|
||||
Full release notes are in [CHANGELOG.md](CHANGELOG.md).
|
||||
|
||||
<details>
|
||||
<summary>Earlier news</summary>
|
||||
|
||||
- [2026-07] **TradingAgents v0.3.1** released with correctness and stability fixes: Alpha Vantage look-ahead filtering, graph-router crash-safety, graph-shape-aware checkpoint resume, working crypto sentiment sources, a configurable LLM retry budget, Bedrock API-key auth, and Claude Sonnet 5 / Fable 5 support.
|
||||
- [2026-06] **TradingAgents v0.3.0** released with a verified data-access contract, an expanded provider registry (NVIDIA, Kimi, Groq, Mistral, Bedrock, and any OpenAI-compatible endpoint), FRED and Polymarket data vendors, a current-generation model catalog, and a CI gate.
|
||||
- [2026-05] **TradingAgents v0.2.5** released with the grounded Sentiment Analyst, GPT-5.5 etc. model coverage, Qwen/GLM/MiniMax dual-region support, `TRADINGAGENTS_*` env-var configurability with API-key auto-detection, remote Ollama support, non-US alpha benchmarks, and ticker path-traversal hardening.
|
||||
- [2026-04] **TradingAgents v0.2.4** released with structured-output agents (Research Manager, Trader, Portfolio Manager), LangGraph checkpoint resume, persistent decision log, DeepSeek/Qwen/GLM/Azure provider support, Docker, and a Windows UTF-8 encoding fix.
|
||||
- [2026-03] **TradingAgents v0.2.3** released with multi-language support, GPT-5.4 family models, unified model catalog, backtesting date fidelity, and proxy support.
|
||||
- [2026-03] **TradingAgents v0.2.2** released with GPT-5.4/Gemini 3.1/Claude 4.6 model coverage, five-tier rating scale, OpenAI Responses API, Anthropic effort control, and cross-platform stability.
|
||||
- [2026-02] **TradingAgents v0.2.0** released with multi-provider LLM support (GPT-5.x, Gemini 3.x, Claude 4.x, Grok 4.x) and improved system architecture.
|
||||
- [2026-01] **Trading-R1** [Technical Report](https://arxiv.org/abs/2509.11420) released, with [Terminal](https://github.com/TauricResearch/Trading-R1) expected to land soon.
|
||||
|
||||
</details>
|
||||
<!-- news:end -->
|
||||
|
||||
<div align="center">
|
||||
|
||||
@@ -25,6 +59,10 @@
|
||||
|
||||
</div>
|
||||
|
||||
> 🎉 **TradingAgents** officially released! We have received numerous inquiries about the work, and we would like to express our thanks for the enthusiasm in our community.
|
||||
>
|
||||
> So we decided to fully open-source the framework. Looking forward to building impactful projects with you!
|
||||
|
||||
## TradingAgents Framework
|
||||
|
||||
TradingAgents is a multi-agent trading framework that mirrors the dynamics of real-world trading firms. By deploying specialized LLM-powered agents: from fundamental analysts, sentiment experts, and technical analysts, to trader, risk management team, the platform collaboratively evaluates market conditions and informs trading decisions. Moreover, these agents engage in dynamic discussions to pinpoint the optimal strategy.
|
||||
@@ -35,11 +73,11 @@ TradingAgents is a multi-agent trading framework that mirrors the dynamics of re
|
||||
|
||||
> TradingAgents framework is designed for research purposes. Trading performance may vary based on many factors, including the chosen backbone language models, model temperature, trading periods, the quality of data, and other non-deterministic factors. [It is not intended as financial, investment, or trading advice.](https://tauric.ai/disclaimer/)
|
||||
|
||||
Our framework decomposes complex trading tasks into specialized roles. This ensures the system achieves a robust, scalable approach to market analysis and decision-making.
|
||||
Our framework decomposes complex trading tasks into specialized roles.
|
||||
|
||||
### Analyst Team
|
||||
- Fundamentals Analyst: Evaluates company financials and performance metrics, identifying intrinsic values and potential red flags.
|
||||
- Sentiment Analyst: Analyzes social media and public sentiment using sentiment scoring algorithms to gauge short-term market mood.
|
||||
- Sentiment Analyst: Aggregates news headlines, StockTwits, and Reddit chatter into a single sentiment read to gauge short-term market mood.
|
||||
- News Analyst: Monitors global news and macroeconomic indicators, interpreting the impact of events on market conditions.
|
||||
- Technical Analyst: Utilizes technical indicators (like MACD and RSI) to detect trading patterns and forecast price movements.
|
||||
|
||||
@@ -55,10 +93,10 @@ Our framework decomposes complex trading tasks into specialized roles. This ensu
|
||||
</p>
|
||||
|
||||
### Trader Agent
|
||||
- Composes reports from the analysts and researchers to make informed trading decisions. It determines the timing and magnitude of trades based on comprehensive market insights.
|
||||
- Composes reports from the analysts and researchers to make informed trading decisions, determining the timing and magnitude of trades.
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/risk.png" width="70%" style="display: inline-block; margin: 0 2%;">
|
||||
<img src="assets/trader.png" width="70%" style="display: inline-block; margin: 0 2%;">
|
||||
</p>
|
||||
|
||||
### Risk Management and Portfolio Manager
|
||||
@@ -66,7 +104,7 @@ Our framework decomposes complex trading tasks into specialized roles. This ensu
|
||||
- The Portfolio Manager approves/rejects the transaction proposal. If approved, the order will be sent to the simulated exchange and executed.
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/trader.png" width="70%" style="display: inline-block; margin: 0 2%;">
|
||||
<img src="assets/risk.png" width="70%" style="display: inline-block; margin: 0 2%;">
|
||||
</p>
|
||||
|
||||
## Installation and CLI
|
||||
@@ -81,34 +119,95 @@ cd TradingAgents
|
||||
|
||||
Create a virtual environment in any of your favorite environment managers:
|
||||
```bash
|
||||
conda create -n tradingagents python=3.13
|
||||
conda create -n tradingagents python=3.12
|
||||
conda activate tradingagents
|
||||
```
|
||||
|
||||
Install dependencies:
|
||||
Or with [uv](https://docs.astral.sh/uv/):
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
uv venv --python 3.12
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
Install the package and its dependencies (`uv pip install .` with uv):
|
||||
```bash
|
||||
pip install .
|
||||
```
|
||||
|
||||
### Docker
|
||||
|
||||
Alternatively, run with Docker:
|
||||
```bash
|
||||
cp .env.example .env # add your API keys
|
||||
docker compose run --rm tradingagents
|
||||
```
|
||||
|
||||
After updating the repository, rebuild the image with `docker compose build`.
|
||||
|
||||
For local models with Ollama:
|
||||
```bash
|
||||
docker compose --profile ollama run --rm tradingagents-ollama
|
||||
```
|
||||
|
||||
### Required APIs
|
||||
|
||||
You will also need the FinnHub API and EODHD API for financial data. All of our code is implemented with the free tier.
|
||||
TradingAgents supports multiple LLM providers. Set the API key for your chosen provider:
|
||||
|
||||
```bash
|
||||
export FINNHUB_API_KEY=$YOUR_FINNHUB_API_KEY
|
||||
export OPENAI_API_KEY=... # OpenAI (GPT)
|
||||
export GOOGLE_API_KEY=... # Google (Gemini)
|
||||
export ANTHROPIC_API_KEY=... # Anthropic (Claude)
|
||||
export XAI_API_KEY=... # xAI (Grok)
|
||||
export DEEPSEEK_API_KEY=... # DeepSeek
|
||||
export DASHSCOPE_API_KEY=... # Qwen — International (dashscope-intl.aliyuncs.com)
|
||||
export DASHSCOPE_CN_API_KEY=... # Qwen — China (dashscope.aliyuncs.com)
|
||||
export ZHIPU_API_KEY=... # GLM via Z.AI (international)
|
||||
export ZHIPU_CN_API_KEY=... # GLM via BigModel (China, open.bigmodel.cn)
|
||||
export MINIMAX_API_KEY=... # MiniMax — Global (api.minimax.io)
|
||||
export MINIMAX_CN_API_KEY=... # MiniMax — China (api.minimaxi.com)
|
||||
export OPENROUTER_API_KEY=... # OpenRouter
|
||||
export MISTRAL_API_KEY=... # Mistral
|
||||
export MOONSHOT_API_KEY=... # Kimi (Moonshot)
|
||||
export GROQ_API_KEY=... # Groq
|
||||
export NVIDIA_API_KEY=... # NVIDIA NIM
|
||||
export FRED_API_KEY=... # FRED macro data (free, optional)
|
||||
export ALPHA_VANTAGE_API_KEY=... # Alpha Vantage
|
||||
export TYPESAFE_API_KEY=... # Jev social-post screening (optional)
|
||||
```
|
||||
|
||||
You will need the OpenAI API for all the agents.
|
||||
For Azure OpenAI, copy `.env.enterprise.example` to `.env.enterprise` and fill in your credentials.
|
||||
|
||||
For AWS Bedrock, install the extra with `pip install ".[bedrock]"`, set `llm_provider: "bedrock"`, configure AWS credentials (environment variables, `~/.aws/credentials`, or an IAM role) and `AWS_DEFAULT_REGION`, and use a Bedrock model ID, e.g. `us.anthropic.claude-opus-4-8-v1:0`.
|
||||
|
||||
For local models, configure Ollama with `llm_provider: "ollama"`. The default endpoint is `http://localhost:11434/v1`; set `OLLAMA_BASE_URL` to point at a remote `ollama-serve`. Pull models with `ollama pull <name>`, and pick "Custom model ID" in the CLI for any model not listed by default.
|
||||
|
||||
For any other OpenAI-compatible server (vLLM, LM Studio, llama.cpp, or a custom relay), use `llm_provider: "openai_compatible"` and set the endpoint via `backend_url` (or `TRADINGAGENTS_LLM_BACKEND_URL`), e.g. `http://localhost:8000/v1` for vLLM or `http://localhost:1234/v1` for LM Studio. The model is whatever your server serves. No key is needed for local servers; set `OPENAI_COMPATIBLE_API_KEY` when the endpoint requires one.
|
||||
|
||||
With `TYPESAFE_API_KEY` set, the Sentiment Analyst screens StockTwits and Reddit posts with TypeSafe's Jev before reading them. Posts that are not about the company are dropped, and each source opens with a count of the remaining posts by stance: bullish, bearish, neutral, or unclear. Without the key, posts pass through unscreened. `jev-latest` moves with new releases; set `TYPESAFE_DEFAULT_MODEL` to a versioned ID such as `jev-1.13.0` to hold it fixed across runs.
|
||||
|
||||
Alternatively, copy `.env.example` to `.env` and fill in your keys:
|
||||
```bash
|
||||
export OPENAI_API_KEY=$YOUR_OPENAI_API_KEY
|
||||
cp .env.example .env
|
||||
```
|
||||
|
||||
### CLI Usage
|
||||
|
||||
You can also try out the CLI directly by running:
|
||||
Launch the interactive CLI:
|
||||
```bash
|
||||
python -m cli.main
|
||||
tradingagents # installed command
|
||||
python -m cli.main # alternative: run directly from source
|
||||
```
|
||||
You will see a screen where you can select your desired tickers, date, LLMs, research depth, etc.
|
||||
You will see a screen where you can select your desired tickers, analysis date, LLM provider, research depth, and more. Your previous run's answers come back as the defaults, so pressing Enter accepts them. The `TRADINGAGENTS_*` variables in `.env` still skip their step entirely.
|
||||
|
||||
### Markets and tickers
|
||||
|
||||
TradingAgents works with any market Yahoo Finance covers, using the exchange-suffixed ticker. Company identity and the alpha benchmark resolve automatically per market.
|
||||
|
||||
- US: `AAPL`, `SPY`
|
||||
- Hong Kong: `0700.HK` · Tokyo: `7203.T` · London: `AZN.L`
|
||||
- India: `RELIANCE.NS`, `.BO` · Canada: `.TO` · Australia: `.AX`
|
||||
- China A-shares: Shanghai `.SS`, Shenzhen `.SZ` (e.g. `600519.SS` for Kweichow Moutai)
|
||||
- Crypto: `BTC-USD`, `ETH-USD`
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/cli/cli_init.png" width="100%" style="display: inline-block; margin: 0 2%;">
|
||||
@@ -128,7 +227,7 @@ An interface will appear showing results as they load, letting you track the age
|
||||
|
||||
### Implementation Details
|
||||
|
||||
We built TradingAgents with LangGraph to ensure flexibility and modularity. We utilize `o1-preview` and `gpt-4o` as our deep thinking and fast thinking LLMs for our experiments. However, for testing purposes, we recommend you use `o4-mini` and `gpt-4.1-mini` to save on costs as our framework makes **lots of** API calls.
|
||||
We built TradingAgents with LangGraph to ensure flexibility and modularity. The framework supports multiple LLM providers: OpenAI, Google, Anthropic, xAI, DeepSeek, Qwen (Alibaba DashScope, international and China endpoints), GLM (Zhipu), MiniMax (global + China), OpenRouter, Ollama for local models, and Azure OpenAI for enterprise.
|
||||
|
||||
### Python Usage
|
||||
|
||||
@@ -136,11 +235,12 @@ To use TradingAgents inside your code, you can import the `tradingagents` module
|
||||
|
||||
```python
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
|
||||
ta = TradingAgentsGraph(debug=True, config=config)
|
||||
ta = TradingAgentsGraph(debug=True, config=DEFAULT_CONFIG.copy())
|
||||
|
||||
# forward propagate
|
||||
_, decision = ta.propagate("NVDA", "2024-05-10")
|
||||
_, decision = ta.propagate("NVDA", "2026-09-01")
|
||||
print(decision)
|
||||
```
|
||||
|
||||
@@ -150,28 +250,129 @@ You can also adjust the default configuration to set your own choice of LLMs, de
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
|
||||
# Create a custom config
|
||||
config = DEFAULT_CONFIG.copy()
|
||||
config["deep_think_llm"] = "gpt-4.1-nano" # Use a different model
|
||||
config["quick_think_llm"] = "gpt-4.1-nano" # Use a different model
|
||||
config["max_debate_rounds"] = 1 # Increase debate rounds
|
||||
config["online_tools"] = True # Use online tools or cached data
|
||||
config["llm_provider"] = "openai" # e.g. openai, google, anthropic, deepseek, groq, ollama; openai_compatible covers any OpenAI-compatible endpoint (vLLM, LM Studio, llama.cpp, ...)
|
||||
config["deep_think_llm"] = "gpt-6-sol" # Model for complex reasoning
|
||||
config["quick_think_llm"] = "gpt-6-luna" # Model for quick tasks
|
||||
config["max_debate_rounds"] = 2
|
||||
|
||||
# Initialize with custom config
|
||||
ta = TradingAgentsGraph(debug=True, config=config)
|
||||
|
||||
# forward propagate
|
||||
_, decision = ta.propagate("NVDA", "2024-05-10")
|
||||
_, decision = ta.propagate("NVDA", "2026-09-01")
|
||||
print(decision)
|
||||
```
|
||||
|
||||
> For `online_tools`, we recommend enabling them for experimentation, as they provide access to real-time data. The agents' offline tools rely on cached data from our **Tauric TradingDB**, a curated dataset we use for backtesting. We're currently in the process of refining this dataset, and we plan to release it soon alongside our upcoming projects. Stay tuned!
|
||||
See `tradingagents/default_config.py` for all configuration options.
|
||||
|
||||
You can view the full list of configurations in `tradingagents/default_config.py`.
|
||||
### Fundamentals as filed
|
||||
|
||||
US company statements can come from SEC EDGAR, which records the date every figure was filed. A run dated in the past then reads the statements exactly as they stood that day: a fiscal year that has ended but has not been filed yet is not served, and a figure restated later still reads as first reported. Apple's 2008 total assets were filed as $39.6B and restated to $36.2B in 2010, so a run dated in between reads $39.6B.
|
||||
|
||||
EDGAR needs no account or API key. Add the vendor to the chain:
|
||||
|
||||
```python
|
||||
config["data_vendors"]["fundamental_data"] = "sec_edgar,yfinance"
|
||||
```
|
||||
|
||||
SEC asks callers to identify themselves and refuses requests that carry no contact address, so a default one is sent. Set your own so SEC can reach you rather than the project:
|
||||
|
||||
```bash
|
||||
SEC_EDGAR_USER_AGENT="Your Name your@email.com"
|
||||
```
|
||||
|
||||
It covers companies that file with the SEC, including foreign companies listed in the US. Anything else, such as Hong Kong or A-share listings, falls through to the next vendor in the chain. EDGAR's machine-readable filings begin in 2009, and a fourth quarter is reported as unavailable rather than derived, because filers publish it only inside the annual figure.
|
||||
|
||||
### Current holdings
|
||||
|
||||
By default the agents do not know what you hold, so their guidance is written for a reader who applies it to their own position. Pass a portfolio to have the trader, the risk analysts and the portfolio manager work against your actual book.
|
||||
|
||||
```python
|
||||
from tradingagents.portfolio import PortfolioContext
|
||||
|
||||
portfolio = PortfolioContext.model_validate({
|
||||
"cash": 25000.0,
|
||||
"currency": "USD",
|
||||
"positions": [{"ticker": "NVDA", "quantity": 120, "average_price": 150.0}],
|
||||
})
|
||||
_, decision = ta.propagate("NVDA", "2026-09-01", portfolio=portfolio)
|
||||
```
|
||||
|
||||
The CLI takes the same content as a JSON file: `tradingagents --portfolio my_book.json`.
|
||||
|
||||
An empty `positions` list means a flat book, which is different from passing nothing. A run without a portfolio is never treated as flat.
|
||||
|
||||
## Persistence and Recovery
|
||||
|
||||
TradingAgents persists two kinds of state across runs.
|
||||
|
||||
### Decision log
|
||||
|
||||
The decision log is always on. Each completed run appends its decision to `~/.tradingagents/memory/trading_memory.md`. On the next run for the same ticker, TradingAgents fetches the realised return (raw, and alpha against the instrument's regional benchmark), generates a one-paragraph reflection, and injects the most recent same-ticker decisions plus recent cross-ticker lessons into the Portfolio Manager prompt, so each analysis carries forward what worked and what didn't.
|
||||
|
||||
Override the path with `TRADINGAGENTS_MEMORY_LOG_PATH`.
|
||||
|
||||
### Checkpoint resume
|
||||
|
||||
Checkpoint resume is opt-in via `--checkpoint`. When enabled, LangGraph saves state after each node so a crashed or interrupted run resumes from the last successful step instead of starting over. The run view says whether it resumed a saved run or started fresh. Checkpoints are cleared automatically on successful completion.
|
||||
|
||||
Per-ticker SQLite databases live at `~/.tradingagents/cache/checkpoints/<TICKER>.db` (override the base with `TRADINGAGENTS_CACHE_DIR`). Use `--clear-checkpoints` to reset all of them before a run.
|
||||
|
||||
```bash
|
||||
tradingagents --checkpoint # enable for this run
|
||||
tradingagents --clear-checkpoints # reset before running
|
||||
```
|
||||
|
||||
```python
|
||||
config = DEFAULT_CONFIG.copy()
|
||||
config["checkpoint_enabled"] = True
|
||||
ta = TradingAgentsGraph(config=config)
|
||||
_, decision = ta.propagate("NVDA", "2026-09-01")
|
||||
```
|
||||
|
||||
## Evaluating decisions over time
|
||||
|
||||
One run gives one decision, which cannot tell you whether the system decides well. `run_backtest` runs the same pipeline over a grid of tickers and dates, writes to a decision log of its own, and scores the decisions whose holding window has since traded.
|
||||
|
||||
```python
|
||||
from tradingagents.backtest import iter_grid, run_backtest, summarize
|
||||
|
||||
dates = iter_grid("2026-06-01", "2026-08-01", every_n_days=7)
|
||||
result = run_backtest(["NVDA", "AAPL"], dates, config, selected_analysts=["market", "news"])
|
||||
print(summarize(result).render())
|
||||
```
|
||||
|
||||
From the CLI:
|
||||
|
||||
```bash
|
||||
tradingagents backtest NVDA,AAPL --start 2026-06-01 --end 2026-08-01 --every 7
|
||||
```
|
||||
|
||||
Each cell is scored on realized alpha against the instrument's regional benchmark, grouped by rating. Your own decision log is never written to, and re-running the same grid with `run_id=result.run_id` skips the cells that already ran, so an interrupted sweep continues where it stopped.
|
||||
|
||||
## Reproducibility
|
||||
|
||||
TradingAgents is LLM-driven, so two runs of the same ticker and date can differ. This is expected for a research tool built on language models, not a defect. The variation comes from a few distinct sources, and it helps to separate them.
|
||||
|
||||
Language model sampling is non-deterministic. Even at a fixed temperature, providers do not guarantee byte-identical output across calls, and reasoning models (the default GPT-6 family, and any thinking-mode model) vary the most because their internal reasoning is itself sampled.
|
||||
|
||||
Live data moves. News, StockTwits, and Reddit return different content as time passes, so a run today sees different inputs than a run last week even for the same historical trade date. Pin the analysis date to hold the price and indicator window fixed, but the social and news sources still reflect "now".
|
||||
|
||||
To reduce variation you can lower the sampling temperature. Set `temperature` in your config (or `TRADINGAGENTS_TEMPERATURE` in `.env`); lower values make models that honor it more repeatable. The current curated models are reasoning-first and largely ignore temperature, so for tighter reproducibility name a non-reasoning model in your config, or in `TRADINGAGENTS_DEEP_THINK_LLM` and `TRADINGAGENTS_QUICK_THINK_LLM`. Any model ID your provider serves is accepted, whether or not the picker lists it.
|
||||
|
||||
```python
|
||||
config = DEFAULT_CONFIG.copy()
|
||||
config["llm_provider"] = "openai"
|
||||
config["temperature"] = 0.0
|
||||
# Reasoning models ignore temperature. For tighter reproducibility, name a
|
||||
# non-reasoning model in deep_think_llm / quick_think_llm.
|
||||
```
|
||||
|
||||
What does not vary anymore: the analyzed company identity is resolved deterministically from the ticker before any agent runs, and the market analyst grounds exact price and indicator claims in a verified data snapshot. Earlier reports of "different companies" or fabricated price levels across runs are addressed by these two mechanisms.
|
||||
|
||||
Backtest results are not guaranteed to match any published figure. Returns depend on the model, the temperature, the date range, data quality, and the sampling above. Treat the framework as a research scaffold for studying multi-agent analysis, not as a strategy with a fixed, replicable return.
|
||||
|
||||
## Contributing
|
||||
|
||||
We welcome contributions from the community! Whether it's fixing a bug, improving documentation, or suggesting a new feature, your input helps make this project better. If you are interested in this line of research, please consider joining our open-source financial AI research community [Tauric Research](https://tauric.ai/).
|
||||
Contributions are welcome: bug fixes, documentation, and feature ideas; past contributions are credited per release in [`CHANGELOG.md`](CHANGELOG.md).
|
||||
|
||||
## Citation
|
||||
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 216 KiB |
@@ -0,0 +1,52 @@
|
||||
import getpass
|
||||
|
||||
import requests
|
||||
from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
|
||||
from cli.config import CLI_CONFIG
|
||||
|
||||
|
||||
def fetch_announcements(url: str = None, timeout: float = None) -> dict:
|
||||
"""Fetch announcements from endpoint. Returns dict with announcements and settings."""
|
||||
endpoint = url or CLI_CONFIG["announcements_url"]
|
||||
timeout = timeout or CLI_CONFIG["announcements_timeout"]
|
||||
fallback = CLI_CONFIG["announcements_fallback"]
|
||||
|
||||
try:
|
||||
response = requests.get(endpoint, timeout=timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return {
|
||||
"announcements": data.get("announcements", [fallback]),
|
||||
"require_attention": data.get("require_attention", False),
|
||||
}
|
||||
except Exception:
|
||||
return {
|
||||
"announcements": [fallback],
|
||||
"require_attention": False,
|
||||
}
|
||||
|
||||
|
||||
def display_announcements(console: Console, data: dict) -> None:
|
||||
"""Display announcements panel. Prompts for Enter if require_attention is True."""
|
||||
announcements = data.get("announcements", [])
|
||||
require_attention = data.get("require_attention", False)
|
||||
|
||||
if not announcements:
|
||||
return
|
||||
|
||||
content = "\n".join(announcements)
|
||||
|
||||
panel = Panel(
|
||||
content,
|
||||
border_style="cyan",
|
||||
padding=(1, 2),
|
||||
title="Announcements",
|
||||
)
|
||||
console.print(panel)
|
||||
|
||||
if require_attention:
|
||||
getpass.getpass("Press Enter to continue...")
|
||||
else:
|
||||
console.print()
|
||||
@@ -0,0 +1,6 @@
|
||||
CLI_CONFIG = {
|
||||
# Announcements
|
||||
"announcements_url": "https://api.tauric.ai/v1/announcements",
|
||||
"announcements_timeout": 1.0,
|
||||
"announcements_fallback": "[cyan]For more information, please visit[/cyan] [link=https://github.com/TauricResearch]https://github.com/TauricResearch[/link]",
|
||||
}
|
||||
+640
@@ -0,0 +1,640 @@
|
||||
"""The live view of a run: message log, agent status, report panels and timings."""
|
||||
|
||||
import datetime
|
||||
import time
|
||||
from collections import deque
|
||||
from time import monotonic
|
||||
|
||||
from rich import box
|
||||
from rich.console import Console
|
||||
from rich.layout import Layout
|
||||
from rich.markdown import Markdown
|
||||
from rich.panel import Panel
|
||||
from rich.rule import Rule
|
||||
from rich.spinner import Spinner
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
from tradingagents.graph.analyst_execution import (
|
||||
ANALYST_NODE_SPECS,
|
||||
AnalystExecutionPlan,
|
||||
)
|
||||
|
||||
console = Console()
|
||||
|
||||
|
||||
class MessageBuffer:
|
||||
# Fixed teams that always run (not user-selectable)
|
||||
FIXED_AGENTS = {
|
||||
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
|
||||
"Trading Team": ["Trader"],
|
||||
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
|
||||
"Portfolio Management": ["Portfolio Manager"],
|
||||
}
|
||||
|
||||
# Analyst name mapping
|
||||
ANALYST_MAPPING = {
|
||||
"market": "Market Analyst",
|
||||
"social": "Sentiment Analyst",
|
||||
"news": "News Analyst",
|
||||
"fundamentals": "Fundamentals Analyst",
|
||||
}
|
||||
|
||||
# Report section mapping: section -> (analyst_key for filtering, finalizing_agent)
|
||||
# analyst_key: which analyst selection controls this section (None = always included)
|
||||
# finalizing_agent: which agent must be "completed" for this report to count as done
|
||||
REPORT_SECTIONS = {
|
||||
"market_report": ("market", "Market Analyst"),
|
||||
"sentiment_report": ("social", "Sentiment Analyst"),
|
||||
"news_report": ("news", "News Analyst"),
|
||||
"fundamentals_report": ("fundamentals", "Fundamentals Analyst"),
|
||||
"investment_plan": (None, "Research Manager"),
|
||||
"trader_investment_plan": (None, "Trader"),
|
||||
"final_trade_decision": (None, "Portfolio Manager"),
|
||||
}
|
||||
|
||||
def __init__(self, max_length=100):
|
||||
self.messages = deque(maxlen=max_length)
|
||||
self.tool_calls = deque(maxlen=max_length)
|
||||
self.current_report = None
|
||||
self.agent_status = {}
|
||||
self.report_sections = {}
|
||||
self.selected_analysts = []
|
||||
self._processed_message_ids = set()
|
||||
|
||||
def init_for_analysis(self, selected_analysts):
|
||||
"""Initialize agent status and report sections based on selected analysts.
|
||||
|
||||
Args:
|
||||
selected_analysts: List of analyst type strings (e.g., ["market", "news"])
|
||||
"""
|
||||
self.selected_analysts = [a.lower() for a in selected_analysts]
|
||||
|
||||
self.agent_status = {}
|
||||
|
||||
for analyst_key in self.selected_analysts:
|
||||
if analyst_key in self.ANALYST_MAPPING:
|
||||
self.agent_status[self.ANALYST_MAPPING[analyst_key]] = "pending"
|
||||
|
||||
for team_agents in self.FIXED_AGENTS.values():
|
||||
for agent in team_agents:
|
||||
self.agent_status[agent] = "pending"
|
||||
|
||||
self.report_sections = {}
|
||||
for section, (analyst_key, _) in self.REPORT_SECTIONS.items():
|
||||
if analyst_key is None or analyst_key in self.selected_analysts:
|
||||
self.report_sections[section] = None
|
||||
|
||||
# Reset other state
|
||||
self.current_report = None
|
||||
self.messages.clear()
|
||||
self.tool_calls.clear()
|
||||
self._processed_message_ids.clear()
|
||||
|
||||
def get_completed_reports_count(self):
|
||||
"""Count reports that are finalized (their finalizing agent is completed).
|
||||
|
||||
A report is considered complete when:
|
||||
1. The report section has content (not None), AND
|
||||
2. The agent responsible for finalizing that report has status "completed"
|
||||
|
||||
This prevents interim updates (like debate rounds) from counting as completed.
|
||||
"""
|
||||
count = 0
|
||||
for section in self.report_sections:
|
||||
if section not in self.REPORT_SECTIONS:
|
||||
continue
|
||||
_, finalizing_agent = self.REPORT_SECTIONS[section]
|
||||
# Report is complete if it has content AND its finalizing agent is done
|
||||
has_content = self.report_sections.get(section) is not None
|
||||
agent_done = self.agent_status.get(finalizing_agent) == "completed"
|
||||
if has_content and agent_done:
|
||||
count += 1
|
||||
return count
|
||||
|
||||
def add_message(self, message_type, content):
|
||||
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
|
||||
self.messages.append((timestamp, message_type, content))
|
||||
|
||||
def add_tool_call(self, tool_name, args):
|
||||
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
|
||||
self.tool_calls.append((timestamp, tool_name, args))
|
||||
|
||||
def update_agent_status(self, agent, status):
|
||||
if agent in self.agent_status:
|
||||
self.agent_status[agent] = status
|
||||
|
||||
def update_report_section(self, section_name, content):
|
||||
if section_name in self.report_sections:
|
||||
self.report_sections[section_name] = content
|
||||
self._update_current_report()
|
||||
|
||||
def _update_current_report(self):
|
||||
# For the panel display, only show the most recently updated section
|
||||
latest_section = None
|
||||
latest_content = None
|
||||
|
||||
for section, content in self.report_sections.items():
|
||||
if content is not None:
|
||||
latest_section = section
|
||||
latest_content = content
|
||||
|
||||
if latest_section and latest_content:
|
||||
section_titles = {
|
||||
"market_report": "Market Analysis",
|
||||
"sentiment_report": "Social Sentiment",
|
||||
"news_report": "News Analysis",
|
||||
"fundamentals_report": "Fundamentals Analysis",
|
||||
"investment_plan": "Research Team Decision",
|
||||
"trader_investment_plan": "Trading Team Plan",
|
||||
"final_trade_decision": "Portfolio Management Decision",
|
||||
}
|
||||
self.current_report = (
|
||||
f"### {section_titles[latest_section]}\n{latest_content}"
|
||||
)
|
||||
|
||||
|
||||
message_buffer = MessageBuffer()
|
||||
|
||||
|
||||
def create_layout():
|
||||
layout = Layout()
|
||||
layout.split_column(
|
||||
Layout(name="header", size=3),
|
||||
Layout(name="main"),
|
||||
Layout(name="footer", size=3),
|
||||
)
|
||||
layout["main"].split_column(
|
||||
Layout(name="upper", ratio=3), Layout(name="analysis", ratio=5)
|
||||
)
|
||||
layout["upper"].split_row(
|
||||
Layout(name="progress", ratio=2), Layout(name="messages", ratio=3)
|
||||
)
|
||||
return layout
|
||||
|
||||
|
||||
def format_tokens(n):
|
||||
"""Format token count for display."""
|
||||
if n >= 1000:
|
||||
return f"{n/1000:.1f}k"
|
||||
return str(n)
|
||||
|
||||
|
||||
def update_display(layout, spinner_text=None, stats_handler=None, start_time=None):
|
||||
# Header with welcome message
|
||||
layout["header"].update(
|
||||
Panel(
|
||||
"[bold green]Welcome to TradingAgents CLI[/bold green]\n"
|
||||
"[dim]© [Tauric Research](https://github.com/TauricResearch)[/dim]",
|
||||
title="Welcome to TradingAgents",
|
||||
border_style="green",
|
||||
padding=(1, 2),
|
||||
expand=True,
|
||||
)
|
||||
)
|
||||
|
||||
# Progress panel showing agent status
|
||||
progress_table = Table(
|
||||
show_header=True,
|
||||
header_style="bold magenta",
|
||||
show_footer=False,
|
||||
box=box.SIMPLE_HEAD, # Use simple header with horizontal lines
|
||||
title=None, # Remove the redundant Progress title
|
||||
padding=(0, 2), # Add horizontal padding
|
||||
expand=True, # Make table expand to fill available space
|
||||
)
|
||||
progress_table.add_column("Team", style="cyan", justify="center", width=20)
|
||||
progress_table.add_column("Agent", style="green", justify="center", width=20)
|
||||
progress_table.add_column("Status", style="yellow", justify="center", width=20)
|
||||
|
||||
# Group agents by team - filter to only include agents in agent_status
|
||||
all_teams = {
|
||||
"Analyst Team": [
|
||||
"Market Analyst",
|
||||
"Sentiment Analyst",
|
||||
"News Analyst",
|
||||
"Fundamentals Analyst",
|
||||
],
|
||||
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
|
||||
"Trading Team": ["Trader"],
|
||||
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
|
||||
"Portfolio Management": ["Portfolio Manager"],
|
||||
}
|
||||
|
||||
teams = {}
|
||||
for team, agents in all_teams.items():
|
||||
active_agents = [a for a in agents if a in message_buffer.agent_status]
|
||||
if active_agents:
|
||||
teams[team] = active_agents
|
||||
|
||||
for team, agents in teams.items():
|
||||
first_agent = agents[0]
|
||||
status = message_buffer.agent_status.get(first_agent, "pending")
|
||||
if status == "in_progress":
|
||||
spinner = Spinner(
|
||||
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
|
||||
)
|
||||
status_cell = spinner
|
||||
else:
|
||||
status_color = {
|
||||
"pending": "yellow",
|
||||
"completed": "green",
|
||||
"error": "red",
|
||||
}.get(status, "white")
|
||||
status_cell = f"[{status_color}]{status}[/{status_color}]"
|
||||
progress_table.add_row(team, first_agent, status_cell)
|
||||
|
||||
for agent in agents[1:]:
|
||||
status = message_buffer.agent_status.get(agent, "pending")
|
||||
if status == "in_progress":
|
||||
spinner = Spinner(
|
||||
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
|
||||
)
|
||||
status_cell = spinner
|
||||
else:
|
||||
status_color = {
|
||||
"pending": "yellow",
|
||||
"completed": "green",
|
||||
"error": "red",
|
||||
}.get(status, "white")
|
||||
status_cell = f"[{status_color}]{status}[/{status_color}]"
|
||||
progress_table.add_row("", agent, status_cell)
|
||||
|
||||
progress_table.add_row("─" * 20, "─" * 20, "─" * 20, style="dim")
|
||||
|
||||
layout["progress"].update(
|
||||
Panel(progress_table, title="Progress", border_style="cyan", padding=(1, 2))
|
||||
)
|
||||
|
||||
# Messages panel showing recent messages and tool calls
|
||||
messages_table = Table(
|
||||
show_header=True,
|
||||
header_style="bold magenta",
|
||||
show_footer=False,
|
||||
expand=True, # Make table expand to fill available space
|
||||
box=box.MINIMAL, # Use minimal box style for a lighter look
|
||||
show_lines=True, # Keep horizontal lines
|
||||
padding=(0, 1), # Add some padding between columns
|
||||
)
|
||||
messages_table.add_column("Time", style="cyan", width=8, justify="center")
|
||||
messages_table.add_column("Type", style="green", width=10, justify="center")
|
||||
messages_table.add_column(
|
||||
"Content", style="white", no_wrap=False, ratio=1
|
||||
) # Make content column expand
|
||||
|
||||
# Combine tool calls and messages
|
||||
all_messages = []
|
||||
|
||||
for timestamp, tool_name, args in message_buffer.tool_calls:
|
||||
formatted_args = format_tool_args(args)
|
||||
all_messages.append((timestamp, "Tool", f"{tool_name}: {formatted_args}"))
|
||||
|
||||
for timestamp, msg_type, content in message_buffer.messages:
|
||||
content_str = str(content) if content else ""
|
||||
if len(content_str) > 200:
|
||||
content_str = content_str[:197] + "..."
|
||||
all_messages.append((timestamp, msg_type, content_str))
|
||||
|
||||
# Sort by timestamp descending (newest first)
|
||||
all_messages.sort(key=lambda x: x[0], reverse=True)
|
||||
|
||||
max_messages = 12
|
||||
|
||||
recent_messages = all_messages[:max_messages]
|
||||
|
||||
for timestamp, msg_type, content in recent_messages:
|
||||
wrapped_content = Text(content, overflow="fold")
|
||||
messages_table.add_row(timestamp, msg_type, wrapped_content)
|
||||
|
||||
layout["messages"].update(
|
||||
Panel(
|
||||
messages_table,
|
||||
title="Messages & Tools",
|
||||
border_style="blue",
|
||||
padding=(1, 2),
|
||||
)
|
||||
)
|
||||
|
||||
# Analysis panel showing current report
|
||||
if message_buffer.current_report:
|
||||
layout["analysis"].update(
|
||||
Panel(
|
||||
Markdown(message_buffer.current_report),
|
||||
title="Current Report",
|
||||
border_style="green",
|
||||
padding=(1, 2),
|
||||
)
|
||||
)
|
||||
else:
|
||||
layout["analysis"].update(
|
||||
Panel(
|
||||
"[italic]Waiting for analysis report...[/italic]",
|
||||
title="Current Report",
|
||||
border_style="green",
|
||||
padding=(1, 2),
|
||||
)
|
||||
)
|
||||
|
||||
# Footer with statistics
|
||||
# Agent progress - derived from agent_status dict
|
||||
agents_completed = sum(
|
||||
1 for status in message_buffer.agent_status.values() if status == "completed"
|
||||
)
|
||||
agents_total = len(message_buffer.agent_status)
|
||||
|
||||
# Report progress - based on agent completion (not just content existence)
|
||||
reports_completed = message_buffer.get_completed_reports_count()
|
||||
reports_total = len(message_buffer.report_sections)
|
||||
|
||||
stats_parts = [f"Agents: {agents_completed}/{agents_total}"]
|
||||
|
||||
# LLM and tool stats from callback handler
|
||||
if stats_handler:
|
||||
stats = stats_handler.get_stats()
|
||||
stats_parts.append(f"LLM: {stats['llm_calls']}")
|
||||
stats_parts.append(f"Tools: {stats['tool_calls']}")
|
||||
|
||||
# Token display with graceful fallback
|
||||
if stats["tokens_in"] > 0 or stats["tokens_out"] > 0:
|
||||
tokens_str = f"Tokens: {format_tokens(stats['tokens_in'])}\u2191 {format_tokens(stats['tokens_out'])}\u2193"
|
||||
else:
|
||||
tokens_str = "Tokens: --"
|
||||
stats_parts.append(tokens_str)
|
||||
|
||||
stats_parts.append(f"Reports: {reports_completed}/{reports_total}")
|
||||
|
||||
# Elapsed time
|
||||
if start_time:
|
||||
elapsed = time.time() - start_time
|
||||
elapsed_str = f"\u23f1 {int(elapsed // 60):02d}:{int(elapsed % 60):02d}"
|
||||
stats_parts.append(elapsed_str)
|
||||
|
||||
stats_table = Table(show_header=False, box=None, padding=(0, 2), expand=True)
|
||||
stats_table.add_column("Stats", justify="center")
|
||||
stats_table.add_row(" | ".join(stats_parts))
|
||||
|
||||
layout["footer"].update(Panel(stats_table, border_style="grey50"))
|
||||
|
||||
|
||||
def display_complete_report(final_state):
|
||||
"""Display the complete analysis report sequentially (avoids truncation)."""
|
||||
console.print()
|
||||
console.print(Rule("Complete Analysis Report", style="bold green"))
|
||||
|
||||
# I. Analyst Team Reports
|
||||
analysts = []
|
||||
if final_state.get("market_report"):
|
||||
analysts.append(("Market Analyst", final_state["market_report"]))
|
||||
if final_state.get("sentiment_report"):
|
||||
analysts.append(("Sentiment Analyst", final_state["sentiment_report"]))
|
||||
if final_state.get("news_report"):
|
||||
analysts.append(("News Analyst", final_state["news_report"]))
|
||||
if final_state.get("fundamentals_report"):
|
||||
analysts.append(("Fundamentals Analyst", final_state["fundamentals_report"]))
|
||||
if analysts:
|
||||
console.print(Panel("[bold]I. Analyst Team Reports[/bold]", border_style="cyan"))
|
||||
for title, content in analysts:
|
||||
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
|
||||
|
||||
# II. Research Team Reports
|
||||
if final_state.get("investment_debate_state"):
|
||||
debate = final_state["investment_debate_state"]
|
||||
research = []
|
||||
if debate.get("bull_history"):
|
||||
research.append(("Bull Researcher", debate["bull_history"]))
|
||||
if debate.get("bear_history"):
|
||||
research.append(("Bear Researcher", debate["bear_history"]))
|
||||
if debate.get("judge_decision"):
|
||||
research.append(("Research Manager", debate["judge_decision"]))
|
||||
if research:
|
||||
console.print(Panel("[bold]II. Research Team Decision[/bold]", border_style="magenta"))
|
||||
for title, content in research:
|
||||
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
|
||||
|
||||
# III. Trading Team
|
||||
if final_state.get("trader_investment_plan"):
|
||||
console.print(Panel("[bold]III. Trading Team Plan[/bold]", border_style="yellow"))
|
||||
console.print(Panel(Markdown(final_state["trader_investment_plan"]), title="Trader", border_style="blue", padding=(1, 2)))
|
||||
|
||||
# IV. Risk Management Team
|
||||
if final_state.get("risk_debate_state"):
|
||||
risk = final_state["risk_debate_state"]
|
||||
risk_reports = []
|
||||
if risk.get("aggressive_history"):
|
||||
risk_reports.append(("Aggressive Analyst", risk["aggressive_history"]))
|
||||
if risk.get("conservative_history"):
|
||||
risk_reports.append(("Conservative Analyst", risk["conservative_history"]))
|
||||
if risk.get("neutral_history"):
|
||||
risk_reports.append(("Neutral Analyst", risk["neutral_history"]))
|
||||
if risk_reports:
|
||||
console.print(Panel("[bold]IV. Risk Management Team Decision[/bold]", border_style="red"))
|
||||
for title, content in risk_reports:
|
||||
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
|
||||
|
||||
# V. Portfolio Manager Decision
|
||||
if risk.get("judge_decision"):
|
||||
console.print(Panel("[bold]V. Portfolio Manager Decision[/bold]", border_style="green"))
|
||||
console.print(Panel(Markdown(risk["judge_decision"]), title="Portfolio Manager", border_style="blue", padding=(1, 2)))
|
||||
|
||||
|
||||
def update_research_team_status(status):
|
||||
"""Update status for research team members (not Trader)."""
|
||||
research_team = ["Bull Researcher", "Bear Researcher", "Research Manager"]
|
||||
for agent in research_team:
|
||||
message_buffer.update_agent_status(agent, status)
|
||||
|
||||
|
||||
# Ordered list of analysts for status transitions
|
||||
ANALYST_ORDER = ["market", "social", "news", "fundamentals"]
|
||||
|
||||
ANALYST_AGENT_NAMES = {
|
||||
"market": "Market Analyst",
|
||||
"social": "Sentiment Analyst",
|
||||
"news": "News Analyst",
|
||||
"fundamentals": "Fundamentals Analyst",
|
||||
}
|
||||
|
||||
ANALYST_REPORT_MAP = {
|
||||
"market": "market_report",
|
||||
"social": "sentiment_report",
|
||||
"news": "news_report",
|
||||
"fundamentals": "fundamentals_report",
|
||||
}
|
||||
|
||||
|
||||
def update_analyst_statuses(message_buffer, chunk, wall_time_tracker=None):
|
||||
"""Update analyst statuses based on accumulated report state.
|
||||
|
||||
Logic:
|
||||
- Store new report content from the current chunk if present
|
||||
- Check accumulated report_sections (not just current chunk) for status
|
||||
- Analysts with reports = completed
|
||||
- First analyst without report = in_progress
|
||||
- Remaining analysts without reports = pending
|
||||
- When all analysts done, set Bull Researcher to in_progress
|
||||
"""
|
||||
selected = message_buffer.selected_analysts
|
||||
found_active = False
|
||||
|
||||
if wall_time_tracker is not None:
|
||||
sync_analyst_tracker_from_chunk(wall_time_tracker, chunk)
|
||||
|
||||
for analyst_key in ANALYST_ORDER:
|
||||
if analyst_key not in selected:
|
||||
continue
|
||||
|
||||
agent_name = ANALYST_AGENT_NAMES[analyst_key]
|
||||
report_key = ANALYST_REPORT_MAP[analyst_key]
|
||||
|
||||
# Capture new report content from current chunk
|
||||
if chunk.get(report_key):
|
||||
message_buffer.update_report_section(report_key, chunk[report_key])
|
||||
|
||||
# Determine status from accumulated sections, not just current chunk
|
||||
has_report = bool(message_buffer.report_sections.get(report_key))
|
||||
|
||||
if has_report:
|
||||
message_buffer.update_agent_status(agent_name, "completed")
|
||||
elif not found_active:
|
||||
message_buffer.update_agent_status(agent_name, "in_progress")
|
||||
found_active = True
|
||||
else:
|
||||
message_buffer.update_agent_status(agent_name, "pending")
|
||||
|
||||
# When all analysts complete, transition research team to in_progress
|
||||
if (
|
||||
not found_active
|
||||
and selected
|
||||
and message_buffer.agent_status.get("Bull Researcher") == "pending"
|
||||
):
|
||||
message_buffer.update_agent_status("Bull Researcher", "in_progress")
|
||||
|
||||
|
||||
def extract_content_string(content):
|
||||
"""Extract string content from various message formats.
|
||||
Returns None if no meaningful text content is found.
|
||||
"""
|
||||
def is_empty(val):
|
||||
"""Whether a value carries nothing to show.
|
||||
|
||||
Text is judged by whether anything was written, not by what it would
|
||||
mean as Python: a report saying "0" or "None" is a message the run
|
||||
produced, and reading it as a falsy literal dropped it from the display.
|
||||
"""
|
||||
if isinstance(val, str):
|
||||
return not val.strip()
|
||||
return val is None or not bool(val)
|
||||
|
||||
if is_empty(content):
|
||||
return None
|
||||
|
||||
if isinstance(content, str):
|
||||
return content.strip()
|
||||
|
||||
if isinstance(content, dict):
|
||||
text = content.get('text', '')
|
||||
return text.strip() if not is_empty(text) else None
|
||||
|
||||
if isinstance(content, list):
|
||||
text_parts = [
|
||||
item.get('text', '').strip() if isinstance(item, dict) and item.get('type') == 'text'
|
||||
else (item.strip() if isinstance(item, str) else '')
|
||||
for item in content
|
||||
]
|
||||
result = ' '.join(t for t in text_parts if t and not is_empty(t))
|
||||
return result if result else None
|
||||
|
||||
return str(content).strip() if not is_empty(content) else None
|
||||
|
||||
|
||||
def classify_message_type(message) -> tuple[str, str | None]:
|
||||
"""Classify LangChain message into display type and extract content.
|
||||
|
||||
Returns:
|
||||
(type, content) - type is one of: User, Agent, Data, Control
|
||||
- content is extracted string or None
|
||||
"""
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
|
||||
content = extract_content_string(getattr(message, 'content', None))
|
||||
|
||||
if isinstance(message, HumanMessage):
|
||||
if content and content.strip() == "Continue":
|
||||
return ("Control", content)
|
||||
return ("User", content)
|
||||
|
||||
if isinstance(message, ToolMessage):
|
||||
return ("Data", content)
|
||||
|
||||
if isinstance(message, AIMessage):
|
||||
return ("Agent", content)
|
||||
|
||||
# Fallback for unknown types
|
||||
return ("System", content)
|
||||
|
||||
|
||||
def format_tool_args(args, max_length=80) -> str:
|
||||
"""Format tool arguments for terminal display."""
|
||||
result = str(args)
|
||||
if len(result) > max_length:
|
||||
return result[:max_length - 3] + "..."
|
||||
return result
|
||||
|
||||
|
||||
class AnalystWallTimeTracker:
|
||||
def __init__(self, plan: AnalystExecutionPlan):
|
||||
self.plan = plan
|
||||
self._started_at: dict[str, float] = {}
|
||||
self._wall_times: dict[str, float] = {}
|
||||
|
||||
def mark_started(self, analyst_key: str, started_at: float | None = None) -> None:
|
||||
if analyst_key not in ANALYST_NODE_SPECS:
|
||||
raise ValueError(f"unknown analyst key: {analyst_key}")
|
||||
self._started_at.setdefault(analyst_key, monotonic() if started_at is None else started_at)
|
||||
|
||||
def mark_completed(
|
||||
self,
|
||||
analyst_key: str,
|
||||
completed_at: float | None = None,
|
||||
) -> None:
|
||||
if analyst_key not in ANALYST_NODE_SPECS:
|
||||
raise ValueError(f"unknown analyst key: {analyst_key}")
|
||||
if analyst_key in self._wall_times:
|
||||
return
|
||||
started_at = self._started_at.get(analyst_key)
|
||||
if started_at is None:
|
||||
return
|
||||
finished_at = monotonic() if completed_at is None else completed_at
|
||||
self._wall_times[analyst_key] = max(0.0, finished_at - started_at)
|
||||
|
||||
def format_summary(self) -> str:
|
||||
parts = []
|
||||
for spec in self.plan.specs:
|
||||
duration = self._wall_times.get(spec.key)
|
||||
if duration is not None:
|
||||
label = spec.agent_node.removesuffix(" Analyst")
|
||||
parts.append(f"{label} {duration:.2f}s")
|
||||
if not parts:
|
||||
return "Analyst wall time: pending"
|
||||
return "Analyst wall time: " + " | ".join(parts)
|
||||
|
||||
|
||||
def sync_analyst_tracker_from_chunk(
|
||||
tracker: AnalystWallTimeTracker,
|
||||
chunk: dict[str, str],
|
||||
now: float | None = None,
|
||||
) -> None:
|
||||
current_time = monotonic() if now is None else now
|
||||
active_found = False
|
||||
|
||||
for spec in tracker.plan.specs:
|
||||
has_report = bool(chunk.get(spec.report_key))
|
||||
|
||||
if has_report:
|
||||
tracker.mark_started(spec.key, started_at=current_time)
|
||||
tracker.mark_completed(spec.key, completed_at=current_time)
|
||||
continue
|
||||
|
||||
if not active_found:
|
||||
tracker.mark_started(spec.key, started_at=current_time)
|
||||
active_found = True
|
||||
+108
-987
File diff suppressed because it is too large
Load Diff
+7
-2
@@ -1,10 +1,15 @@
|
||||
from enum import Enum
|
||||
from typing import List, Optional, Dict
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class AnalystType(str, Enum):
|
||||
MARKET = "market"
|
||||
# Wire value stays "social" for saved-config and string-keyed-caller
|
||||
# back-compat; the user-facing label is "Sentiment Analyst".
|
||||
SOCIAL = "social"
|
||||
NEWS = "news"
|
||||
FUNDAMENTALS = "fundamentals"
|
||||
|
||||
|
||||
class AssetType(str, Enum):
|
||||
STOCK = "stock"
|
||||
CRYPTO = "crypto"
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
"""What the last run chose, offered back as the next run's defaults.
|
||||
|
||||
The interactive flow asks the same questions every time, and only some of them
|
||||
have an environment variable to skip them (the analyst set has none). Remembered
|
||||
answers prefill the prompts so Enter accepts them; they never skip a step, so a
|
||||
run always starts on choices the user has seen.
|
||||
|
||||
Only answers that are stable between runs are kept. The ticker and the analysis
|
||||
date are not: they change every run, and a remembered date would quietly offer a
|
||||
stale one.
|
||||
|
||||
Every value is checked against the current choices on the way out, because
|
||||
models and providers are added and retired between versions. A remembered model
|
||||
that is no longer offered is dropped rather than shown.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from cli.models import AnalystType, AssetType
|
||||
from cli.prompts import _llm_provider_table, filter_analysts_for_asset_type
|
||||
from tradingagents.llm_clients.model_catalog import get_model_options
|
||||
|
||||
_PREFS_PATH = Path(os.path.expanduser("~")) / ".tradingagents" / "cli_prefs.json"
|
||||
|
||||
REMEMBERED = (
|
||||
"output_language", "analysts", "research_depth", "llm_provider",
|
||||
"quick_think_llm", "deep_think_llm", "backend_url",
|
||||
)
|
||||
|
||||
|
||||
def load_last_run() -> dict:
|
||||
"""The previous run's answers, or an empty dict when there is nothing usable.
|
||||
|
||||
Convenience state: an unreadable or corrupt file means no defaults, never an
|
||||
error in the user's way.
|
||||
"""
|
||||
try:
|
||||
data = json.loads(_PREFS_PATH.read_text(encoding="utf-8"))
|
||||
return data if isinstance(data, dict) else {}
|
||||
except (OSError, ValueError):
|
||||
return {}
|
||||
|
||||
|
||||
def save_last_run(selections: dict) -> None:
|
||||
"""Record the answers worth offering next time; failure is never fatal."""
|
||||
kept = {k: v for k, v in selections.items() if k in REMEMBERED and v not in (None, "", [])}
|
||||
kept["analysts"] = [getattr(a, "value", a) for a in kept.get("analysts", [])] or None
|
||||
kept = {k: v for k, v in kept.items() if v is not None}
|
||||
try:
|
||||
_PREFS_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
temp = _PREFS_PATH.with_suffix(".tmp")
|
||||
temp.write_text(json.dumps(kept, indent=2), encoding="utf-8")
|
||||
os.replace(temp, _PREFS_PATH) # a concurrent run reads one file or the other
|
||||
except OSError:
|
||||
return
|
||||
|
||||
|
||||
def sanitize(prefs: dict, asset_type) -> dict:
|
||||
"""Keep only the remembered answers that are still choosable now."""
|
||||
kept: dict = {}
|
||||
if isinstance(prefs.get("output_language"), str):
|
||||
kept["output_language"] = prefs["output_language"]
|
||||
if prefs.get("research_depth") in (1, 3, 5):
|
||||
kept["research_depth"] = prefs["research_depth"]
|
||||
|
||||
known = {a.value for a in AnalystType}
|
||||
analysts = [a for a in prefs.get("analysts") or [] if a in known]
|
||||
allowed = filter_analysts_for_asset_type([AnalystType(a) for a in analysts], AssetType(asset_type))
|
||||
if allowed:
|
||||
kept["analysts"] = [a.value for a in allowed]
|
||||
|
||||
provider = prefs.get("llm_provider")
|
||||
# Region-specific providers (qwen-cn) are picked in a second prompt, so the
|
||||
# base key is what the provider menu matches.
|
||||
base = (provider or "").split("-cn")[0]
|
||||
if base and base in {key for _, key, _ in _llm_provider_table()}:
|
||||
kept["llm_provider"] = provider
|
||||
if isinstance(prefs.get("backend_url"), str) and prefs["backend_url"]:
|
||||
kept["backend_url"] = prefs["backend_url"]
|
||||
for field, mode in (("quick_think_llm", "quick"), ("deep_think_llm", "deep")):
|
||||
try:
|
||||
offered = {model for _, model in get_model_options(base, mode)}
|
||||
except KeyError:
|
||||
continue
|
||||
if prefs.get(field) in offered:
|
||||
kept[field] = prefs[field]
|
||||
return kept
|
||||
+683
@@ -0,0 +1,683 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import questionary
|
||||
from dotenv import find_dotenv, set_key
|
||||
|
||||
from cli.display import console
|
||||
from cli.models import AnalystType, AssetType
|
||||
from tradingagents.llm_clients.api_key_env import get_api_key_env
|
||||
from tradingagents.llm_clients.model_catalog import get_model_options
|
||||
|
||||
TICKER_INPUT_EXAMPLES = "SPY, 0700.HK, BTC-USD"
|
||||
|
||||
ANALYST_CHOICES = [
|
||||
("Market Analyst", AnalystType.MARKET),
|
||||
("Sentiment Analyst", AnalystType.SOCIAL),
|
||||
("News Analyst", AnalystType.NEWS),
|
||||
("Fundamentals Analyst", AnalystType.FUNDAMENTALS),
|
||||
]
|
||||
|
||||
CRYPTO_SUFFIXES = ("-USD", "-USDT", "-USDC", "-BTC", "-ETH")
|
||||
|
||||
|
||||
def is_valid_ticker_input(value: str) -> bool:
|
||||
"""Whether a ticker entry is acceptable (charset + length).
|
||||
|
||||
Allows the characters Yahoo symbols use, including ``=`` for futures/forex
|
||||
like ``GC=F`` and ``EURUSD=X`` (#980), and ``^`` for indices. Empty input is
|
||||
allowed (it defaults to SPY downstream).
|
||||
"""
|
||||
v = value.strip()
|
||||
return not v or (all(ch.isalnum() or ch in "._-^=" for ch in v) and len(v) <= 32)
|
||||
|
||||
|
||||
def get_ticker() -> str:
|
||||
"""Prompt the user to enter a ticker symbol, preserving exchange suffixes.
|
||||
|
||||
Uses questionary.text (not typer.prompt, which strips trailing dot-suffixes
|
||||
like ``000404.SH`` on some shells) and validates the symbol charset so an
|
||||
obvious typo is caught before the run starts.
|
||||
"""
|
||||
ticker = questionary.text(
|
||||
f"Enter ticker symbol (e.g. {TICKER_INPUT_EXAMPLES}):",
|
||||
validate=lambda x: (
|
||||
is_valid_ticker_input(x)
|
||||
or "Please enter a valid ticker symbol, e.g. AAPL, 000404.SZ, 0700.HK, GC=F."
|
||||
),
|
||||
style=questionary.Style(
|
||||
[
|
||||
("text", "fg:green"),
|
||||
("highlighted", "noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if ticker is None:
|
||||
console.print("\n[red]No ticker symbol provided. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return normalize_ticker_symbol(ticker) if ticker.strip() else "SPY"
|
||||
|
||||
|
||||
def normalize_ticker_symbol(ticker: str) -> str:
|
||||
"""Resolve user input to its canonical Yahoo symbol (single source of truth).
|
||||
|
||||
Delegates to the data layer's ``normalize_symbol`` so the symbol the CLI
|
||||
passes through the pipeline is exactly the one the data path will price
|
||||
(e.g. ``BTCUSD`` -> ``BTC-USD``, ``XAUUSD`` -> ``GC=F``). Falls back to the
|
||||
plain upper-case if the data layer is unavailable.
|
||||
"""
|
||||
try:
|
||||
from tradingagents.dataflows.symbols import normalize_symbol
|
||||
|
||||
return normalize_symbol(ticker)
|
||||
except Exception:
|
||||
return ticker.strip().upper()
|
||||
|
||||
|
||||
def detect_asset_type(ticker: str) -> AssetType:
|
||||
"""Classify on the canonical symbol so e.g. BTCUSD and BTC-USDT both read as
|
||||
crypto (#981/#982), matching what the data path will actually fetch."""
|
||||
canonical = normalize_ticker_symbol(ticker)
|
||||
if canonical.endswith(CRYPTO_SUFFIXES):
|
||||
return AssetType.CRYPTO
|
||||
return AssetType.STOCK
|
||||
|
||||
|
||||
def filter_analysts_for_asset_type(
|
||||
analysts: list[AnalystType], asset_type: AssetType
|
||||
) -> list[AnalystType]:
|
||||
if asset_type != AssetType.CRYPTO:
|
||||
return analysts
|
||||
return [
|
||||
analyst
|
||||
for analyst in analysts
|
||||
if analyst != AnalystType.FUNDAMENTALS
|
||||
]
|
||||
|
||||
|
||||
def _matching_choice(options, default):
|
||||
"""The option value equal to ``default``, or None to leave the menu as is."""
|
||||
return next((value for _, value in options if value == default), None)
|
||||
|
||||
|
||||
def select_analysts(asset_type: AssetType = AssetType.STOCK, default=None) -> list[AnalystType]:
|
||||
"""Select analysts using an interactive checkbox.
|
||||
|
||||
``default`` pre-checks the previous run's analysts; the prompt still shows.
|
||||
"""
|
||||
available_analysts = filter_analysts_for_asset_type(
|
||||
[value for _, value in ANALYST_CHOICES],
|
||||
asset_type,
|
||||
)
|
||||
choices = questionary.checkbox(
|
||||
"Select Your [Analysts Team]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value, checked=value.value in (default or []))
|
||||
for display, value in ANALYST_CHOICES
|
||||
if value in available_analysts
|
||||
],
|
||||
instruction="\n- Press Space to select/unselect analysts\n- Press 'a' to select/unselect all\n- Press Enter when done",
|
||||
validate=lambda x: len(x) > 0 or "You must select at least one analyst.",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("checkbox-selected", "fg:green"),
|
||||
("selected", "fg:green noinherit"),
|
||||
("highlighted", "noinherit"),
|
||||
("pointer", "noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if not choices:
|
||||
console.print("\n[red]No analysts selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return choices
|
||||
|
||||
|
||||
def select_research_depth(default=None) -> int:
|
||||
"""Select research depth using an interactive selection."""
|
||||
|
||||
DEPTH_OPTIONS = [
|
||||
("Shallow - Quick research, few debate and strategy discussion rounds", 1),
|
||||
("Medium - Middle ground, moderate debate rounds and strategy discussion", 3),
|
||||
("Deep - Comprehensive research, in depth debate and strategy discussion", 5),
|
||||
]
|
||||
|
||||
choice = questionary.select(
|
||||
"Select Your [Research Depth]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value) for display, value in DEPTH_OPTIONS
|
||||
],
|
||||
default=_matching_choice(DEPTH_OPTIONS, default),
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:yellow noinherit"),
|
||||
("highlighted", "fg:yellow noinherit"),
|
||||
("pointer", "fg:yellow noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print("\n[red]No research depth selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return choice
|
||||
|
||||
|
||||
# Mainstream OpenRouter chat-LLM provider namespaces. We surface the newest
|
||||
# models from these rather than the universal-newest, which is dominated by
|
||||
# niche/experimental releases. These are the general-purpose chat providers;
|
||||
# more enterprise/specialised namespaces (nvidia, cohere, amazon, ...) tend to
|
||||
# ship research/safety variants as their newest, so they're left out of the
|
||||
# shortlist. Provider names are stable (unlike model IDs), so this rarely needs
|
||||
# touching; anything not here is still reachable via Custom ID.
|
||||
_OPENROUTER_MAINSTREAM = {
|
||||
"openai", "anthropic", "google", "deepseek", "qwen", "mistralai",
|
||||
"meta-llama", "x-ai", "z-ai", "minimax", "moonshotai",
|
||||
}
|
||||
|
||||
|
||||
def _fetch_openrouter_models() -> list[tuple[str, str]]:
|
||||
"""Fetch available models from the OpenRouter API."""
|
||||
import requests
|
||||
try:
|
||||
resp = requests.get("https://openrouter.ai/api/v1/models", timeout=10)
|
||||
resp.raise_for_status()
|
||||
models = resp.json().get("data", [])
|
||||
# Newest first so the top-N shown really is the latest available — the
|
||||
# API currently returns this order, but sort explicitly so the prompt's
|
||||
# "latest available" label holds regardless of response ordering.
|
||||
models.sort(key=lambda m: m.get("created") or 0, reverse=True)
|
||||
return [(m.get("name") or m["id"], m["id"]) for m in models]
|
||||
except Exception as e:
|
||||
console.print(f"\n[yellow]Could not fetch OpenRouter models: {e}[/yellow]")
|
||||
return []
|
||||
|
||||
|
||||
def _require_text(message: str, hint: str) -> str:
|
||||
"""Prompt for a required value; exit cleanly if the user cancels.
|
||||
|
||||
``questionary.text(...).ask()`` returns None on Ctrl-C/Esc; mirror the
|
||||
exit-on-cancel behavior of the other required selections so a cancelled
|
||||
prompt never returns an empty model/deployment that would fail downstream.
|
||||
"""
|
||||
response = questionary.text(
|
||||
message,
|
||||
validate=lambda x: len(x.strip()) > 0 or hint,
|
||||
).ask()
|
||||
if response is None:
|
||||
console.print("\n[red]Cancelled. Exiting...[/red]")
|
||||
exit(1)
|
||||
return response.strip()
|
||||
|
||||
|
||||
def select_openrouter_model(mode: str) -> str:
|
||||
"""Select an OpenRouter model from the newest available, or enter a custom ID.
|
||||
|
||||
``mode`` ("quick"/"deep") labels the prompt so the two consecutive
|
||||
OpenRouter selections are distinguishable, like the other providers (#1000).
|
||||
"""
|
||||
models = _fetch_openrouter_models() # newest first
|
||||
# Prefer the newest from mainstream providers so the shortlist isn't crowded
|
||||
# out by niche/experimental releases; fall back to all if none match.
|
||||
mainstream = [
|
||||
(name, mid) for name, mid in models
|
||||
if not mid.startswith("~") # skip variant/alias duplicate routes
|
||||
and mid.split("/", 1)[0] in _OPENROUTER_MAINSTREAM
|
||||
]
|
||||
top = (mainstream or models)[:5]
|
||||
|
||||
choices = [questionary.Choice(name, value=mid) for name, mid in top]
|
||||
choices.append(questionary.Choice("Custom model ID", value="custom"))
|
||||
|
||||
choice = questionary.select(
|
||||
f"Select Your [{mode.title()}-Thinking] OpenRouter Model (latest available):",
|
||||
choices=choices,
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style([
|
||||
("selected", "fg:magenta noinherit"),
|
||||
("highlighted", "fg:magenta noinherit"),
|
||||
("pointer", "fg:magenta noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print("\n[red]No model selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
if choice == "custom":
|
||||
return _require_text(
|
||||
"Enter OpenRouter model ID (e.g. google/gemma-4-26b-a4b-it):",
|
||||
"Please enter a model ID.",
|
||||
)
|
||||
return choice
|
||||
|
||||
|
||||
def _prompt_custom_model_id() -> str:
|
||||
"""Prompt user to type a custom model ID."""
|
||||
return _require_text("Enter model ID:", "Please enter a model ID.")
|
||||
|
||||
|
||||
def _select_model(provider: str, mode: str, default=None) -> str:
|
||||
"""Select a model for the given provider and mode (quick/deep)."""
|
||||
if provider.lower() == "openrouter":
|
||||
return select_openrouter_model(mode)
|
||||
|
||||
if provider.lower() == "azure":
|
||||
return _require_text(
|
||||
f"Enter Azure deployment name ({mode}-thinking):",
|
||||
"Please enter a deployment name.",
|
||||
)
|
||||
|
||||
choice = questionary.select(
|
||||
f"Select Your [{mode.title()}-Thinking LLM Engine]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value)
|
||||
for display, value in get_model_options(provider, mode)
|
||||
],
|
||||
default=_matching_choice(get_model_options(provider, mode), default),
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:magenta noinherit"),
|
||||
("highlighted", "fg:magenta noinherit"),
|
||||
("pointer", "fg:magenta noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print(f"\n[red]No {mode} thinking llm engine selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
if choice == "custom":
|
||||
return _prompt_custom_model_id()
|
||||
|
||||
return choice
|
||||
|
||||
|
||||
def select_shallow_thinking_agent(provider, default=None) -> str:
|
||||
"""Select shallow thinking llm engine using an interactive selection."""
|
||||
return _select_model(provider, "quick", default)
|
||||
|
||||
|
||||
def select_deep_thinking_agent(provider, default=None) -> str:
|
||||
"""Select deep thinking llm engine using an interactive selection."""
|
||||
return _select_model(provider, "deep", default)
|
||||
|
||||
|
||||
def _llm_provider_table() -> list[tuple[str, str, str | None]]:
|
||||
"""(display_name, provider_key, base_url) for every supported provider.
|
||||
|
||||
Shared by the interactive picker and by env-driven configuration so an
|
||||
env-set provider resolves to the same default endpoint the menu uses.
|
||||
Ollama users can point at a remote ollama-serve via OLLAMA_BASE_URL
|
||||
(convention from the broader Ollama ecosystem); falls back to the
|
||||
localhost default when unset.
|
||||
"""
|
||||
ollama_url = os.environ.get("OLLAMA_BASE_URL") or "http://localhost:11434/v1"
|
||||
return [
|
||||
("OpenAI", "openai", "https://api.openai.com/v1"),
|
||||
("Google", "google", None),
|
||||
("Anthropic", "anthropic", "https://api.anthropic.com/"),
|
||||
("xAI", "xai", "https://api.x.ai/v1"),
|
||||
("DeepSeek", "deepseek", "https://api.deepseek.com"),
|
||||
("Qwen", "qwen", "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"),
|
||||
# Z.AI international, the platform ZHIPU_API_KEY belongs to; the CN
|
||||
# platform is the separate glm-cn key, chosen in the region prompt.
|
||||
("GLM", "glm", "https://api.z.ai/api/paas/v4/"),
|
||||
("MiniMax", "minimax", "https://api.minimax.io/v1"),
|
||||
("OpenRouter", "openrouter", "https://openrouter.ai/api/v1"),
|
||||
("Mistral", "mistral", "https://api.mistral.ai/v1"),
|
||||
("Kimi (Moonshot)", "kimi", "https://api.moonshot.ai/v1"),
|
||||
("Groq", "groq", "https://api.groq.com/openai/v1"),
|
||||
("NVIDIA NIM", "nvidia", "https://integrate.api.nvidia.com/v1"),
|
||||
("Azure OpenAI", "azure", None),
|
||||
("Amazon Bedrock", "bedrock", None),
|
||||
("Ollama", "ollama", ollama_url),
|
||||
("OpenAI-compatible (vLLM, LM Studio, llama.cpp, custom relay)", "openai_compatible", None),
|
||||
]
|
||||
|
||||
|
||||
def provider_default_url(provider_key: str) -> str | None:
|
||||
"""Return the default backend URL for a provider key, or None if unknown."""
|
||||
key = provider_key.lower()
|
||||
for _, pk, url in _llm_provider_table():
|
||||
if pk == key:
|
||||
return url
|
||||
return None
|
||||
|
||||
|
||||
def resolve_backend_url(
|
||||
provider: str, menu_url: str | None = None, env_url: str | None = None
|
||||
) -> str | None:
|
||||
"""Resolve the backend URL with the correct precedence.
|
||||
|
||||
An explicit env override (``env_url``, from ``TRADINGAGENTS_LLM_BACKEND_URL``
|
||||
via ``DEFAULT_CONFIG['backend_url']``) is honored regardless of how the
|
||||
provider was chosen — interactively or from the environment (#978).
|
||||
Otherwise the menu/region URL, then the provider's default.
|
||||
"""
|
||||
return env_url or menu_url or provider_default_url(provider)
|
||||
|
||||
|
||||
def prompt_openai_compatible_url(default=None) -> str:
|
||||
"""Prompt for a custom OpenAI-compatible endpoint base URL."""
|
||||
url = questionary.text(
|
||||
"Enter the OpenAI-compatible base URL "
|
||||
"(e.g. http://localhost:8000/v1 for vLLM, http://localhost:1234/v1 for LM Studio):",
|
||||
default=default or "",
|
||||
validate=lambda x: x.strip().startswith(("http://", "https://"))
|
||||
or "Enter a URL starting with http:// or https://",
|
||||
).ask()
|
||||
if not url:
|
||||
console.print("\n[red]No endpoint URL provided. Exiting...[/red]")
|
||||
exit(1)
|
||||
return url.strip()
|
||||
|
||||
|
||||
def select_llm_provider(default=None) -> tuple[str, str | None]:
|
||||
"""Select the LLM provider and its API endpoint."""
|
||||
PROVIDERS = _llm_provider_table()
|
||||
# A region-specific key (qwen-cn) is chosen in a later prompt; the menu
|
||||
# lists the base provider.
|
||||
base = (default or "").split("-cn")[0]
|
||||
preselected = next(
|
||||
((key, url) for _, key, url in PROVIDERS if key == base), None
|
||||
)
|
||||
|
||||
choice = questionary.select(
|
||||
"Select your LLM Provider:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=(provider_key, url))
|
||||
for display, provider_key, url in PROVIDERS
|
||||
],
|
||||
default=preselected,
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:magenta noinherit"),
|
||||
("highlighted", "fg:magenta noinherit"),
|
||||
("pointer", "fg:magenta noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print("\n[red]No LLM provider selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
provider, url = choice
|
||||
return provider, url
|
||||
|
||||
|
||||
def ask_openai_reasoning_effort() -> str:
|
||||
"""Ask for OpenAI reasoning effort level."""
|
||||
choices = [
|
||||
questionary.Choice("Medium (Default)", "medium"),
|
||||
questionary.Choice("High (More thorough)", "high"),
|
||||
questionary.Choice("Low (Faster)", "low"),
|
||||
]
|
||||
return questionary.select(
|
||||
"Select Reasoning Effort:",
|
||||
choices=choices,
|
||||
style=questionary.Style([
|
||||
("selected", "fg:cyan noinherit"),
|
||||
("highlighted", "fg:cyan noinherit"),
|
||||
("pointer", "fg:cyan noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def ask_anthropic_effort() -> str | None:
|
||||
"""Ask for Anthropic effort level.
|
||||
|
||||
Controls token usage and response thoroughness on Claude 4.5 / 4.6 / 4.7
|
||||
models. The API also accepts "max"; we expose low/medium/high as the
|
||||
common selection range.
|
||||
"""
|
||||
return questionary.select(
|
||||
"Select Effort Level:",
|
||||
choices=[
|
||||
questionary.Choice("High (recommended)", "high"),
|
||||
questionary.Choice("Medium (balanced)", "medium"),
|
||||
questionary.Choice("Low (faster, cheaper)", "low"),
|
||||
],
|
||||
style=questionary.Style([
|
||||
("selected", "fg:cyan noinherit"),
|
||||
("highlighted", "fg:cyan noinherit"),
|
||||
("pointer", "fg:cyan noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def ask_gemini_thinking_config() -> str | None:
|
||||
"""Ask for Gemini thinking configuration.
|
||||
|
||||
Returns thinking_level: "high" or "minimal".
|
||||
Client maps to appropriate API param based on model series.
|
||||
"""
|
||||
return questionary.select(
|
||||
"Select Thinking Mode:",
|
||||
choices=[
|
||||
questionary.Choice("Enable Thinking (recommended)", "high"),
|
||||
questionary.Choice("Minimal/Disable Thinking", "minimal"),
|
||||
],
|
||||
style=questionary.Style([
|
||||
("selected", "fg:green noinherit"),
|
||||
("highlighted", "fg:green noinherit"),
|
||||
("pointer", "fg:green noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def ask_glm_region() -> tuple[str, str]:
|
||||
"""Ask which GLM platform (Z.AI international vs BigModel China) to use.
|
||||
|
||||
Zhipu serves the same GLM models under two brands with separate
|
||||
accounts; keys aren't interchangeable. Returns (provider_key, backend_url).
|
||||
"""
|
||||
return questionary.select(
|
||||
"Select GLM platform:",
|
||||
choices=[
|
||||
questionary.Choice(
|
||||
"Z.AI — api.z.ai (international, uses ZHIPU_API_KEY)",
|
||||
value=("glm", "https://api.z.ai/api/paas/v4/"),
|
||||
),
|
||||
questionary.Choice(
|
||||
"BigModel — open.bigmodel.cn (China, uses ZHIPU_CN_API_KEY)",
|
||||
value=("glm-cn", "https://open.bigmodel.cn/api/paas/v4/"),
|
||||
),
|
||||
],
|
||||
style=questionary.Style([
|
||||
("selected", "fg:cyan noinherit"),
|
||||
("highlighted", "fg:cyan noinherit"),
|
||||
("pointer", "fg:cyan noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def ask_qwen_region() -> tuple[str, str]:
|
||||
"""Ask which Qwen region (international vs China) to use.
|
||||
|
||||
Alibaba DashScope exposes two endpoints with separate accounts —
|
||||
a key from one region does NOT authenticate against the other
|
||||
(fixes #758). Returns (provider_key, backend_url).
|
||||
"""
|
||||
return questionary.select(
|
||||
"Select Qwen region:",
|
||||
choices=[
|
||||
questionary.Choice(
|
||||
"International — dashscope-intl.aliyuncs.com (uses DASHSCOPE_API_KEY)",
|
||||
value=("qwen", "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"),
|
||||
),
|
||||
questionary.Choice(
|
||||
"China — dashscope.aliyuncs.com (uses DASHSCOPE_CN_API_KEY)",
|
||||
value=("qwen-cn", "https://dashscope.aliyuncs.com/compatible-mode/v1"),
|
||||
),
|
||||
],
|
||||
style=questionary.Style([
|
||||
("selected", "fg:cyan noinherit"),
|
||||
("highlighted", "fg:cyan noinherit"),
|
||||
("pointer", "fg:cyan noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def ask_minimax_region() -> tuple[str, str]:
|
||||
"""Ask which MiniMax region (global vs China) to use.
|
||||
|
||||
MiniMax exposes two endpoints with separate accounts — a key from
|
||||
one region does NOT authenticate against the other. Returns
|
||||
(provider_key, backend_url).
|
||||
"""
|
||||
return questionary.select(
|
||||
"Select MiniMax region:",
|
||||
choices=[
|
||||
questionary.Choice(
|
||||
"Global — api.minimax.io (uses MINIMAX_API_KEY)",
|
||||
value=("minimax", "https://api.minimax.io/v1"),
|
||||
),
|
||||
questionary.Choice(
|
||||
"China — api.minimaxi.com (uses MINIMAX_CN_API_KEY)",
|
||||
value=("minimax-cn", "https://api.minimaxi.com/v1"),
|
||||
),
|
||||
],
|
||||
style=questionary.Style([
|
||||
("selected", "fg:cyan noinherit"),
|
||||
("highlighted", "fg:cyan noinherit"),
|
||||
("pointer", "fg:cyan noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
|
||||
def confirm_ollama_endpoint(url: str) -> None:
|
||||
"""Show the resolved Ollama endpoint after provider selection.
|
||||
|
||||
Surfaces three things the user benefits from seeing before model
|
||||
selection: which URL we'll actually hit, where it came from
|
||||
(`OLLAMA_BASE_URL` vs default), and a soft warning if the URL is
|
||||
missing the scheme/port that ollama-serve expects. The warning is
|
||||
advisory only — we don't reject malformed input, since the user may
|
||||
be doing something deliberately unusual (e.g. a reverse-proxy path).
|
||||
"""
|
||||
from_env = os.environ.get("OLLAMA_BASE_URL")
|
||||
origin = " (from OLLAMA_BASE_URL)" if from_env and from_env == url else ""
|
||||
console.print(f"[green]✓ Using Ollama at {url}{origin}[/green]")
|
||||
|
||||
if not url.startswith(("http://", "https://")):
|
||||
console.print(
|
||||
f"[yellow]Note: {url!r} is missing a scheme. "
|
||||
f"Ollama-serve typically expects a URL like "
|
||||
f"http://<host>:11434/v1.[/yellow]"
|
||||
)
|
||||
elif ":11434" not in url and "://localhost" not in url and "://127.0.0.1" not in url:
|
||||
# Soft hint when the port differs from the ollama-serve default
|
||||
# and the host isn't local (where users sometimes proxy on :80).
|
||||
console.print(
|
||||
f"[yellow]Note: {url!r} doesn't include port 11434. "
|
||||
f"Make sure your remote ollama-serve listens on the port "
|
||||
f"shown above.[/yellow]"
|
||||
)
|
||||
|
||||
|
||||
def ensure_api_key(provider: str) -> str | None:
|
||||
"""Make sure the API key for `provider` is available in the environment.
|
||||
|
||||
If the env var is already set, returns its value untouched. Otherwise
|
||||
interactively prompts the user, persists the value to the project's
|
||||
.env file via python-dotenv's set_key (creating .env if needed), and
|
||||
exports it into os.environ so the current process picks it up.
|
||||
|
||||
Returns None for providers that do not require a key (e.g. ollama)
|
||||
and for providers not found in the canonical mapping.
|
||||
"""
|
||||
env_var = get_api_key_env(provider)
|
||||
if env_var is None:
|
||||
return None # ollama / unknown — no key check possible
|
||||
|
||||
# Key-optional providers (generic OpenAI-compatible / local servers) read the
|
||||
# key when present but must never force an interactive prompt.
|
||||
from tradingagents.llm_clients.openai_client import OPENAI_COMPATIBLE_PROVIDERS
|
||||
spec = OPENAI_COMPATIBLE_PROVIDERS.get(provider.lower())
|
||||
if spec is not None and spec.key_optional:
|
||||
return os.environ.get(env_var)
|
||||
|
||||
existing = os.environ.get(env_var)
|
||||
if existing:
|
||||
return existing
|
||||
|
||||
console.print(
|
||||
f"\n[yellow]{env_var} is not set in your environment.[/yellow]"
|
||||
)
|
||||
key = questionary.password(
|
||||
f"Paste your {env_var} (will be saved to .env):",
|
||||
style=questionary.Style([
|
||||
("text", "fg:cyan"),
|
||||
("highlighted", "noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
if not key:
|
||||
console.print(
|
||||
f"[red]Skipped. API calls will fail until {env_var} is set.[/red]"
|
||||
)
|
||||
return None
|
||||
|
||||
env_path = find_dotenv(usecwd=True) or str(Path.cwd() / ".env")
|
||||
# The file holds credentials, so make it owner-only before writing: create
|
||||
# it 0600 when absent, and tighten an existing one (set_key keeps the mode).
|
||||
if not os.path.exists(env_path):
|
||||
os.close(os.open(env_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o600))
|
||||
os.chmod(env_path, 0o600)
|
||||
set_key(env_path, env_var, key)
|
||||
os.environ[env_var] = key
|
||||
console.print(f"[green]Saved {env_var} to {env_path}[/green]")
|
||||
return key
|
||||
|
||||
|
||||
def ask_output_language(default=None) -> str:
|
||||
"""Ask for report output language.
|
||||
|
||||
``default`` is offered only when it is one of the listed languages: a custom
|
||||
one entered last time is free text, which the menu cannot preselect.
|
||||
"""
|
||||
choices = [
|
||||
questionary.Choice("English (default)", "English"),
|
||||
questionary.Choice("Chinese (中文)", "Chinese"),
|
||||
questionary.Choice("Japanese (日本語)", "Japanese"),
|
||||
questionary.Choice("Korean (한국어)", "Korean"),
|
||||
questionary.Choice("Hindi (हिन्दी)", "Hindi"),
|
||||
questionary.Choice("Spanish (Español)", "Spanish"),
|
||||
questionary.Choice("Portuguese (Português)", "Portuguese"),
|
||||
questionary.Choice("French (Français)", "French"),
|
||||
questionary.Choice("German (Deutsch)", "German"),
|
||||
questionary.Choice("Arabic (العربية)", "Arabic"),
|
||||
questionary.Choice("Russian (Русский)", "Russian"),
|
||||
questionary.Choice("Custom language", "custom"),
|
||||
]
|
||||
choice = questionary.select(
|
||||
"Select Output Language:",
|
||||
choices=choices,
|
||||
default=_matching_choice([(c.title, c.value) for c in choices], default),
|
||||
style=questionary.Style([
|
||||
("selected", "fg:yellow noinherit"),
|
||||
("highlighted", "fg:yellow noinherit"),
|
||||
("pointer", "fg:yellow noinherit"),
|
||||
]),
|
||||
).ask()
|
||||
|
||||
# Output language has a sensible default, so a cancel falls back to English
|
||||
# rather than exiting the run (unlike the required model/provider prompts).
|
||||
if choice is None:
|
||||
return "English"
|
||||
if choice == "custom":
|
||||
return (questionary.text(
|
||||
"Enter language name (e.g. Turkish, Vietnamese, Thai, Indonesian):",
|
||||
validate=lambda x: len(x.strip()) > 0 or "Please enter a language name.",
|
||||
).ask() or "").strip() or "English"
|
||||
|
||||
return choice
|
||||
+393
@@ -0,0 +1,393 @@
|
||||
"""Running one analysis from the CLI: build the graph, stream it into the live view, save the report."""
|
||||
|
||||
import datetime
|
||||
import os
|
||||
import time
|
||||
from functools import wraps
|
||||
from pathlib import Path
|
||||
|
||||
import typer
|
||||
from rich.live import Live
|
||||
|
||||
from cli.display import (
|
||||
ANALYST_ORDER,
|
||||
AnalystWallTimeTracker,
|
||||
classify_message_type,
|
||||
console,
|
||||
create_layout,
|
||||
display_complete_report,
|
||||
message_buffer,
|
||||
update_analyst_statuses,
|
||||
update_display,
|
||||
update_research_team_status,
|
||||
)
|
||||
from cli.selections import get_user_selections
|
||||
from cli.stats_handler import StatsCallbackHandler
|
||||
from tradingagents.agents.rating import is_review
|
||||
from tradingagents.dataflows.symbols import safe_ticker_component
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
from tradingagents.graph.analyst_execution import (
|
||||
build_analyst_execution_plan,
|
||||
)
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
from tradingagents.reporting import write_report_tree
|
||||
|
||||
|
||||
def _run_directory(config: dict, ticker: str, trade_date: str) -> Path:
|
||||
"""Where this run writes, with the ticker validated as a path component.
|
||||
|
||||
Every other path that interpolates a ticker checks it first; a value of
|
||||
".." here would place the run outside the results directory.
|
||||
"""
|
||||
return Path(config["results_dir"]) / safe_ticker_component(ticker) / trade_date
|
||||
|
||||
|
||||
def _announce_checkpoint_state(graph, ticker: str, trade_date: str) -> None:
|
||||
"""Say whether this run resumed a saved one, where the user can see it.
|
||||
|
||||
The graph logs this, but nothing in the CLI configures logging and the live
|
||||
view owns the screen, so a resume was invisible.
|
||||
"""
|
||||
if getattr(graph, "_resuming", False):
|
||||
message_buffer.add_message(
|
||||
"System", f"Resuming the saved run for {ticker} on {trade_date}"
|
||||
)
|
||||
else:
|
||||
message_buffer.add_message("System", f"Starting fresh for {ticker} on {trade_date}")
|
||||
|
||||
|
||||
def _build_run_config(selections: dict, checkpoint: bool | None) -> dict:
|
||||
"""Assemble the run config from interactive selections, honoring env precedence.
|
||||
|
||||
Round counts and checkpoint follow "explicit env/flag wins": an env-applied
|
||||
value on DEFAULT_CONFIG is preserved unless the user overrode it on the CLI.
|
||||
"""
|
||||
config = DEFAULT_CONFIG.copy()
|
||||
# Research depth sets both round counts, but an explicit env override
|
||||
# (TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS) wins over the
|
||||
# interactive selection — leave the env-applied value in place (#977).
|
||||
for env_var, key in (("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "max_debate_rounds"),
|
||||
("TRADINGAGENTS_MAX_RISK_ROUNDS", "max_risk_discuss_rounds")):
|
||||
if os.environ.get(env_var):
|
||||
# The depth prompt still appeared (it is skipped only when both are
|
||||
# set), so say which half of the answer the environment overrode.
|
||||
console.print(
|
||||
f"[green]✓ {key} from environment:[/green] {config[key]} "
|
||||
f"(set by {env_var}, so the research depth you chose does not apply to it)"
|
||||
)
|
||||
else:
|
||||
config[key] = selections["research_depth"]
|
||||
config["quick_think_llm"] = selections["quick_think_llm"]
|
||||
config["deep_think_llm"] = selections["deep_think_llm"]
|
||||
config["backend_url"] = selections["backend_url"]
|
||||
config["llm_provider"] = selections["llm_provider"].lower()
|
||||
# Provider-specific thinking configuration
|
||||
config["google_thinking_level"] = selections.get("google_thinking_level")
|
||||
config["openai_reasoning_effort"] = selections.get("openai_reasoning_effort")
|
||||
config["anthropic_effort"] = selections.get("anthropic_effort")
|
||||
config["output_language"] = selections.get("output_language", "English")
|
||||
# --checkpoint/--no-checkpoint overrides only when explicitly given; omitting
|
||||
# the flag preserves TRADINGAGENTS_CHECKPOINT_ENABLED / the default (#976).
|
||||
if checkpoint is not None:
|
||||
config["checkpoint_enabled"] = checkpoint
|
||||
return config
|
||||
|
||||
|
||||
def run_analysis(checkpoint: bool | None = None, portfolio=None):
|
||||
# First get all user selections
|
||||
selections = get_user_selections()
|
||||
|
||||
config = _build_run_config(selections, checkpoint)
|
||||
|
||||
stats_handler = StatsCallbackHandler()
|
||||
|
||||
# Normalize analyst selection to predefined order (selection is a 'set', order is fixed)
|
||||
selected_set = {analyst.value for analyst in selections["analysts"]}
|
||||
selected_analyst_keys = [a for a in ANALYST_ORDER if a in selected_set]
|
||||
analyst_execution_plan = build_analyst_execution_plan(selected_analyst_keys)
|
||||
analyst_wall_time_tracker = AnalystWallTimeTracker(analyst_execution_plan)
|
||||
|
||||
graph = TradingAgentsGraph(
|
||||
selected_analyst_keys,
|
||||
config=config,
|
||||
debug=True,
|
||||
callbacks=[stats_handler],
|
||||
)
|
||||
|
||||
message_buffer.init_for_analysis(selected_analyst_keys)
|
||||
|
||||
# Track start time for elapsed display
|
||||
start_time = time.time()
|
||||
|
||||
results_dir = _run_directory(config, selections["ticker"], selections["analysis_date"])
|
||||
results_dir.mkdir(parents=True, exist_ok=True)
|
||||
report_dir = results_dir / "reports"
|
||||
report_dir.mkdir(parents=True, exist_ok=True)
|
||||
log_file = results_dir / "message_tool.log"
|
||||
log_file.touch(exist_ok=True)
|
||||
|
||||
def save_message_decorator(obj, func_name):
|
||||
func = getattr(obj, func_name)
|
||||
|
||||
@wraps(func)
|
||||
def wrapper(*args, **kwargs):
|
||||
func(*args, **kwargs)
|
||||
timestamp, message_type, content = obj.messages[-1]
|
||||
content = content.replace("\n", " ") # Replace newlines with spaces
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(f"{timestamp} [{message_type}] {content}\n")
|
||||
return wrapper
|
||||
|
||||
def save_tool_call_decorator(obj, func_name):
|
||||
func = getattr(obj, func_name)
|
||||
|
||||
@wraps(func)
|
||||
def wrapper(*args, **kwargs):
|
||||
func(*args, **kwargs)
|
||||
timestamp, tool_name, args = obj.tool_calls[-1]
|
||||
args_str = ", ".join(f"{k}={v}" for k, v in args.items())
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(f"{timestamp} [Tool Call] {tool_name}({args_str})\n")
|
||||
return wrapper
|
||||
|
||||
def save_report_section_decorator(obj, func_name):
|
||||
func = getattr(obj, func_name)
|
||||
|
||||
@wraps(func)
|
||||
def wrapper(section_name, content):
|
||||
func(section_name, content)
|
||||
if section_name in obj.report_sections and obj.report_sections[section_name] is not None:
|
||||
content = obj.report_sections[section_name]
|
||||
if content:
|
||||
file_name = f"{section_name}.md"
|
||||
text = "\n".join(str(item) for item in content) if isinstance(content, list) else content
|
||||
with open(report_dir / file_name, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
return wrapper
|
||||
|
||||
message_buffer.add_message = save_message_decorator(message_buffer, "add_message")
|
||||
message_buffer.add_tool_call = save_tool_call_decorator(message_buffer, "add_tool_call")
|
||||
message_buffer.update_report_section = save_report_section_decorator(message_buffer, "update_report_section")
|
||||
|
||||
layout = create_layout()
|
||||
|
||||
# The alternate screen keeps a layout taller than the window from redrawing
|
||||
# by scrolling; the final report prints after this block, on the normal screen.
|
||||
with Live(layout, refresh_per_second=4, screen=True):
|
||||
# Initial display
|
||||
update_display(layout, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
message_buffer.add_message("System", f"Selected ticker: {selections['ticker']}")
|
||||
if selections["asset_type"] != "stock":
|
||||
message_buffer.add_message("System", f"Detected asset type: {selections['asset_type']}")
|
||||
message_buffer.add_message(
|
||||
"System", f"Analysis date: {selections['analysis_date']}"
|
||||
)
|
||||
message_buffer.add_message(
|
||||
"System",
|
||||
f"Selected analysts: {', '.join(analyst.value for analyst in selections['analysts'])}",
|
||||
)
|
||||
update_display(layout, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
first_analyst = analyst_execution_plan.specs[0].agent_node
|
||||
message_buffer.update_agent_status(first_analyst, "in_progress")
|
||||
analyst_wall_time_tracker.mark_started(selected_analyst_keys[0])
|
||||
update_display(layout, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
spinner_text = (
|
||||
f"Analyzing {selections['ticker']} on {selections['analysis_date']}..."
|
||||
)
|
||||
update_display(layout, spinner_text, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
# The same initial state propagate() builds: settled decision log, past
|
||||
# context and resolved instrument identity.
|
||||
init_agent_state = graph.create_run_state(
|
||||
selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio
|
||||
)
|
||||
# Pass callbacks to graph config for tool execution tracking
|
||||
# (LLM tracking is handled separately via LLM constructor)
|
||||
args = graph.propagator.get_graph_args(callbacks=[stats_handler])
|
||||
|
||||
# Recompile with a checkpointer and inject the thread_id so --checkpoint
|
||||
# actually saves and resumes on the CLI path (#1249); a no-op when
|
||||
# checkpointing is disabled. Torn down in the finally below.
|
||||
checkpoint_tid = graph.begin_checkpoint(
|
||||
selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio
|
||||
)
|
||||
if checkpoint_tid is not None:
|
||||
args.setdefault("config", {}).setdefault("configurable", {})["thread_id"] = checkpoint_tid
|
||||
_announce_checkpoint_state(graph, selections["ticker"], selections["analysis_date"])
|
||||
|
||||
# Stream the analysis. On resume, feed None so LangGraph continues the
|
||||
# interrupted run instead of re-appending the initial state (#1249); the
|
||||
# try/finally tears the checkpointer down even if the stream raises.
|
||||
trace = []
|
||||
try:
|
||||
for chunk in graph.graph.stream(graph.checkpoint_input(init_agent_state), **args):
|
||||
for message in chunk.get("messages", []):
|
||||
msg_id = getattr(message, "id", None)
|
||||
if msg_id is not None:
|
||||
if msg_id in message_buffer._processed_message_ids:
|
||||
continue
|
||||
message_buffer._processed_message_ids.add(msg_id)
|
||||
|
||||
msg_type, content = classify_message_type(message)
|
||||
if content and content.strip():
|
||||
message_buffer.add_message(msg_type, content)
|
||||
|
||||
if hasattr(message, "tool_calls") and message.tool_calls:
|
||||
for tool_call in message.tool_calls:
|
||||
if isinstance(tool_call, dict):
|
||||
message_buffer.add_tool_call(tool_call["name"], tool_call["args"])
|
||||
else:
|
||||
message_buffer.add_tool_call(tool_call.name, tool_call.args)
|
||||
|
||||
update_analyst_statuses(
|
||||
message_buffer,
|
||||
chunk,
|
||||
wall_time_tracker=analyst_wall_time_tracker,
|
||||
)
|
||||
|
||||
# Research Team - Handle Investment Debate State
|
||||
if chunk.get("investment_debate_state"):
|
||||
debate_state = chunk["investment_debate_state"]
|
||||
bull_hist = debate_state.get("bull_history", "").strip()
|
||||
bear_hist = debate_state.get("bear_history", "").strip()
|
||||
judge = debate_state.get("judge_decision", "").strip()
|
||||
|
||||
# Only update status when there's actual content
|
||||
if bull_hist or bear_hist:
|
||||
update_research_team_status("in_progress")
|
||||
if bull_hist:
|
||||
message_buffer.update_report_section(
|
||||
"investment_plan", f"### Bull Researcher Analysis\n{bull_hist}"
|
||||
)
|
||||
if bear_hist:
|
||||
message_buffer.update_report_section(
|
||||
"investment_plan", f"### Bear Researcher Analysis\n{bear_hist}"
|
||||
)
|
||||
if judge:
|
||||
message_buffer.update_report_section(
|
||||
"investment_plan", f"### Research Manager Decision\n{judge}"
|
||||
)
|
||||
update_research_team_status("completed")
|
||||
message_buffer.update_agent_status("Trader", "in_progress")
|
||||
|
||||
# Trading Team
|
||||
if chunk.get("trader_investment_plan"):
|
||||
message_buffer.update_report_section(
|
||||
"trader_investment_plan", chunk["trader_investment_plan"]
|
||||
)
|
||||
if message_buffer.agent_status.get("Trader") != "completed":
|
||||
message_buffer.update_agent_status("Trader", "completed")
|
||||
message_buffer.update_agent_status("Aggressive Analyst", "in_progress")
|
||||
|
||||
# Risk Management Team - Handle Risk Debate State
|
||||
if chunk.get("risk_debate_state"):
|
||||
risk_state = chunk["risk_debate_state"]
|
||||
agg_hist = risk_state.get("aggressive_history", "").strip()
|
||||
con_hist = risk_state.get("conservative_history", "").strip()
|
||||
neu_hist = risk_state.get("neutral_history", "").strip()
|
||||
judge = risk_state.get("judge_decision", "").strip()
|
||||
|
||||
if agg_hist:
|
||||
if message_buffer.agent_status.get("Aggressive Analyst") != "completed":
|
||||
message_buffer.update_agent_status("Aggressive Analyst", "in_progress")
|
||||
message_buffer.update_report_section(
|
||||
"final_trade_decision", f"### Aggressive Analyst Analysis\n{agg_hist}"
|
||||
)
|
||||
if con_hist:
|
||||
if message_buffer.agent_status.get("Conservative Analyst") != "completed":
|
||||
message_buffer.update_agent_status("Conservative Analyst", "in_progress")
|
||||
message_buffer.update_report_section(
|
||||
"final_trade_decision", f"### Conservative Analyst Analysis\n{con_hist}"
|
||||
)
|
||||
if neu_hist:
|
||||
if message_buffer.agent_status.get("Neutral Analyst") != "completed":
|
||||
message_buffer.update_agent_status("Neutral Analyst", "in_progress")
|
||||
message_buffer.update_report_section(
|
||||
"final_trade_decision", f"### Neutral Analyst Analysis\n{neu_hist}"
|
||||
)
|
||||
if judge and message_buffer.agent_status.get("Portfolio Manager") != "completed":
|
||||
message_buffer.update_agent_status("Portfolio Manager", "in_progress")
|
||||
message_buffer.update_report_section(
|
||||
"final_trade_decision", f"### Portfolio Manager Decision\n{judge}"
|
||||
)
|
||||
message_buffer.update_agent_status("Aggressive Analyst", "completed")
|
||||
message_buffer.update_agent_status("Conservative Analyst", "completed")
|
||||
message_buffer.update_agent_status("Neutral Analyst", "completed")
|
||||
message_buffer.update_agent_status("Portfolio Manager", "completed")
|
||||
|
||||
update_display(layout, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
trace.append(chunk)
|
||||
|
||||
# Streamed chunks are per-node deltas, not full state. Merge them
|
||||
# so every report field populated across the run is present.
|
||||
final_state = {}
|
||||
for chunk in trace:
|
||||
final_state.update(chunk)
|
||||
|
||||
# Clean run: log the decision, then drop this run's checkpoint so a
|
||||
# later run starts fresh. A mid-stream failure skips both, keeping
|
||||
# the checkpoint for resume.
|
||||
graph.record_decision(selections["ticker"], selections["analysis_date"], final_state)
|
||||
graph.clear_checkpoint_on_success(
|
||||
selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio
|
||||
)
|
||||
finally:
|
||||
# Always restore the plain uncheckpointed graph, even on failure.
|
||||
graph.end_checkpoint()
|
||||
|
||||
for agent in message_buffer.agent_status:
|
||||
message_buffer.update_agent_status(agent, "completed")
|
||||
|
||||
message_buffer.add_message(
|
||||
"System", f"Completed analysis for {selections['analysis_date']}"
|
||||
)
|
||||
message_buffer.add_message("System", analyst_wall_time_tracker.format_summary())
|
||||
|
||||
for section in message_buffer.report_sections:
|
||||
if section in final_state:
|
||||
message_buffer.update_report_section(section, final_state[section])
|
||||
|
||||
update_display(layout, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
# Post-analysis prompts (outside Live context for clean interaction)
|
||||
console.print("\n[bold cyan]Analysis Complete![/bold cyan]\n")
|
||||
|
||||
# A decision nobody can read is not a position. Say so here rather than
|
||||
# leaving the run to look like a normal result.
|
||||
if is_review(graph.process_signal(final_state.get("final_trade_decision", ""))):
|
||||
console.print(
|
||||
"[yellow]No rating could be read from the final decision, so this run "
|
||||
"is recorded for review rather than as a position. Re-run, or read the "
|
||||
"decision text below and judge it yourself.[/yellow]\n"
|
||||
)
|
||||
console.print(f"[dim]{analyst_wall_time_tracker.format_summary()}[/dim]")
|
||||
|
||||
# Prompt to save report
|
||||
save_choice = typer.prompt("Save report?", default="Y").strip().upper()
|
||||
if save_choice in ("Y", "YES", ""):
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
# Under results_dir, not the working directory: in Docker the working
|
||||
# directory is inside the container and the report goes with it, while
|
||||
# results_dir is the mounted volume the rest of the run already writes to.
|
||||
default_path = (Path(config["results_dir"]) / "reports"
|
||||
/ f"{safe_ticker_component(selections['ticker'])}_{timestamp}")
|
||||
save_path_str = typer.prompt(
|
||||
"Save path (press Enter for default)",
|
||||
default=str(default_path)
|
||||
).strip()
|
||||
save_path = Path(save_path_str)
|
||||
try:
|
||||
report_file = write_report_tree(final_state, selections["ticker"], save_path)
|
||||
console.print(f"\n[green]✓ Report saved to:[/green] {save_path.resolve()}")
|
||||
console.print(f" [dim]Complete report:[/dim] {report_file.name}")
|
||||
except Exception as e:
|
||||
console.print(f"[red]Error saving report: {e}[/red]")
|
||||
|
||||
# Prompt to display full report
|
||||
display_choice = typer.prompt("\nDisplay full report on screen?", default="Y").strip().upper()
|
||||
if display_choice in ("Y", "YES", ""):
|
||||
display_complete_report(final_state)
|
||||
@@ -0,0 +1,315 @@
|
||||
"""The interactive choices for a run: ticker, date, analysts, depth, provider and models."""
|
||||
|
||||
import datetime
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import typer
|
||||
from rich.align import Align
|
||||
from rich.panel import Panel
|
||||
|
||||
from cli.announcements import display_announcements, fetch_announcements
|
||||
from cli.display import (
|
||||
console,
|
||||
)
|
||||
from cli.prefs import load_last_run, sanitize, save_last_run
|
||||
from cli.prompts import (
|
||||
ask_anthropic_effort,
|
||||
ask_gemini_thinking_config,
|
||||
ask_glm_region,
|
||||
ask_minimax_region,
|
||||
ask_openai_reasoning_effort,
|
||||
ask_output_language,
|
||||
ask_qwen_region,
|
||||
confirm_ollama_endpoint,
|
||||
detect_asset_type,
|
||||
ensure_api_key,
|
||||
get_ticker,
|
||||
prompt_openai_compatible_url,
|
||||
resolve_backend_url,
|
||||
select_analysts,
|
||||
select_deep_thinking_agent,
|
||||
select_llm_provider,
|
||||
select_research_depth,
|
||||
select_shallow_thinking_agent,
|
||||
)
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
|
||||
|
||||
def get_user_selections():
|
||||
"""Ask for the run's settings, offering the previous run's answers."""
|
||||
selections = _prompt_selections(load_last_run())
|
||||
save_last_run(selections)
|
||||
return selections
|
||||
|
||||
|
||||
def _prompt_selections(prefs):
|
||||
"""Walk the selection steps. ``prefs`` prefills, the environment skips."""
|
||||
with open(Path(__file__).parent / "static" / "welcome.txt", encoding="utf-8") as f:
|
||||
welcome_ascii = f.read()
|
||||
|
||||
welcome_content = f"{welcome_ascii}\n"
|
||||
welcome_content += "[bold green]TradingAgents: Multi-Agents LLM Financial Trading Framework - CLI[/bold green]\n\n"
|
||||
welcome_content += "[bold]Workflow Steps:[/bold]\n"
|
||||
welcome_content += "I. Analyst Team → II. Research Team → III. Trader → IV. Risk Management → V. Portfolio Management\n\n"
|
||||
welcome_content += (
|
||||
"[dim]Built by [Tauric Research](https://github.com/TauricResearch)[/dim]"
|
||||
)
|
||||
|
||||
welcome_box = Panel(
|
||||
welcome_content,
|
||||
border_style="green",
|
||||
padding=(1, 2),
|
||||
title="Welcome to TradingAgents",
|
||||
subtitle="Multi-Agents LLM Financial Trading Framework",
|
||||
)
|
||||
console.print(Align.center(welcome_box))
|
||||
console.print()
|
||||
console.print() # Add vertical space before announcements
|
||||
|
||||
# Fetch and display announcements (silent on failure)
|
||||
announcements = fetch_announcements()
|
||||
display_announcements(console, announcements)
|
||||
|
||||
def create_question_box(title, prompt, default=None):
|
||||
box_content = f"[bold]{title}[/bold]\n"
|
||||
box_content += f"[dim]{prompt}[/dim]"
|
||||
if default:
|
||||
box_content += f"\n[dim]Default: {default}[/dim]"
|
||||
return Panel(box_content, border_style="blue", padding=(1, 2))
|
||||
|
||||
def thinking_value_or_prompt(env_var, config_key, label, box_title, box_body, prompt_fn):
|
||||
"""Return the env-configured reasoning/thinking value, or prompt for it.
|
||||
|
||||
When ``env_var`` is set the interactive choice is skipped and the value
|
||||
the env overlay placed on DEFAULT_CONFIG is used — mirroring the
|
||||
env-precedence rule applied to the other selection steps.
|
||||
"""
|
||||
if os.environ.get(env_var):
|
||||
value = DEFAULT_CONFIG[config_key]
|
||||
console.print(f"[green]✓ {label} from environment:[/green] {value}")
|
||||
return value
|
||||
console.print(create_question_box(box_title, box_body))
|
||||
return prompt_fn()
|
||||
|
||||
# Step 1: Ticker symbol
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 1: Ticker Symbol",
|
||||
"Enter the ticker, with exchange suffix when needed (e.g. SPY, 0700.HK, BTC-USD)",
|
||||
"SPY",
|
||||
)
|
||||
)
|
||||
selected_ticker = get_ticker()
|
||||
asset_type = detect_asset_type(selected_ticker)
|
||||
# Only announce when it's not the default stock path, to avoid printing
|
||||
# "stock" on every run.
|
||||
if asset_type.value != "stock":
|
||||
console.print(
|
||||
f"[green]Detected asset type:[/green] {asset_type.value}"
|
||||
)
|
||||
|
||||
# Step 2: Analysis date
|
||||
default_date = datetime.datetime.now().strftime("%Y-%m-%d")
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 2: Analysis Date",
|
||||
"Enter the analysis date (YYYY-MM-DD)",
|
||||
default_date,
|
||||
)
|
||||
)
|
||||
analysis_date = get_analysis_date()
|
||||
|
||||
# Step 3: Output language (skipped when set via TRADINGAGENTS_OUTPUT_LANGUAGE)
|
||||
if os.environ.get("TRADINGAGENTS_OUTPUT_LANGUAGE"):
|
||||
output_language = DEFAULT_CONFIG["output_language"]
|
||||
console.print(
|
||||
f"[green]✓ Output language from environment:[/green] {output_language}"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 3: Output Language",
|
||||
"Select the language for analyst reports and final decision"
|
||||
)
|
||||
)
|
||||
output_language = ask_output_language(prefs.get("output_language"))
|
||||
|
||||
# Step 4: Select analysts
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 4: Analysts Team", "Select your LLM analyst agents for the analysis"
|
||||
)
|
||||
)
|
||||
prefs = sanitize(prefs, asset_type.value)
|
||||
selected_analysts = select_analysts(asset_type, prefs.get("analysts"))
|
||||
console.print(
|
||||
f"[green]Selected analysts:[/green] {', '.join(analyst.value for analyst in selected_analysts)}"
|
||||
)
|
||||
|
||||
# Step 5: Research depth (skipped when both round counts are set via env).
|
||||
# Research depth maps to the debate + risk round counts; when both are
|
||||
# supplied through TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS we keep
|
||||
# the run non-interactive and honor the env values (#977).
|
||||
depth_from_env = bool(os.environ.get("TRADINGAGENTS_MAX_DEBATE_ROUNDS")) and bool(
|
||||
os.environ.get("TRADINGAGENTS_MAX_RISK_ROUNDS")
|
||||
)
|
||||
if depth_from_env:
|
||||
selected_research_depth = DEFAULT_CONFIG["max_debate_rounds"]
|
||||
console.print(
|
||||
f"[green]✓ Research depth from environment:[/green] "
|
||||
f"{DEFAULT_CONFIG['max_debate_rounds']} debate / "
|
||||
f"{DEFAULT_CONFIG['max_risk_discuss_rounds']} risk rounds"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 5: Research Depth", "Select your research depth level"
|
||||
)
|
||||
)
|
||||
selected_research_depth = select_research_depth(prefs.get("research_depth"))
|
||||
|
||||
# Step 6: LLM Provider (skipped when set via TRADINGAGENTS_LLM_PROVIDER).
|
||||
# The backend URL comes from TRADINGAGENTS_LLM_BACKEND_URL when set,
|
||||
# otherwise the provider's default endpoint — the same value the menu
|
||||
# would have picked.
|
||||
provider_from_env = bool(os.environ.get("TRADINGAGENTS_LLM_PROVIDER"))
|
||||
if provider_from_env:
|
||||
selected_llm_provider = DEFAULT_CONFIG["llm_provider"].lower()
|
||||
backend_url = resolve_backend_url(
|
||||
selected_llm_provider, env_url=DEFAULT_CONFIG["backend_url"]
|
||||
)
|
||||
console.print(f"[green]✓ LLM provider from environment:[/green] {selected_llm_provider}")
|
||||
console.print(f"[green]✓ Backend URL:[/green] {backend_url}")
|
||||
# Still confirm/persist the API key so the run doesn't fail later.
|
||||
ensure_api_key(selected_llm_provider)
|
||||
else:
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 6: LLM Provider", "Select your LLM provider"
|
||||
)
|
||||
)
|
||||
selected_llm_provider, backend_url = select_llm_provider(prefs.get("llm_provider"))
|
||||
|
||||
# Providers with regional endpoints prompt for the region as a secondary
|
||||
# step so the main dropdown stays clean (mainland China and international
|
||||
# accounts cannot share API keys).
|
||||
if selected_llm_provider == "qwen":
|
||||
selected_llm_provider, backend_url = ask_qwen_region()
|
||||
elif selected_llm_provider == "minimax":
|
||||
selected_llm_provider, backend_url = ask_minimax_region()
|
||||
elif selected_llm_provider == "glm":
|
||||
selected_llm_provider, backend_url = ask_glm_region()
|
||||
|
||||
# Honor an explicit env backend URL even when the provider was chosen
|
||||
# interactively, so it isn't overwritten by the menu default (#978).
|
||||
backend_url = resolve_backend_url(
|
||||
selected_llm_provider, backend_url, env_url=DEFAULT_CONFIG["backend_url"]
|
||||
)
|
||||
|
||||
# The generic OpenAI-compatible endpoint has no default; ask for it if
|
||||
# neither the menu nor the environment supplied one.
|
||||
if selected_llm_provider == "openai_compatible" and not backend_url:
|
||||
remembered_url = (prefs.get("backend_url")
|
||||
if prefs.get("llm_provider") == selected_llm_provider else None)
|
||||
backend_url = prompt_openai_compatible_url(remembered_url)
|
||||
|
||||
# For Ollama, surface the resolved endpoint (OLLAMA_BASE_URL vs default)
|
||||
# before model selection so it's obvious where we're connecting.
|
||||
if selected_llm_provider == "ollama":
|
||||
confirm_ollama_endpoint(backend_url)
|
||||
|
||||
# Confirm the provider's API key is present; prompt the user to paste
|
||||
# one and persist it to .env if it's missing, so the analysis run
|
||||
# doesn't fail later at the first API call.
|
||||
ensure_api_key(selected_llm_provider)
|
||||
|
||||
# Step 7: Thinking agents (skipped when either model is set via environment)
|
||||
if os.environ.get("TRADINGAGENTS_QUICK_THINK_LLM") or os.environ.get("TRADINGAGENTS_DEEP_THINK_LLM"):
|
||||
selected_shallow_thinker = DEFAULT_CONFIG["quick_think_llm"]
|
||||
selected_deep_thinker = DEFAULT_CONFIG["deep_think_llm"]
|
||||
console.print(
|
||||
f"[green]✓ Thinking agents from environment:[/green] "
|
||||
f"quick={selected_shallow_thinker}, deep={selected_deep_thinker}"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
create_question_box(
|
||||
"Step 7: Thinking Agents", "Select your thinking agents for analysis"
|
||||
)
|
||||
)
|
||||
remembered = prefs if prefs.get("llm_provider") == selected_llm_provider else {}
|
||||
selected_shallow_thinker = select_shallow_thinking_agent(
|
||||
selected_llm_provider, remembered.get("quick_think_llm")
|
||||
)
|
||||
selected_deep_thinker = select_deep_thinking_agent(
|
||||
selected_llm_provider, remembered.get("deep_think_llm")
|
||||
)
|
||||
|
||||
# Step 8: Provider-specific reasoning/thinking configuration. Each knob is
|
||||
# settable via its TRADINGAGENTS_* env var; when that var is set (or the
|
||||
# provider itself came from env) the prompt is skipped and the configured
|
||||
# value is used — same env-precedence rule as the steps above. None = each
|
||||
# provider's own default.
|
||||
thinking_level = None
|
||||
reasoning_effort = None
|
||||
anthropic_effort = None
|
||||
|
||||
provider_lower = selected_llm_provider.lower()
|
||||
if provider_from_env:
|
||||
thinking_level = DEFAULT_CONFIG["google_thinking_level"]
|
||||
reasoning_effort = DEFAULT_CONFIG["openai_reasoning_effort"]
|
||||
anthropic_effort = DEFAULT_CONFIG["anthropic_effort"]
|
||||
elif provider_lower == "google":
|
||||
thinking_level = thinking_value_or_prompt(
|
||||
"TRADINGAGENTS_GOOGLE_THINKING_LEVEL", "google_thinking_level",
|
||||
"Gemini thinking mode", "Step 8: Thinking Mode",
|
||||
"Configure Gemini thinking mode", ask_gemini_thinking_config,
|
||||
)
|
||||
elif provider_lower == "openai":
|
||||
reasoning_effort = thinking_value_or_prompt(
|
||||
"TRADINGAGENTS_OPENAI_REASONING_EFFORT", "openai_reasoning_effort",
|
||||
"Reasoning effort", "Step 8: Reasoning Effort",
|
||||
"Configure OpenAI reasoning effort level", ask_openai_reasoning_effort,
|
||||
)
|
||||
elif provider_lower == "anthropic":
|
||||
anthropic_effort = thinking_value_or_prompt(
|
||||
"TRADINGAGENTS_ANTHROPIC_EFFORT", "anthropic_effort",
|
||||
"Claude effort", "Step 8: Effort Level",
|
||||
"Configure Claude effort level", ask_anthropic_effort,
|
||||
)
|
||||
|
||||
return {
|
||||
"ticker": selected_ticker,
|
||||
"asset_type": asset_type.value,
|
||||
"analysis_date": analysis_date,
|
||||
"analysts": selected_analysts,
|
||||
"research_depth": selected_research_depth,
|
||||
"llm_provider": selected_llm_provider.lower(),
|
||||
"backend_url": backend_url,
|
||||
"quick_think_llm": selected_shallow_thinker,
|
||||
"deep_think_llm": selected_deep_thinker,
|
||||
"google_thinking_level": thinking_level,
|
||||
"openai_reasoning_effort": reasoning_effort,
|
||||
"anthropic_effort": anthropic_effort,
|
||||
"output_language": output_language,
|
||||
}
|
||||
|
||||
|
||||
def get_analysis_date():
|
||||
"""Get the analysis date from user input."""
|
||||
while True:
|
||||
date_str = typer.prompt(
|
||||
"", default=datetime.datetime.now().strftime("%Y-%m-%d")
|
||||
)
|
||||
try:
|
||||
# Validate date format and ensure it's not in the future
|
||||
analysis_date = datetime.datetime.strptime(date_str, "%Y-%m-%d")
|
||||
if analysis_date.date() > datetime.datetime.now().date():
|
||||
console.print("[red]Error: Analysis date cannot be in the future[/red]")
|
||||
continue
|
||||
return date_str
|
||||
except ValueError:
|
||||
console.print(
|
||||
"[red]Error: Invalid date format. Please use YYYY-MM-DD[/red]"
|
||||
)
|
||||
@@ -0,0 +1,76 @@
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.callbacks import BaseCallbackHandler
|
||||
from langchain_core.messages import AIMessage
|
||||
from langchain_core.outputs import LLMResult
|
||||
|
||||
|
||||
class StatsCallbackHandler(BaseCallbackHandler):
|
||||
"""Callback handler that tracks LLM calls, tool calls, and token usage."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._lock = threading.Lock()
|
||||
self.llm_calls = 0
|
||||
self.tool_calls = 0
|
||||
self.tokens_in = 0
|
||||
self.tokens_out = 0
|
||||
|
||||
def on_llm_start(
|
||||
self,
|
||||
serialized: dict[str, Any],
|
||||
prompts: list[str],
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Increment LLM call counter when an LLM starts."""
|
||||
with self._lock:
|
||||
self.llm_calls += 1
|
||||
|
||||
def on_chat_model_start(
|
||||
self,
|
||||
serialized: dict[str, Any],
|
||||
messages: list[list[Any]],
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Increment LLM call counter when a chat model starts."""
|
||||
with self._lock:
|
||||
self.llm_calls += 1
|
||||
|
||||
def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:
|
||||
"""Extract token usage from LLM response."""
|
||||
try:
|
||||
generation = response.generations[0][0]
|
||||
except (IndexError, TypeError):
|
||||
return
|
||||
|
||||
usage_metadata = None
|
||||
if hasattr(generation, "message"):
|
||||
message = generation.message
|
||||
if isinstance(message, AIMessage) and hasattr(message, "usage_metadata"):
|
||||
usage_metadata = message.usage_metadata
|
||||
|
||||
if usage_metadata:
|
||||
with self._lock:
|
||||
self.tokens_in += usage_metadata.get("input_tokens", 0)
|
||||
self.tokens_out += usage_metadata.get("output_tokens", 0)
|
||||
|
||||
def on_tool_start(
|
||||
self,
|
||||
serialized: dict[str, Any],
|
||||
input_str: str,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Increment tool call counter when a tool starts."""
|
||||
with self._lock:
|
||||
self.tool_calls += 1
|
||||
|
||||
def get_stats(self) -> dict[str, Any]:
|
||||
"""Return current statistics."""
|
||||
with self._lock:
|
||||
return {
|
||||
"llm_calls": self.llm_calls,
|
||||
"tool_calls": self.tool_calls,
|
||||
"tokens_in": self.tokens_in,
|
||||
"tokens_out": self.tokens_out,
|
||||
}
|
||||
-195
@@ -1,195 +0,0 @@
|
||||
import questionary
|
||||
from typing import List, Optional, Tuple, Dict
|
||||
|
||||
from cli.models import AnalystType
|
||||
|
||||
ANALYST_ORDER = [
|
||||
("Market Analyst", AnalystType.MARKET),
|
||||
("Social Media Analyst", AnalystType.SOCIAL),
|
||||
("News Analyst", AnalystType.NEWS),
|
||||
("Fundamentals Analyst", AnalystType.FUNDAMENTALS),
|
||||
]
|
||||
|
||||
|
||||
def get_ticker() -> str:
|
||||
"""Prompt the user to enter a ticker symbol."""
|
||||
ticker = questionary.text(
|
||||
"Enter the ticker symbol to analyze:",
|
||||
validate=lambda x: len(x.strip()) > 0 or "Please enter a valid ticker symbol.",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("text", "fg:green"),
|
||||
("highlighted", "noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if not ticker:
|
||||
console.print("\n[red]No ticker symbol provided. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return ticker.strip().upper()
|
||||
|
||||
|
||||
def get_analysis_date() -> str:
|
||||
"""Prompt the user to enter a date in YYYY-MM-DD format."""
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
def validate_date(date_str: str) -> bool:
|
||||
if not re.match(r"^\d{4}-\d{2}-\d{2}$", date_str):
|
||||
return False
|
||||
try:
|
||||
datetime.strptime(date_str, "%Y-%m-%d")
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
date = questionary.text(
|
||||
"Enter the analysis date (YYYY-MM-DD):",
|
||||
validate=lambda x: validate_date(x.strip())
|
||||
or "Please enter a valid date in YYYY-MM-DD format.",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("text", "fg:green"),
|
||||
("highlighted", "noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if not date:
|
||||
console.print("\n[red]No date provided. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return date.strip()
|
||||
|
||||
|
||||
def select_analysts() -> List[AnalystType]:
|
||||
"""Select analysts using an interactive checkbox."""
|
||||
choices = questionary.checkbox(
|
||||
"Select Your [Analysts Team]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value) for display, value in ANALYST_ORDER
|
||||
],
|
||||
instruction="\n- Press Space to select/unselect analysts\n- Press 'a' to select/unselect all\n- Press Enter when done",
|
||||
validate=lambda x: len(x) > 0 or "You must select at least one analyst.",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("checkbox-selected", "fg:green"),
|
||||
("selected", "fg:green noinherit"),
|
||||
("highlighted", "noinherit"),
|
||||
("pointer", "noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if not choices:
|
||||
console.print("\n[red]No analysts selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return choices
|
||||
|
||||
|
||||
def select_research_depth() -> int:
|
||||
"""Select research depth using an interactive selection."""
|
||||
|
||||
# Define research depth options with their corresponding values
|
||||
DEPTH_OPTIONS = [
|
||||
("Shallow - Quick research, few debate and strategy discussion rounds", 1),
|
||||
("Medium - Middle ground, moderate debate rounds and strategy discussion", 3),
|
||||
("Deep - Comprehensive research, in depth debate and strategy discussion", 5),
|
||||
]
|
||||
|
||||
choice = questionary.select(
|
||||
"Select Your [Research Depth]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value) for display, value in DEPTH_OPTIONS
|
||||
],
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:yellow noinherit"),
|
||||
("highlighted", "fg:yellow noinherit"),
|
||||
("pointer", "fg:yellow noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print("\n[red]No research depth selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return choice
|
||||
|
||||
|
||||
def select_shallow_thinking_agent() -> str:
|
||||
"""Select shallow thinking llm engine using an interactive selection."""
|
||||
|
||||
# Define shallow thinking llm engine options with their corresponding model names
|
||||
SHALLOW_AGENT_OPTIONS = [
|
||||
("GPT-4o-mini - Fast and efficient for quick tasks", "gpt-4o-mini"),
|
||||
("GPT-4.1-nano - Ultra-lightweight model for basic operations", "gpt-4.1-nano"),
|
||||
("GPT-4.1-mini - Compact model with good performance", "gpt-4.1-mini"),
|
||||
("GPT-4o - Standard model with solid capabilities", "gpt-4o"),
|
||||
]
|
||||
|
||||
choice = questionary.select(
|
||||
"Select Your [Quick-Thinking LLM Engine]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value)
|
||||
for display, value in SHALLOW_AGENT_OPTIONS
|
||||
],
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:magenta noinherit"),
|
||||
("highlighted", "fg:magenta noinherit"),
|
||||
("pointer", "fg:magenta noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print(
|
||||
"\n[red]No shallow thinking llm engine selected. Exiting...[/red]"
|
||||
)
|
||||
exit(1)
|
||||
|
||||
return choice
|
||||
|
||||
|
||||
def select_deep_thinking_agent() -> str:
|
||||
"""Select deep thinking llm engine using an interactive selection."""
|
||||
|
||||
# Define deep thinking llm engine options with their corresponding model names
|
||||
DEEP_AGENT_OPTIONS = [
|
||||
("GPT-4.1-nano - Ultra-lightweight model for basic operations", "gpt-4.1-nano"),
|
||||
("GPT-4.1-mini - Compact model with good performance", "gpt-4.1-mini"),
|
||||
("GPT-4o - Standard model with solid capabilities", "gpt-4o"),
|
||||
("o4-mini - Specialized reasoning model (compact)", "o4-mini"),
|
||||
("o3-mini - Advanced reasoning model (lightweight)", "o3-mini"),
|
||||
("o3 - Full advanced reasoning model", "o3"),
|
||||
("o1 - Premier reasoning and problem-solving model", "o1"),
|
||||
]
|
||||
|
||||
choice = questionary.select(
|
||||
"Select Your [Deep-Thinking LLM Engine]:",
|
||||
choices=[
|
||||
questionary.Choice(display, value=value)
|
||||
for display, value in DEEP_AGENT_OPTIONS
|
||||
],
|
||||
instruction="\n- Use arrow keys to navigate\n- Press Enter to select",
|
||||
style=questionary.Style(
|
||||
[
|
||||
("selected", "fg:magenta noinherit"),
|
||||
("highlighted", "fg:magenta noinherit"),
|
||||
("pointer", "fg:magenta noinherit"),
|
||||
]
|
||||
),
|
||||
).ask()
|
||||
|
||||
if choice is None:
|
||||
console.print("\n[red]No deep thinking llm engine selected. Exiting...[/red]")
|
||||
exit(1)
|
||||
|
||||
return choice
|
||||
@@ -0,0 +1,36 @@
|
||||
services:
|
||||
tradingagents:
|
||||
build: .
|
||||
env_file:
|
||||
- .env
|
||||
volumes:
|
||||
- tradingagents_data:/home/appuser/.tradingagents
|
||||
tty: true
|
||||
stdin_open: true
|
||||
|
||||
ollama:
|
||||
image: ollama/ollama:latest
|
||||
volumes:
|
||||
- ollama_data:/root/.ollama
|
||||
profiles:
|
||||
- ollama
|
||||
|
||||
tradingagents-ollama:
|
||||
build: .
|
||||
env_file:
|
||||
- .env
|
||||
environment:
|
||||
- TRADINGAGENTS_LLM_PROVIDER=ollama
|
||||
- OLLAMA_BASE_URL=http://ollama:11434/v1
|
||||
volumes:
|
||||
- tradingagents_data:/home/appuser/.tradingagents
|
||||
depends_on:
|
||||
- ollama
|
||||
tty: true
|
||||
stdin_open: true
|
||||
profiles:
|
||||
- ollama
|
||||
|
||||
volumes:
|
||||
tradingagents_data:
|
||||
ollama_data:
|
||||
@@ -1,19 +1,18 @@
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
# Create a custom config
|
||||
# DEFAULT_CONFIG already applies TRADINGAGENTS_* env-var overrides
|
||||
# (llm_provider, deep_think_llm, quick_think_llm, backend_url, etc.),
|
||||
# so users can switch models or endpoints purely via .env without
|
||||
# editing this script. Override individual keys here only when you
|
||||
# want a hard-coded value that should ignore the environment.
|
||||
config = DEFAULT_CONFIG.copy()
|
||||
config["deep_think_llm"] = "gpt-4.1-nano" # Use a different model
|
||||
config["quick_think_llm"] = "gpt-4.1-nano" # Use a different model
|
||||
config["max_debate_rounds"] = 1 # Increase debate rounds
|
||||
config["online_tools"] = True # Increase debate rounds
|
||||
|
||||
# Initialize with custom config
|
||||
ta = TradingAgentsGraph(debug=True, config=config)
|
||||
|
||||
# forward propagate
|
||||
_, decision = ta.propagate("NVDA", "2024-05-10")
|
||||
_, decision = ta.propagate("NVDA", "2026-09-01")
|
||||
print(decision)
|
||||
|
||||
# Memorize mistakes and reflect
|
||||
# ta.reflect_and_remember(1000) # parameter is the position returns
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "tradingagents"
|
||||
version = "0.5.1"
|
||||
description = "TradingAgents: Multi-Agents LLM Financial Trading Framework"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"langchain-core>=0.3.81",
|
||||
"langchain-anthropic>=0.3.15",
|
||||
"langchain-google-genai>=4.0.0",
|
||||
"langchain-openai>=0.3.23",
|
||||
"langgraph>=0.4.8",
|
||||
"langgraph-checkpoint-sqlite>=2.0.0",
|
||||
"pandas>=2.3.0",
|
||||
"python-dotenv>=1.0.0",
|
||||
"pytz>=2025.2",
|
||||
"questionary>=2.1.0",
|
||||
"requests>=2.32.4",
|
||||
"rich>=14.0.0",
|
||||
"typer>=0.21.0",
|
||||
"stockstats>=0.6.5",
|
||||
"typing-extensions>=4.14.0",
|
||||
"yfinance>=1.4.1",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
"ruff>=0.15",
|
||||
"pytest>=8.0",
|
||||
"pytest-subtests>=0.13",
|
||||
]
|
||||
# Amazon Bedrock support (AWS SigV4 auth + boto3). Optional so the core install
|
||||
# stays lean: pip install "tradingagents[bedrock]".
|
||||
bedrock = [
|
||||
"langchain-aws>=1.5.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
tradingagents = "cli.main:app"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
include = ["tradingagents*", "cli*"]
|
||||
|
||||
[tool.setuptools.package-data]
|
||||
cli = ["static/*"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
addopts = "-ra --strict-markers"
|
||||
markers = [
|
||||
"unit: fast isolated unit tests",
|
||||
"integration: tests requiring external services",
|
||||
"smoke: quick sanity-check tests",
|
||||
]
|
||||
filterwarnings = [
|
||||
"ignore::DeprecationWarning",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
target-version = "py310"
|
||||
extend-exclude = ["results"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
# Standard "good defaults" rule set (pyflakes + pycodestyle + isort + bugbear +
|
||||
# pyupgrade + comprehensions/simplify). Line length (E501) and layout are owned
|
||||
# by the formatter; whole-repo `ruff format` adoption is deferred until the
|
||||
# open-PR backlog clears, to avoid mass merge conflicts.
|
||||
select = ["E", "W", "F", "I", "B", "UP", "C4", "SIM"]
|
||||
ignore = ["E501"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"**/__init__.py" = ["F401"] # intentional re-exports
|
||||
|
||||
[tool.ruff.lint.isort]
|
||||
# Keep multiple aliased names from one module in a single combined import block
|
||||
# (e.g. the vendor imports in router.py) instead of one statement per name.
|
||||
combine-as-imports = true
|
||||
+1
-24
@@ -1,24 +1 @@
|
||||
typing-extensions
|
||||
langchain-openai
|
||||
langchain-experimental
|
||||
pandas
|
||||
yfinance
|
||||
praw
|
||||
feedparser
|
||||
stockstats
|
||||
eodhd
|
||||
langgraph
|
||||
chromadb
|
||||
setuptools
|
||||
backtrader
|
||||
akshare
|
||||
tushare
|
||||
finnhub-python
|
||||
parsel
|
||||
requests
|
||||
tqdm
|
||||
pytz
|
||||
redis
|
||||
chainlit
|
||||
rich
|
||||
questionary
|
||||
.
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
"""
|
||||
Setup script for the TradingAgents package.
|
||||
"""
|
||||
|
||||
from setuptools import setup, find_packages
|
||||
|
||||
setup(
|
||||
name="tradingagents",
|
||||
version="0.1.0",
|
||||
description="Multi-Agents LLM Financial Trading Framework",
|
||||
author="TradingAgents Team",
|
||||
author_email="yijia.xiao@cs.ucla.edu",
|
||||
url="https://github.com/TauricResearch",
|
||||
packages=find_packages(),
|
||||
install_requires=[
|
||||
"langchain>=0.1.0",
|
||||
"langchain-openai>=0.0.2",
|
||||
"langchain-experimental>=0.0.40",
|
||||
"langgraph>=0.0.20",
|
||||
"numpy>=1.24.0",
|
||||
"pandas>=2.0.0",
|
||||
"praw>=7.7.0",
|
||||
"stockstats>=0.5.4",
|
||||
"yfinance>=0.2.31",
|
||||
"typer>=0.9.0",
|
||||
"rich>=13.0.0",
|
||||
"questionary>=2.0.1",
|
||||
],
|
||||
python_requires=">=3.10",
|
||||
entry_points={
|
||||
"console_scripts": [
|
||||
"tradingagents=cli.main:app",
|
||||
],
|
||||
},
|
||||
classifiers=[
|
||||
"Development Status :: 3 - Alpha",
|
||||
"Intended Audience :: Financial and Trading Industry",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Topic :: Office/Business :: Financial :: Investment",
|
||||
],
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Shared pytest fixtures that prevent CI hangs when API keys are absent."""
|
||||
|
||||
import os
|
||||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _blank_settings_overlay():
|
||||
"""Blank every TRADINGAGENTS_* setting before the package is imported.
|
||||
|
||||
The package loads .env on import and folds these variables into
|
||||
DEFAULT_CONFIG, so a contributor's own settings would become the defaults
|
||||
the suite asserts on. A blank value is still present, so load_dotenv leaves
|
||||
it alone, and the overlay reads it as unset. Tests of the overlay set their own.
|
||||
"""
|
||||
from dotenv import dotenv_values, find_dotenv
|
||||
|
||||
names = set(os.environ)
|
||||
for filename in (".env", ".env.enterprise"):
|
||||
names |= set(dotenv_values(find_dotenv(filename, usecwd=True)))
|
||||
for name in names:
|
||||
if name.startswith("TRADINGAGENTS_"):
|
||||
os.environ[name] = ""
|
||||
|
||||
|
||||
_blank_settings_overlay()
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
for marker in ("unit", "integration", "smoke"):
|
||||
config.addinivalue_line("markers", f"{marker}: {marker}-level tests")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_network(request, monkeypatch):
|
||||
"""Tests do not reach the network; one that must is marked integration."""
|
||||
if request.node.get_closest_marker("integration"):
|
||||
return
|
||||
|
||||
def refuse(self, address):
|
||||
raise OSError(f"test tried to reach the network: {address}")
|
||||
|
||||
monkeypatch.setattr(socket.socket, "connect", refuse)
|
||||
monkeypatch.setattr(socket.socket, "connect_ex", refuse)
|
||||
|
||||
|
||||
_API_KEY_ENV_VARS = (
|
||||
"OPENAI_API_KEY",
|
||||
"GOOGLE_API_KEY",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"XAI_API_KEY",
|
||||
"DEEPSEEK_API_KEY",
|
||||
"DASHSCOPE_API_KEY",
|
||||
"DASHSCOPE_CN_API_KEY",
|
||||
"ZHIPU_API_KEY",
|
||||
"ZHIPU_CN_API_KEY",
|
||||
"MINIMAX_API_KEY",
|
||||
"MINIMAX_CN_API_KEY",
|
||||
"OPENROUTER_API_KEY",
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
"ALPHA_VANTAGE_API_KEY",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _dummy_api_keys(monkeypatch):
|
||||
for env_var in _API_KEY_ENV_VARS:
|
||||
# `or` not a .get default: an env var present but empty (e.g. a key left
|
||||
# blank in a .env copied from .env.example) must still get the placeholder.
|
||||
monkeypatch.setenv(env_var, os.environ.get(env_var) or "placeholder")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_config():
|
||||
"""Reset the global dataflows config before and after each test.
|
||||
|
||||
``set_config`` merges (it never clears keys absent from the override), so a
|
||||
test that sets e.g. ``tool_vendors`` would otherwise leak into later tests
|
||||
and make routing behavior order-dependent. Replace the global outright so
|
||||
every test starts from a clean DEFAULT_CONFIG.
|
||||
"""
|
||||
import copy
|
||||
|
||||
import tradingagents.dataflows.config as config_module
|
||||
import tradingagents.default_config as default_config
|
||||
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
yield
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
@@ -0,0 +1,213 @@
|
||||
"""Alpha Vantage request hardening.
|
||||
|
||||
Regressions for #990 (no request timeout -> can hang), #991 (invalid-key
|
||||
responses mislabeled as rate limits and silently treated as transient), and
|
||||
#1115 (fundamentals look-ahead filter never ran because the payload is a JSON
|
||||
string, not a dict), and the date trim that keeps post-end_date bars out of a
|
||||
historical run.
|
||||
"""
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.dataflows.net as net
|
||||
import tradingagents.dataflows.vendors.alpha_vantage.common as av
|
||||
import tradingagents.dataflows.vendors.alpha_vantage.fundamentals as avf
|
||||
import tradingagents.dataflows.vendors.alpha_vantage.stock as avs
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
status_code = 200
|
||||
|
||||
def __init__(self, text):
|
||||
self.text = text
|
||||
|
||||
def raise_for_status(self):
|
||||
pass
|
||||
|
||||
|
||||
def _patched_get(body, capture=None):
|
||||
def fake_get(url, params=None, **kwargs):
|
||||
if capture is not None:
|
||||
capture.update(kwargs)
|
||||
return _FakeResponse(body)
|
||||
return fake_get
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_passes_timeout(monkeypatch):
|
||||
captured = {}
|
||||
monkeypatch.setattr(net.requests, "get", _patched_get("Date,Close\n2025-01-02,1.0", captured))
|
||||
av._make_api_request("TIME_SERIES_DAILY", {"symbol": "AAPL"})
|
||||
assert captured.get("timeout") == av.REQUEST_TIMEOUT # #990
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_rate_limit_detected(monkeypatch):
|
||||
body = '{"Information": "Our standard API rate limit is 25 requests per day. ... your API key ..."}'
|
||||
monkeypatch.setattr(net.requests, "get", _patched_get(body))
|
||||
with pytest.raises(av.AlphaVantageRateLimitError):
|
||||
av._make_api_request("TIME_SERIES_DAILY", {"symbol": "AAPL"})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_invalid_key_not_mislabeled_as_rate_limit(monkeypatch):
|
||||
# AV's invalid-key notice mentions "API key"; it must NOT be treated as a
|
||||
# (transient) rate limit, but surface as a real configuration error (#991).
|
||||
body = ('{"Information": "the parameter apikey is invalid or missing. '
|
||||
'Please claim your free API key on (https://www.alphavantage.co/support/#api-key)."}')
|
||||
monkeypatch.setattr(net.requests, "get", _patched_get(body))
|
||||
with pytest.raises(av.AlphaVantageNotConfiguredError):
|
||||
av._make_api_request("TIME_SERIES_DAILY", {"symbol": "AAPL"})
|
||||
with pytest.raises(av.AlphaVantageRateLimitError): # sanity: rate-limit path still distinct
|
||||
monkeypatch.setattr(net.requests, "get", _patched_get('{"Note": "API call frequency is 5 calls per minute."}'))
|
||||
av._make_api_request("TIME_SERIES_DAILY", {"symbol": "AAPL"})
|
||||
|
||||
|
||||
_FUNDAMENTALS_JSON = json.dumps({
|
||||
"symbol": "AAPL",
|
||||
"annualReports": [
|
||||
{"fiscalDateEnding": "2025-12-31", "totalAssets": "1"}, # future -> must drop
|
||||
{"fiscalDateEnding": "2023-12-31", "totalAssets": "2"}, # past -> must keep
|
||||
],
|
||||
"quarterlyReports": [
|
||||
{"fiscalDateEnding": "2024-06-30", "totalAssets": "3"}, # future -> must drop
|
||||
{"fiscalDateEnding": "2023-09-30", "totalAssets": "4"}, # past -> must keep
|
||||
],
|
||||
})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_fundamentals_look_ahead_filter_runs_on_json_string(monkeypatch):
|
||||
# #1115: the payload arrives as a JSON *string*; the old dict-only guard let
|
||||
# future-dated fiscal periods leak into historical runs.
|
||||
monkeypatch.setattr(avf, "_make_api_request", lambda fn, params: _FUNDAMENTALS_JSON)
|
||||
out = avf.get_balance_sheet("AAPL", curr_date="2024-01-01")
|
||||
assert isinstance(out, str) # callers still receive a str
|
||||
parsed = json.loads(out)
|
||||
assert [r["fiscalDateEnding"] for r in parsed["annualReports"]] == ["2023-12-31"]
|
||||
assert [r["fiscalDateEnding"] for r in parsed["quarterlyReports"]] == ["2023-09-30"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_fundamentals_no_curr_date_passes_through(monkeypatch):
|
||||
monkeypatch.setattr(avf, "_make_api_request", lambda fn, params: _FUNDAMENTALS_JSON)
|
||||
assert avf.get_income_statement("AAPL") == _FUNDAMENTALS_JSON
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_fundamentals_non_json_body_unchanged(monkeypatch):
|
||||
monkeypatch.setattr(avf, "_make_api_request", lambda fn, params: "not-json")
|
||||
assert avf.get_cashflow("AAPL", curr_date="2024-01-01") == "not-json"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Date trim (see the rationale on the unguarded trim in alpha_vantage_common)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_DAILY_CSV = (
|
||||
"timestamp,open,high,low,close,volume\n"
|
||||
"2024-05-13,1,1,1,1,10\n" # after end_date -> must never be served
|
||||
"2024-05-10,1,1,1,1,10\n"
|
||||
"2024-05-09,1,1,1,1,10\n"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stock_data_is_trimmed_to_the_requested_window(monkeypatch):
|
||||
monkeypatch.setattr(avs, "_make_api_request", lambda *a, **k: _DAILY_CSV)
|
||||
out = avs.get_stock("IBM", "2024-05-09", "2024-05-10")
|
||||
assert "2024-05-10" in out and "2024-05-09" in out
|
||||
assert "2024-05-13" not in out, "bar after end_date leaked into the window"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_unparseable_body_is_never_served_untrimmed(monkeypatch):
|
||||
"""The trim used to swallow the failure and return the whole body, putting
|
||||
bars after end_date into a backtest. It must raise instead."""
|
||||
monkeypatch.setattr(avs, "_make_api_request",
|
||||
lambda *a, **k: "timestamp,close\nnot-a-date,1\n")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
avs.get_stock("IBM", "2024-05-09", "2024-05-10")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_empty_body_still_passes_through(monkeypatch):
|
||||
monkeypatch.setattr(avs, "_make_api_request", lambda *a, **k: "")
|
||||
assert avs.get_stock("IBM", "2024-05-09", "2024-05-10") == ""
|
||||
|
||||
|
||||
def test_request_error_message_carries_no_key(monkeypatch):
|
||||
# Alpha Vantage also sends its key in the URL (#1324).
|
||||
import requests
|
||||
key = "AVKEY1234567890XYZ"
|
||||
monkeypatch.setenv("ALPHA_VANTAGE_API_KEY", key)
|
||||
|
||||
def boom(*a, **k):
|
||||
raise requests.Timeout(f"Read timed out. url: https://www.alphavantage.co/query?apikey={key}")
|
||||
|
||||
monkeypatch.setattr(net.requests, "get", boom)
|
||||
with pytest.raises(requests.Timeout) as caught:
|
||||
av._make_api_request("OVERVIEW", {"symbol": "IBM"})
|
||||
assert key not in str(caught.value)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_global_news_omitted_optionals_use_the_configured_defaults(monkeypatch):
|
||||
"""The tool passes None for an omitted look_back_days or limit (#1326)."""
|
||||
from tradingagents.dataflows.vendors.alpha_vantage import news as alpha_vantage_news
|
||||
|
||||
monkeypatch.setattr(alpha_vantage_news, "get_config",
|
||||
lambda: {"global_news_lookback_days": 3, "global_news_article_limit": 9})
|
||||
seen = {}
|
||||
monkeypatch.setattr(alpha_vantage_news, "_make_api_request", lambda fn, params: seen.update(params) or "{}")
|
||||
|
||||
alpha_vantage_news.get_global_news("2026-08-14", None, None)
|
||||
|
||||
assert seen["time_from"].startswith("20260811") and seen["limit"] == "9"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_news_window_includes_the_analysis_day(monkeypatch):
|
||||
"""time_to was midnight at the start of the end date, so everything
|
||||
published during the analysis day, the most decision-relevant day, was
|
||||
excluded. The yfinance path includes it."""
|
||||
from tradingagents.dataflows.vendors.alpha_vantage import news as alpha_vantage_news
|
||||
|
||||
seen = {}
|
||||
monkeypatch.setattr(alpha_vantage_news, "_make_api_request",
|
||||
lambda fn, params: seen.update(params) or "{}")
|
||||
|
||||
alpha_vantage_news.get_news("AAPL", "2026-03-10", "2026-03-14")
|
||||
|
||||
assert seen["time_from"] == "20260310T0000"
|
||||
assert seen["time_to"] == "20260314T2359"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("indicator", ["vwma", "mfi"])
|
||||
def test_an_indicator_this_vendor_lacks_lets_the_next_one_serve_it(indicator):
|
||||
"""Returning prose counts as success to the router, so the chain stops at a
|
||||
vendor that cannot compute the indicator while the next one can."""
|
||||
from tradingagents.dataflows.errors import VendorError
|
||||
from tradingagents.dataflows.vendors.alpha_vantage import indicator as alpha_vantage_indicator
|
||||
|
||||
with pytest.raises(VendorError):
|
||||
alpha_vantage_indicator.get_indicator("AAPL", indicator, "2026-05-08", 30)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_ticker_news_asks_for_only_as_many_articles_as_configured(monkeypatch):
|
||||
"""The endpoint returns 50 articles with per-article sentiment arrays by
|
||||
default, and the whole payload went into the prompt."""
|
||||
from tradingagents.dataflows.vendors.alpha_vantage import news as alpha_vantage_news
|
||||
|
||||
monkeypatch.setattr(alpha_vantage_news, "get_config", lambda: {"news_article_limit": 8})
|
||||
seen = {}
|
||||
monkeypatch.setattr(alpha_vantage_news, "_make_api_request",
|
||||
lambda fn, params: seen.update(params) or "{}")
|
||||
|
||||
alpha_vantage_news.get_news("AAPL", "2026-03-10", "2026-03-14")
|
||||
|
||||
assert seen["limit"] == "8"
|
||||
@@ -0,0 +1,30 @@
|
||||
import unittest
|
||||
|
||||
from tradingagents.graph.analyst_execution import (
|
||||
build_analyst_execution_plan,
|
||||
)
|
||||
|
||||
|
||||
class AnalystExecutionPlanTests(unittest.TestCase):
|
||||
def test_build_plan_preserves_selected_order(self):
|
||||
plan = build_analyst_execution_plan(["news", "market"])
|
||||
|
||||
self.assertEqual([spec.key for spec in plan.specs], ["news", "market"])
|
||||
self.assertEqual(plan.specs[0].agent_node, "News Analyst")
|
||||
self.assertEqual(plan.specs[0].tool_node, "tools_news")
|
||||
self.assertEqual(plan.specs[0].clear_node, "Msg Clear News")
|
||||
|
||||
def test_rejects_unknown_analyst_keys(self):
|
||||
with self.assertRaises(ValueError):
|
||||
build_analyst_execution_plan(["market", "macro"])
|
||||
|
||||
def test_social_key_displays_as_sentiment_analyst(self):
|
||||
# The wire key stays "social" for saved-config back-compat, but the
|
||||
# user-visible agent_node label must match the v0.2.5 rename so the
|
||||
# wall-time summary and any future consumer of agent_node says
|
||||
# "Sentiment Analyst" rather than the legacy "Social Analyst".
|
||||
plan = build_analyst_execution_plan(["social"])
|
||||
spec = plan.specs[0]
|
||||
self.assertEqual(spec.key, "social")
|
||||
self.assertEqual(spec.agent_node, "Sentiment Analyst")
|
||||
self.assertEqual(spec.report_key, "sentiment_report")
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Tests for Anthropic effort-parameter gating (#831).
|
||||
|
||||
Haiku (any version) and Sonnet 4.5 reject the ``effort`` parameter with a
|
||||
400. Only Opus 4.5+ and Sonnet 4.6+ accept it. The gate uses a per-family
|
||||
minimum version so future ``claude-{opus,sonnet}-X-Y`` releases inherit
|
||||
support automatically.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients import anthropic_client as mod
|
||||
|
||||
|
||||
def _capture_kwargs(monkeypatch):
|
||||
captured: dict = {}
|
||||
monkeypatch.setattr(
|
||||
mod, "NormalizedChatAnthropic",
|
||||
lambda **kwargs: captured.setdefault("kwargs", kwargs),
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEffortGate:
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-haiku-4-5", "claude-haiku-5-0", "claude-haiku-4-7-preview",
|
||||
# Sonnet 4.5 (and earlier) 400 on effort — only Sonnet 4.6+ supports it.
|
||||
"claude-sonnet-4-5", "claude-sonnet-4-0",
|
||||
],
|
||||
)
|
||||
def test_unsupported_models_do_not_receive_effort(self, monkeypatch, model):
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(model=model, effort="medium", api_key="x").get_llm()
|
||||
assert "effort" not in captured["kwargs"]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-opus-4-5", "claude-opus-4-6", "claude-opus-4-7",
|
||||
"claude-sonnet-4-6",
|
||||
],
|
||||
)
|
||||
def test_current_opus_and_sonnet_receive_effort(self, monkeypatch, model):
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(model=model, effort="high", api_key="x").get_llm()
|
||||
assert captured["kwargs"]["effort"] == "high"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["claude-opus-5-0", "claude-opus-4-8", "claude-sonnet-5-0"],
|
||||
)
|
||||
def test_future_opus_sonnet_inherit_effort_via_pattern(self, monkeypatch, model):
|
||||
"""Forward-compat: new Opus/Sonnet versions don't need a code change."""
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(model=model, effort="low", api_key="x").get_llm()
|
||||
assert captured["kwargs"]["effort"] == "low"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
# Claude 5 family uses single-number version IDs; all are effort-capable.
|
||||
["claude-sonnet-5", "claude-fable-5", "claude-mythos-5", "claude-opus-5", "claude-opus-5-5",
|
||||
"claude-fable-5-1"],
|
||||
)
|
||||
def test_claude_5_family_receives_effort(self, monkeypatch, model):
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(model=model, effort="high", api_key="x").get_llm()
|
||||
assert captured["kwargs"]["effort"] == "high"
|
||||
|
||||
def test_mythos_preview_receives_effort(self, monkeypatch):
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(
|
||||
model="claude-mythos-preview", effort="medium", api_key="x"
|
||||
).get_llm()
|
||||
assert captured["kwargs"]["effort"] == "medium"
|
||||
|
||||
def test_unknown_anthropic_model_does_not_receive_effort(self, monkeypatch):
|
||||
"""Default is conservative — unknown models don't get effort to avoid 400s."""
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(
|
||||
model="claude-experimental-x", effort="medium", api_key="x"
|
||||
).get_llm()
|
||||
assert "effort" not in captured["kwargs"]
|
||||
|
||||
def test_other_kwargs_still_forwarded_when_effort_skipped(self, monkeypatch):
|
||||
"""Skipping effort must not break other passthrough kwargs."""
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
mod.AnthropicClient(
|
||||
model="claude-haiku-4-5",
|
||||
effort="medium",
|
||||
api_key="placeholder",
|
||||
max_tokens=1024,
|
||||
timeout=30,
|
||||
).get_llm()
|
||||
assert captured["kwargs"]["api_key"] == "placeholder"
|
||||
assert captured["kwargs"]["max_tokens"] == 1024
|
||||
assert captured["kwargs"]["timeout"] == 30
|
||||
assert "effort" not in captured["kwargs"]
|
||||
@@ -0,0 +1,192 @@
|
||||
"""Tests for the canonical provider->env-var mapping and the CLI key-prompt helper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import stat
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.api_key_env import PROVIDER_API_KEY_ENV, get_api_key_env
|
||||
|
||||
# ---- Mapping coverage -----------------------------------------------------
|
||||
|
||||
|
||||
def test_every_select_llm_provider_choice_has_an_entry():
|
||||
"""select_llm_provider() must not present a provider the mapping doesn't know about."""
|
||||
# Mirrors the dropdown order in cli/prompts.select_llm_provider so the two
|
||||
# stay in lockstep. Region-specific keys (qwen-cn / minimax-cn / glm-cn)
|
||||
# are reached via the secondary region prompt, so they must also be present.
|
||||
expected = {
|
||||
"openai", "google", "anthropic", "xai", "deepseek",
|
||||
"qwen", "qwen-cn",
|
||||
"glm", "glm-cn",
|
||||
"minimax", "minimax-cn",
|
||||
"openrouter", "azure", "ollama",
|
||||
}
|
||||
assert expected.issubset(PROVIDER_API_KEY_ENV.keys())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider,env_var",
|
||||
[
|
||||
("openai", "OPENAI_API_KEY"),
|
||||
("anthropic", "ANTHROPIC_API_KEY"),
|
||||
("google", "GOOGLE_API_KEY"),
|
||||
("azure", "AZURE_OPENAI_API_KEY"),
|
||||
("xai", "XAI_API_KEY"),
|
||||
("deepseek", "DEEPSEEK_API_KEY"),
|
||||
("qwen", "DASHSCOPE_API_KEY"),
|
||||
("qwen-cn", "DASHSCOPE_CN_API_KEY"),
|
||||
("glm", "ZHIPU_API_KEY"),
|
||||
("glm-cn", "ZHIPU_CN_API_KEY"),
|
||||
("minimax", "MINIMAX_API_KEY"),
|
||||
("minimax-cn", "MINIMAX_CN_API_KEY"),
|
||||
("openrouter", "OPENROUTER_API_KEY"),
|
||||
],
|
||||
)
|
||||
def test_known_providers_resolve(provider, env_var):
|
||||
assert get_api_key_env(provider) == env_var
|
||||
|
||||
|
||||
def test_ollama_has_no_key():
|
||||
assert get_api_key_env("ollama") is None
|
||||
|
||||
|
||||
def test_unknown_provider_returns_none():
|
||||
assert get_api_key_env("not-a-real-provider") is None
|
||||
|
||||
|
||||
def test_case_insensitive_lookup():
|
||||
assert get_api_key_env("OpenAI") == "OPENAI_API_KEY"
|
||||
assert get_api_key_env("QWEN-CN") == "DASHSCOPE_CN_API_KEY"
|
||||
|
||||
|
||||
# ---- ensure_api_key behavior ---------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def prompts(monkeypatch):
|
||||
"""Import cli.prompts with a fresh environment so module-level state is consistent."""
|
||||
import importlib
|
||||
|
||||
import cli.prompts as prompts_module
|
||||
return importlib.reload(prompts_module)
|
||||
|
||||
|
||||
def test_ensure_api_key_returns_existing(monkeypatch, prompts):
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-already-set")
|
||||
result = prompts.ensure_api_key("openai")
|
||||
assert result == "sk-already-set"
|
||||
|
||||
|
||||
def test_ensure_api_key_no_op_for_ollama(monkeypatch, prompts):
|
||||
# Even with no env var set, ollama should not prompt and should return None.
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
with patch.object(prompts, "questionary") as mock_q:
|
||||
result = prompts.ensure_api_key("ollama")
|
||||
assert result is None
|
||||
mock_q.password.assert_not_called()
|
||||
|
||||
|
||||
def test_ensure_api_key_unknown_provider_no_prompt(monkeypatch, prompts):
|
||||
with patch.object(prompts, "questionary") as mock_q:
|
||||
result = prompts.ensure_api_key("totally-fake-provider")
|
||||
assert result is None
|
||||
mock_q.password.assert_not_called()
|
||||
|
||||
|
||||
def test_ensure_api_key_prompts_and_writes_to_env(monkeypatch, tmp_path, prompts):
|
||||
"""When key is missing, user-pasted value must be written to .env AND os.environ."""
|
||||
monkeypatch.delenv("DEEPSEEK_API_KEY", raising=False)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-deepseek-test")})()
|
||||
with patch.object(prompts.questionary, "password", return_value=fake_prompt):
|
||||
result = prompts.ensure_api_key("deepseek")
|
||||
|
||||
assert result == "sk-deepseek-test"
|
||||
assert os.environ["DEEPSEEK_API_KEY"] == "sk-deepseek-test"
|
||||
env_file = tmp_path / ".env"
|
||||
assert env_file.exists()
|
||||
assert "DEEPSEEK_API_KEY" in env_file.read_text()
|
||||
assert "sk-deepseek-test" in env_file.read_text()
|
||||
|
||||
|
||||
def test_ensure_api_key_user_cancels_returns_none(monkeypatch, tmp_path, prompts):
|
||||
"""Empty prompt response (user cancelled) must not write to .env."""
|
||||
monkeypatch.delenv("XAI_API_KEY", raising=False)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
|
||||
fake_prompt = type("P", (), {"ask": staticmethod(lambda: None)})()
|
||||
with patch.object(prompts.questionary, "password", return_value=fake_prompt):
|
||||
result = prompts.ensure_api_key("xai")
|
||||
|
||||
assert result is None
|
||||
assert "XAI_API_KEY" not in os.environ
|
||||
# .env may or may not exist depending on find_dotenv's walk, but if it
|
||||
# does it must not contain the key.
|
||||
env_file = tmp_path / ".env"
|
||||
if env_file.exists():
|
||||
assert "XAI_API_KEY" not in env_file.read_text()
|
||||
|
||||
|
||||
def test_ensure_api_key_updates_existing_env_file(monkeypatch, tmp_path, prompts):
|
||||
"""An existing .env with other keys must be preserved on writeback."""
|
||||
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
|
||||
monkeypatch.chdir(tmp_path)
|
||||
env_file = tmp_path / ".env"
|
||||
env_file.write_text("OPENAI_API_KEY=sk-existing\nOTHER=value\n")
|
||||
|
||||
fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-openrouter-new")})()
|
||||
with patch.object(prompts.questionary, "password", return_value=fake_prompt):
|
||||
prompts.ensure_api_key("openrouter")
|
||||
|
||||
content = env_file.read_text()
|
||||
assert "OPENAI_API_KEY" in content and "sk-existing" in content
|
||||
assert "OTHER=value" in content
|
||||
assert "OPENROUTER_API_KEY" in content and "sk-openrouter-new" in content
|
||||
|
||||
|
||||
def _prompt_key(prompts, monkeypatch, tmp_path, key="sk-typed-in"):
|
||||
monkeypatch.chdir(tmp_path)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
monkeypatch.setattr(prompts, "find_dotenv", lambda **k: "")
|
||||
with patch.object(prompts, "questionary") as mock_q:
|
||||
mock_q.password.return_value.ask.return_value = key
|
||||
prompts.ensure_api_key("openai")
|
||||
|
||||
|
||||
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
|
||||
def test_saved_key_file_is_owner_only(monkeypatch, prompts, tmp_path):
|
||||
# The prompt writes a real credential; the file must not be readable by
|
||||
# other local users whatever the umask is.
|
||||
old = os.umask(0o002)
|
||||
try:
|
||||
_prompt_key(prompts, monkeypatch, tmp_path)
|
||||
finally:
|
||||
os.umask(old)
|
||||
env = tmp_path / ".env"
|
||||
assert "sk-typed-in" in env.read_text()
|
||||
assert stat.S_IMODE(env.stat().st_mode) == 0o600
|
||||
|
||||
|
||||
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
|
||||
def test_existing_key_file_is_tightened_before_writing(monkeypatch, prompts, tmp_path):
|
||||
env = tmp_path / ".env"
|
||||
env.write_text("OTHER=1\n")
|
||||
os.chmod(env, 0o664)
|
||||
_prompt_key(prompts, monkeypatch, tmp_path)
|
||||
assert stat.S_IMODE(env.stat().st_mode) == 0o600
|
||||
assert "OTHER=1" in env.read_text()
|
||||
|
||||
|
||||
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
|
||||
def test_read_only_key_file_is_still_updated(monkeypatch, prompts, tmp_path):
|
||||
env = tmp_path / ".env"
|
||||
env.write_text("OTHER=1\n")
|
||||
os.chmod(env, 0o400)
|
||||
_prompt_key(prompts, monkeypatch, tmp_path)
|
||||
assert "sk-typed-in" in env.read_text()
|
||||
assert stat.S_IMODE(env.stat().st_mode) == 0o600
|
||||
@@ -0,0 +1,276 @@
|
||||
"""Backtesting: many single-shot decisions, scored by the decision log.
|
||||
|
||||
A run already records its rating and later settles it with realized and alpha
|
||||
return against the regional benchmark. A backtest is that machinery over a grid
|
||||
of tickers and dates, aggregated. It evaluates decision quality; it does not
|
||||
simulate a portfolio, so there is no execution, no fees and no equity curve.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.backtest import iter_grid, run_backtest, summarize
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
|
||||
DECISION = "Rating: Buy\n\nbuy it"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_grid_spacing_and_canonical_dates():
|
||||
assert iter_grid("2026-01-05", "2026-01-20", every_n_days=7) == ["2026-01-05", "2026-01-12", "2026-01-19"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_grid_stops_at_today(monkeypatch):
|
||||
import tradingagents.backtest as bt
|
||||
|
||||
monkeypatch.setattr(bt, "get_current_date", lambda: "2026-01-10")
|
||||
assert iter_grid("2026-01-05", "2026-02-20", every_n_days=5) == ["2026-01-05", "2026-01-10"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_grid_rejects_a_non_canonical_date():
|
||||
with pytest.raises(ValueError, match="YYYY-MM-DD"):
|
||||
iter_grid("2026-1-5", "2026-01-20")
|
||||
|
||||
|
||||
class _FakeGraph:
|
||||
"""Stands in for TradingAgentsGraph, writing to the log the harness gave it."""
|
||||
|
||||
instances: list = []
|
||||
fail_on: set = set()
|
||||
|
||||
def __init__(self, selected_analysts=None, config=None, **kw):
|
||||
self.analysts = list(selected_analysts) if selected_analysts else None
|
||||
self.config = config
|
||||
self.memory_log = TradingMemoryLog(config)
|
||||
self.calls = []
|
||||
self.settled = []
|
||||
_FakeGraph.instances.append(self)
|
||||
|
||||
def propagate(self, ticker, trade_date, asset_type="stock", portfolio=None):
|
||||
self.calls.append((ticker, trade_date))
|
||||
if (ticker, trade_date) in _FakeGraph.fail_on:
|
||||
raise RuntimeError("vendor exploded")
|
||||
self.memory_log.store_decision(ticker, trade_date, DECISION)
|
||||
return {"final_trade_decision": DECISION}, "Buy"
|
||||
|
||||
def settle_pending(self, ticker):
|
||||
self.settled.append(ticker)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fake_graph(monkeypatch, tmp_path):
|
||||
import tradingagents.backtest as bt
|
||||
|
||||
_FakeGraph.instances = []
|
||||
_FakeGraph.fail_on = set()
|
||||
monkeypatch.setattr(bt, "TradingAgentsGraph", _FakeGraph)
|
||||
return _FakeGraph
|
||||
|
||||
|
||||
def _config(tmp_path):
|
||||
return {"results_dir": str(tmp_path / "results"),
|
||||
"memory_log_path": str(tmp_path / "live_trading_memory.md")}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_live_decision_log_is_never_written(tmp_path):
|
||||
config = _config(tmp_path)
|
||||
result = run_backtest(["NVDA"], ["2026-01-05", "2026-01-12"], config)
|
||||
|
||||
assert not (tmp_path / "live_trading_memory.md").exists()
|
||||
assert result.log_path.exists() and result.cells_run == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_cell_already_in_the_log_is_not_run_again(tmp_path):
|
||||
config = _config(tmp_path)
|
||||
first = run_backtest(["NVDA"], ["2026-01-05"], config)
|
||||
|
||||
again = run_backtest(["NVDA"], ["2026-01-05", "2026-01-12"], config, run_id=first.run_id)
|
||||
|
||||
assert again.cells_run == 1 and again.skipped == 1
|
||||
assert _FakeGraph.instances[-1].calls == [("NVDA", "2026-01-12")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_every_ticker_is_settled_after_the_grid(tmp_path):
|
||||
"""Settlement runs at the start of the next same-ticker run, so the last
|
||||
date of each ticker would stay pending without an explicit pass."""
|
||||
run_backtest(["NVDA", "AAPL"], ["2026-01-05", "2026-01-12"], _config(tmp_path))
|
||||
assert sorted(_FakeGraph.instances[-1].settled) == ["AAPL", "NVDA"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_failed_cell_does_not_abort_the_sweep(tmp_path):
|
||||
_FakeGraph.fail_on = {("NVDA", "2026-01-05")}
|
||||
result = run_backtest(["NVDA"], ["2026-01-05", "2026-01-12"], _config(tmp_path))
|
||||
|
||||
assert result.cells_run == 1
|
||||
assert result.failures == [("NVDA", "2026-01-05", "vendor exploded")]
|
||||
|
||||
|
||||
# --- reading the result ------------------------------------------------------
|
||||
|
||||
def _log_with(tmp_path, rows):
|
||||
log = TradingMemoryLog({"memory_log_path": str(tmp_path / "m.md")})
|
||||
for ticker, date, decision, outcome in rows:
|
||||
log.store_decision(ticker, date, decision)
|
||||
if outcome is not None:
|
||||
log.update_with_outcome(ticker, date, outcome[0], outcome[1], 5, "note", "2026-02-01")
|
||||
return tmp_path / "m.md"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_summary_scores_resolved_cells_and_keeps_pending_out_of_the_average(tmp_path):
|
||||
log = _log_with(tmp_path, [
|
||||
("NVDA", "2026-01-05", "Rating: Buy\n\nx", (0.10, 0.04)),
|
||||
("NVDA", "2026-01-12", "Rating: Buy\n\nx", (-0.02, -0.02)),
|
||||
("AAPL", "2026-01-05", "Rating: Sell\n\nx", None),
|
||||
])
|
||||
|
||||
summary = summarize(log)
|
||||
|
||||
assert summary.resolved == 2 and summary.pending == 1
|
||||
buys = summary.by_rating["Buy"]
|
||||
assert buys.count == 2 and buys.hit_rate == 0.5 and round(buys.mean_alpha, 4) == 0.01
|
||||
assert "Sell" not in summary.by_rating # unsettled: nothing to score yet
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_summary_states_what_it_cannot_prove(tmp_path):
|
||||
text = summarize(_log_with(tmp_path, [("NVDA", "2026-01-05", DECISION, (0.1, 0.05))])).render()
|
||||
assert "not archived" in text
|
||||
assert "one" in text.lower() and "sampl" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_analyst_set_under_test_is_the_one_that_runs(tmp_path):
|
||||
"""A backtest of a two-analyst setup must not silently run four."""
|
||||
run_backtest(["NVDA"], ["2026-01-05"], _config(tmp_path), selected_analysts=["market", "news"])
|
||||
assert _FakeGraph.instances[-1].analysts == ["market", "news"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_run_id_cannot_escape_the_results_directory(tmp_path):
|
||||
"""run_id becomes a path segment, so it is validated like a ticker is."""
|
||||
with pytest.raises(ValueError):
|
||||
run_backtest(["NVDA"], ["2026-01-05"], _config(tmp_path), run_id="../../escaped")
|
||||
with pytest.raises(ValueError):
|
||||
run_backtest(["NVDA"], ["2026-01-05"], _config(tmp_path), run_id="/etc/cron.d/x")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_failed_settlement_does_not_lose_the_remaining_tickers(tmp_path, monkeypatch):
|
||||
"""Settlement reflects with an LLM, so it can fail; the sweep still returns
|
||||
its result and every other ticker still gets settled."""
|
||||
settled = []
|
||||
|
||||
def _settle(self, ticker):
|
||||
if ticker == "NVDA":
|
||||
raise RuntimeError("reflector timed out")
|
||||
settled.append(ticker)
|
||||
|
||||
monkeypatch.setattr(_FakeGraph, "settle_pending", _settle, raising=False)
|
||||
result = run_backtest(["NVDA", "AAPL"], ["2026-01-05"], _config(tmp_path))
|
||||
|
||||
assert result.cells_run == 2
|
||||
assert settled == ["AAPL"]
|
||||
assert result.settlement_failures == [("NVDA", "reflector timed out")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_pending_note_appears_only_when_something_is_pending(tmp_path):
|
||||
settled = [("NVDA", "2026-01-05", DECISION, (0.1, 0.05))]
|
||||
assert "Pending" not in summarize(_log_with(tmp_path, settled)).render()
|
||||
assert "Pending" in summarize(_log_with(tmp_path, settled + [("AAPL", "2026-01-05", DECISION, None)])).render()
|
||||
|
||||
|
||||
# --- scoring reads the direction the rating claimed ---------------------------
|
||||
|
||||
def _scored(tmp_path, rows):
|
||||
log = _log_with(tmp_path, rows)
|
||||
return summarize(log).by_rating
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_bearish_call_that_fell_counts_as_right(tmp_path):
|
||||
"""Alpha below the benchmark is the outcome a Sell predicted; scoring it as
|
||||
a miss reported the system as wrong exactly when it was right."""
|
||||
scores = _scored(tmp_path, [
|
||||
("NVDA", "2026-01-05", "**Rating**: Sell\n\nx", (-0.08, -0.05)),
|
||||
("AAPL", "2026-01-05", "**Rating**: Underweight\n\nx", (-0.03, -0.02)),
|
||||
])
|
||||
assert scores["Sell"].hit_rate == 1.0
|
||||
assert scores["Underweight"].hit_rate == 1.0
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_bearish_call_that_rose_counts_as_wrong(tmp_path):
|
||||
scores = _scored(tmp_path, [("NVDA", "2026-01-05", "**Rating**: Sell\n\nx", (0.08, 0.05))])
|
||||
assert scores["Sell"].hit_rate == 0.0
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_bullish_call_is_scored_the_same_way_as_before(tmp_path):
|
||||
scores = _scored(tmp_path, [
|
||||
("NVDA", "2026-01-05", "**Rating**: Buy\n\nx", (0.10, 0.04)),
|
||||
("AAPL", "2026-01-05", "**Rating**: Buy\n\nx", (-0.02, -0.02)),
|
||||
])
|
||||
assert scores["Buy"].hit_rate == 0.5
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_hold_claims_no_direction_so_it_gets_no_hit_rate(tmp_path):
|
||||
scores = _scored(tmp_path, [("NVDA", "2026-01-05", "**Rating**: Hold\n\nx", (0.01, 0.005))])
|
||||
assert scores["Hold"].hit_rate is None
|
||||
assert scores["Hold"].mean_alpha == 0.005
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_report_names_the_window_the_scores_cover(tmp_path):
|
||||
text = summarize(_log_with(tmp_path, [
|
||||
("NVDA", "2026-01-05", "**Rating**: Buy\n\nx", (0.1, 0.05))])).render()
|
||||
assert "5" in text and "day" in text.lower()
|
||||
assert "Hold" not in text or "no direction" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_window_reported_is_the_one_the_outcomes_used(tmp_path):
|
||||
"""The log records the window each outcome was measured over; the summary
|
||||
must not claim a different one."""
|
||||
log = TradingMemoryLog({"memory_log_path": str(tmp_path / "m.md")})
|
||||
log.store_decision("NVDA", "2026-01-05", "**Rating**: Buy\n\nx")
|
||||
log.update_with_outcome("NVDA", "2026-01-05", 0.1, 0.04, 21, "note", "2026-02-01")
|
||||
|
||||
assert "21 trading days" in summarize(tmp_path / "m.md").render()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_backtest_result_is_summarized_directly(tmp_path):
|
||||
"""The result names its own log, so a caller never builds the log to score it."""
|
||||
from tradingagents.backtest import BacktestResult
|
||||
|
||||
path = _log_with(tmp_path, [("NVDA", "2026-01-05", "Rating: Buy\n\nx", (0.10, 0.04))])
|
||||
|
||||
assert summarize(BacktestResult(run_id="r", log_path=path)).resolved == 1
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_log_path_that_does_not_exist_is_an_error_not_an_empty_summary(tmp_path):
|
||||
missing = tmp_path / "no-such-dir" / "m.md"
|
||||
|
||||
with pytest.raises(FileNotFoundError):
|
||||
summarize(missing)
|
||||
assert not missing.parent.exists()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_result_whose_cells_all_failed_summarizes_as_empty(tmp_path):
|
||||
from tradingagents.backtest import BacktestResult
|
||||
|
||||
result = BacktestResult(run_id="r", log_path=tmp_path / "never-written.md")
|
||||
|
||||
assert summarize(result).resolved == 0
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Amazon Bedrock — first-class native client via the optional langchain-aws extra.
|
||||
|
||||
Auth uses the AWS credential chain (no single key env); the model is a Bedrock
|
||||
model ID / inference profile ID; langchain-aws is imported lazily with a clear
|
||||
install hint when the [bedrock] extra is absent.
|
||||
"""
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.api_key_env import get_api_key_env
|
||||
from tradingagents.llm_clients.factory import create_llm_client
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_factory_routes_bedrock():
|
||||
client = create_llm_client("bedrock", "us.anthropic.claude-opus-4-8-v1:0")
|
||||
assert type(client).__name__ == "BedrockClient"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_bedrock_any_model_and_no_key_env():
|
||||
assert validate_model("bedrock", "any.model-id:0") is True
|
||||
# Bedrock uses the AWS credential chain, so there is no single key env.
|
||||
assert get_api_key_env("bedrock") is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_helpful_error_when_langchain_aws_absent(monkeypatch):
|
||||
import tradingagents.llm_clients.bedrock_client as bc
|
||||
monkeypatch.setattr(bc, "_BEDROCK_CLASS", None)
|
||||
monkeypatch.setitem(sys.modules, "langchain_aws", None) # force ImportError on import
|
||||
with pytest.raises(ImportError, match=r"bedrock"):
|
||||
create_llm_client("bedrock", "m").get_llm()
|
||||
|
||||
|
||||
def _capture_kwargs(monkeypatch):
|
||||
"""Stub _bedrock_class so the constructor kwargs are testable without the
|
||||
optional langchain-aws extra installed."""
|
||||
import tradingagents.llm_clients.bedrock_client as bc
|
||||
captured = {}
|
||||
|
||||
class _FakeChat:
|
||||
def __init__(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
|
||||
monkeypatch.setattr(bc, "_bedrock_class", lambda: _FakeChat)
|
||||
return captured
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_bearer_token_passed_as_api_key(monkeypatch):
|
||||
# #1103: a Bedrock API key authenticates without AWS access keys.
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bt-secret")
|
||||
monkeypatch.setenv("AWS_DEFAULT_REGION", "us-east-1")
|
||||
create_llm_client("bedrock", "us.anthropic.claude-opus-4-8-v1:0").get_llm()
|
||||
assert captured["api_key"] == "bt-secret"
|
||||
assert captured["region_name"] == "us-east-1"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_no_bearer_token_omits_api_key(monkeypatch):
|
||||
# Without a token, fall back to the AWS credential chain (no api_key kwarg).
|
||||
captured = _capture_kwargs(monkeypatch)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
create_llm_client("bedrock", "us.anthropic.claude-opus-4-8-v1:0").get_llm()
|
||||
assert "api_key" not in captured
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_construction_when_extra_installed(monkeypatch):
|
||||
pytest.importorskip("langchain_aws")
|
||||
import tradingagents.llm_clients.bedrock_client as bc
|
||||
monkeypatch.setattr(bc, "_BEDROCK_CLASS", None)
|
||||
monkeypatch.setenv("AWS_DEFAULT_REGION", "eu-west-1")
|
||||
llm = create_llm_client("bedrock", "us.anthropic.claude-sonnet-5").get_llm()
|
||||
assert type(llm).__name__ == "NormalizedChatBedrockConverse"
|
||||
assert llm.region_name == "eu-west-1"
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Unit tests for the LLM capability table."""
|
||||
|
||||
from dataclasses import FrozenInstanceError
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.capabilities import (
|
||||
get_capabilities,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestExactIdMatches:
|
||||
def test_deepseek_chat_supports_tool_choice(self):
|
||||
caps = get_capabilities("deepseek-chat")
|
||||
assert caps.supports_tool_choice is True
|
||||
|
||||
def test_deepseek_reasoner_rejects_tool_choice(self):
|
||||
caps = get_capabilities("deepseek-reasoner")
|
||||
assert caps.supports_tool_choice is False
|
||||
assert caps.requires_reasoning_content_roundtrip is True
|
||||
|
||||
def test_deepseek_v4_flash_rejects_tool_choice(self):
|
||||
caps = get_capabilities("deepseek-v4-flash")
|
||||
assert caps.supports_tool_choice is False
|
||||
assert caps.requires_reasoning_content_roundtrip is True
|
||||
|
||||
def test_deepseek_v4_pro_rejects_tool_choice(self):
|
||||
caps = get_capabilities("deepseek-v4-pro")
|
||||
assert caps.supports_tool_choice is False
|
||||
assert caps.requires_reasoning_content_roundtrip is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPatternMatches:
|
||||
"""Forward-compat regex patterns catch unknown DeepSeek and MiniMax variants."""
|
||||
|
||||
def test_future_deepseek_v5_inherits_thinking_quirks(self):
|
||||
caps = get_capabilities("deepseek-v5-flash")
|
||||
assert caps.supports_tool_choice is False
|
||||
assert caps.requires_reasoning_content_roundtrip is True
|
||||
|
||||
def test_future_deepseek_v9_inherits_thinking_quirks(self):
|
||||
caps = get_capabilities("deepseek-v9-anything")
|
||||
assert caps.supports_tool_choice is False
|
||||
|
||||
def test_reasoner_variant_inherits_thinking_quirks(self):
|
||||
caps = get_capabilities("deepseek-reasoner-pro")
|
||||
assert caps.supports_tool_choice is False
|
||||
|
||||
def test_minimax_m3_inherits_thinking_quirks(self):
|
||||
caps = get_capabilities("MiniMax-M3")
|
||||
assert caps.supports_tool_choice is False
|
||||
|
||||
def test_future_minimax_m4_highspeed_inherits_thinking_quirks(self):
|
||||
caps = get_capabilities("MiniMax-M4-highspeed")
|
||||
assert caps.supports_tool_choice is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMinimaxExactMatches:
|
||||
"""MiniMax M2.x models reject langchain's function-spec dict tool_choice
|
||||
(official API enum: none/auto only)."""
|
||||
|
||||
def test_m2_7_rejects_tool_choice(self):
|
||||
caps = get_capabilities("MiniMax-M2.7")
|
||||
assert caps.supports_tool_choice is False
|
||||
assert caps.supports_json_mode is False # only MiniMax-Text-01 supports json_object
|
||||
|
||||
def test_m2_7_highspeed_rejects_tool_choice(self):
|
||||
assert get_capabilities("MiniMax-M2.7-highspeed").supports_tool_choice is False
|
||||
|
||||
def test_m2_1_rejects_tool_choice(self):
|
||||
assert get_capabilities("MiniMax-M2.1").supports_tool_choice is False
|
||||
|
||||
def test_m2_base_rejects_tool_choice(self):
|
||||
assert get_capabilities("MiniMax-M2").supports_tool_choice is False
|
||||
|
||||
def test_m2_x_requires_reasoning_split(self):
|
||||
# M2.x reasoning models need reasoning_split=True so <think> blocks
|
||||
# land in reasoning_details instead of content (#826).
|
||||
for model in ("MiniMax-M2.7", "MiniMax-M2.5-highspeed", "MiniMax-M2"):
|
||||
assert get_capabilities(model).requires_reasoning_split is True
|
||||
|
||||
def test_future_m3_inherits_reasoning_split(self):
|
||||
assert get_capabilities("MiniMax-M3-highspeed").requires_reasoning_split is True
|
||||
|
||||
def test_non_reasoning_minimax_does_not_get_reasoning_split(self):
|
||||
# Coding Plan, MiniMax-Text-01, and any non-M2-prefixed MiniMax model
|
||||
# reject the reasoning_split kwarg via the openai SDK's strict
|
||||
# validation (#826). Default capability has it disabled.
|
||||
for model in ("minimax-text-01", "MiniMax-Coding-Plan", "abab6.5-chat"):
|
||||
assert get_capabilities(model).requires_reasoning_split is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestDefault:
|
||||
"""Unknown / non-DeepSeek models get the permissive default."""
|
||||
|
||||
def test_gpt_default(self):
|
||||
caps = get_capabilities("gpt-4.1")
|
||||
assert caps.supports_tool_choice is True
|
||||
assert caps.preferred_structured_method == "function_calling"
|
||||
|
||||
def test_grok_default(self):
|
||||
caps = get_capabilities("grok-4-0709")
|
||||
assert caps.supports_tool_choice is True
|
||||
|
||||
def test_unknown_model_default(self):
|
||||
caps = get_capabilities("totally-made-up-model-id")
|
||||
assert caps.supports_tool_choice is True
|
||||
|
||||
def test_exact_match_precedes_pattern(self):
|
||||
"""deepseek-chat must NOT match the v\\d regex."""
|
||||
caps = get_capabilities("deepseek-chat")
|
||||
assert caps.supports_tool_choice is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOpenRouterDeepSeekNamespace:
|
||||
"""OpenRouter namespaces DeepSeek as ``deepseek/<id>``; strip it so the
|
||||
same quirks apply as the native provider (#1199)."""
|
||||
|
||||
def test_prefixed_v4_flash_suppresses_tool_choice(self):
|
||||
# Was falling through to _DEFAULT (tool_choice on) -> slow object-form call.
|
||||
assert get_capabilities("deepseek/deepseek-v4-flash").supports_tool_choice is False
|
||||
|
||||
def test_prefixed_reasoner_suppresses_tool_choice(self):
|
||||
assert get_capabilities("deepseek/deepseek-reasoner").supports_tool_choice is False
|
||||
|
||||
def test_prefixed_chat_selects_deepseek_chat_not_default(self):
|
||||
# Must resolve to _DEEPSEEK_CHAT, not _DEFAULT: supports_json_schema=False
|
||||
# is what distinguishes them (both keep tool_choice).
|
||||
caps = get_capabilities("deepseek/deepseek-chat")
|
||||
assert caps.supports_tool_choice is True
|
||||
assert caps.supports_json_schema is False # _DEEPSEEK_CHAT, not _DEFAULT
|
||||
|
||||
def test_only_official_namespace_is_stripped(self):
|
||||
# A third-party publisher whose model name WOULD match a deepseek pattern
|
||||
# must stay _DEFAULT: proves we strip only "deepseek/", not any "*/".
|
||||
caps = get_capabilities("tngtech/deepseek-v4-flash")
|
||||
assert caps.supports_tool_choice is True # not thinking
|
||||
assert caps.supports_json_schema is True # _DEFAULT
|
||||
|
||||
def test_native_ids_unchanged(self):
|
||||
assert get_capabilities("deepseek-v4-flash").supports_tool_choice is False
|
||||
assert get_capabilities("deepseek-chat").supports_tool_choice is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_capabilities_dataclass_is_frozen():
|
||||
"""Capability rows are immutable so they can be safely shared."""
|
||||
caps = get_capabilities("deepseek-chat")
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
caps.supports_tool_choice = False # type: ignore[misc]
|
||||
@@ -0,0 +1,170 @@
|
||||
"""The checkpoint lifecycle is reusable so --checkpoint works on the CLI path (#1249).
|
||||
|
||||
Checkpoint setup previously lived only inside ``propagate``; the CLI streamed the
|
||||
checkpointer-less graph, so ``--checkpoint`` neither saved nor resumed. The
|
||||
lifecycle is now ``begin_checkpoint`` / ``end_checkpoint`` /
|
||||
``clear_checkpoint_on_success`` on TradingAgentsGraph, used by both paths. These
|
||||
tests drive that lifecycle exactly as the CLI does (begin -> stream self.graph ->
|
||||
clear/end) and prove state is saved and resumed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
from typing import TypedDict
|
||||
|
||||
import pytest
|
||||
from langgraph.graph import END, StateGraph
|
||||
|
||||
from tradingagents.graph.checkpointer import checkpoint_step
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
_should_crash = False
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
count: int
|
||||
|
||||
|
||||
def _node_a(state: _State) -> dict:
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
|
||||
def _node_b(state: _State) -> dict:
|
||||
if _should_crash:
|
||||
raise RuntimeError("simulated mid-stream crash")
|
||||
return {"count": state["count"] + 10}
|
||||
|
||||
|
||||
def _workflow() -> StateGraph:
|
||||
b = StateGraph(_State)
|
||||
b.add_node("analyst", _node_a)
|
||||
b.add_node("trader", _node_b)
|
||||
b.set_entry_point("analyst")
|
||||
b.add_edge("analyst", "trader")
|
||||
b.add_edge("trader", END)
|
||||
return b
|
||||
|
||||
|
||||
def _bare_graph(tmpdir, *, enabled=True):
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
g.config = {
|
||||
"checkpoint_enabled": enabled, "data_cache_dir": tmpdir,
|
||||
"max_debate_rounds": 1, "max_risk_discuss_rounds": 1,
|
||||
}
|
||||
g.selected_analysts = ("market",)
|
||||
g.workflow = _workflow()
|
||||
g.graph = g.workflow.compile()
|
||||
g._checkpointer_ctx = None
|
||||
return g
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_disabled_is_a_noop():
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
g = _bare_graph(tmp, enabled=False)
|
||||
plain = g.graph
|
||||
assert g.begin_checkpoint("AAPL", "2026-05-08", "stock") is None
|
||||
assert g.graph is plain # graph not recompiled
|
||||
g.end_checkpoint() # safe no-op
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_begin_returns_thread_id_and_recompiles():
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
g = _bare_graph(tmp)
|
||||
plain = g.graph
|
||||
tid = g.begin_checkpoint("AAPL", "2026-05-08", "stock")
|
||||
try:
|
||||
assert tid # a real thread_id
|
||||
assert g.graph is not plain # recompiled with a checkpointer
|
||||
finally:
|
||||
g.end_checkpoint()
|
||||
assert g._checkpointer_ctx is None # restored
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_checkpoint_input_is_none_only_when_resuming():
|
||||
global _should_crash
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
init = {"count": 0}
|
||||
args = ("AAPL", "2026-05-08", "stock")
|
||||
# Fresh run: no checkpoint yet -> stream the initial state, then crash.
|
||||
g1 = _bare_graph(tmp)
|
||||
tid = g1.begin_checkpoint(*args)
|
||||
try:
|
||||
assert g1._resuming is False
|
||||
assert g1.checkpoint_input(init) is init # not resuming -> initial state
|
||||
_should_crash = True
|
||||
with pytest.raises(RuntimeError):
|
||||
for _ in g1.graph.stream(init, config={"configurable": {"thread_id": tid}}):
|
||||
pass
|
||||
finally:
|
||||
g1.end_checkpoint()
|
||||
assert g1.checkpoint_input(init) is init # reset after teardown
|
||||
|
||||
# A later run finds the checkpoint -> resume by feeding None, not the
|
||||
# initial state (re-passing it would duplicate messages, #1249).
|
||||
_should_crash = False
|
||||
g2 = _bare_graph(tmp)
|
||||
g2.begin_checkpoint(*args)
|
||||
try:
|
||||
assert g2._resuming is True
|
||||
assert g2.checkpoint_input(init) is None
|
||||
finally:
|
||||
g2.end_checkpoint()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cli_style_usage_saves_then_resumes():
|
||||
global _should_crash
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
cfg_args = ("AAPL", "2026-05-08", "stock")
|
||||
|
||||
# Run 1 (the CLI path): begin -> stream self.graph -> crash at 'trader'.
|
||||
_should_crash = True
|
||||
g1 = _bare_graph(tmp)
|
||||
tid = g1.begin_checkpoint(*cfg_args)
|
||||
args = {"config": {"configurable": {"thread_id": tid}}}
|
||||
try:
|
||||
with pytest.raises(RuntimeError):
|
||||
for _ in g1.graph.stream({"count": 0}, **args):
|
||||
pass
|
||||
finally:
|
||||
g1.end_checkpoint()
|
||||
|
||||
# A checkpoint was saved for this run signature (so --checkpoint works).
|
||||
|
||||
sig = g1._run_signature("stock")
|
||||
assert checkpoint_step(tmp, "AAPL", "2026-05-08", sig) is not None
|
||||
|
||||
# Run 2 (fresh graph, as a new CLI invocation): resume and finish.
|
||||
_should_crash = False
|
||||
g2 = _bare_graph(tmp)
|
||||
tid2 = g2.begin_checkpoint(*cfg_args)
|
||||
assert tid2 == tid # stable id -> same thread resumes
|
||||
try:
|
||||
result = g2.graph.invoke(None, config={"configurable": {"thread_id": tid2}})
|
||||
assert result["count"] == 11 # analyst(+1) resumed into trader(+10)
|
||||
g2.clear_checkpoint_on_success(*cfg_args)
|
||||
finally:
|
||||
g2.end_checkpoint()
|
||||
|
||||
# Cleared on success -> a later run starts fresh.
|
||||
assert checkpoint_step(tmp, "AAPL", "2026-05-08", sig) is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_clearing_removes_the_database_sidecars(tmp_path):
|
||||
"""SQLite writes -wal and -shm next to the database; leaving them behind
|
||||
means a cleared checkpoint still has committed state on disk."""
|
||||
from tradingagents.graph.checkpointer import clear_all_checkpoints
|
||||
|
||||
cp = tmp_path / "checkpoints"
|
||||
cp.mkdir(parents=True)
|
||||
for suffix in (".db", ".db-wal", ".db-shm"):
|
||||
(cp / f"NVDA{suffix}").write_text("x")
|
||||
|
||||
cleared = clear_all_checkpoints(str(tmp_path))
|
||||
|
||||
assert cleared == 1
|
||||
assert list(cp.iterdir()) == []
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Test checkpoint resume: crash mid-analysis, re-run resumes from last node."""
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from typing import TypedDict
|
||||
|
||||
from langgraph.graph import END, StateGraph
|
||||
|
||||
from tradingagents.graph.checkpointer import (
|
||||
checkpoint_step,
|
||||
clear_checkpoint,
|
||||
get_checkpointer,
|
||||
thread_id,
|
||||
)
|
||||
|
||||
# Mutable flag to simulate crash on first run
|
||||
_should_crash = False
|
||||
|
||||
|
||||
class _SimpleState(TypedDict):
|
||||
count: int
|
||||
|
||||
|
||||
def _node_a(state: _SimpleState) -> dict:
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
|
||||
def _node_b(state: _SimpleState) -> dict:
|
||||
if _should_crash:
|
||||
raise RuntimeError("simulated mid-analysis crash")
|
||||
return {"count": state["count"] + 10}
|
||||
|
||||
|
||||
def _build_graph() -> StateGraph:
|
||||
builder = StateGraph(_SimpleState)
|
||||
builder.add_node("analyst", _node_a)
|
||||
builder.add_node("trader", _node_b)
|
||||
builder.set_entry_point("analyst")
|
||||
builder.add_edge("analyst", "trader")
|
||||
builder.add_edge("trader", END)
|
||||
return builder
|
||||
|
||||
|
||||
class TestCheckpointResume(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.tmpdir = tempfile.mkdtemp()
|
||||
self.ticker = "TEST"
|
||||
self.date = "2026-04-20"
|
||||
|
||||
def test_crash_and_resume(self):
|
||||
"""Crash at 'trader' node, then resume from checkpoint."""
|
||||
global _should_crash
|
||||
builder = _build_graph()
|
||||
tid = thread_id(self.ticker, self.date)
|
||||
cfg = {"configurable": {"thread_id": tid}}
|
||||
|
||||
# Run 1: crash at trader node
|
||||
_should_crash = True
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph.invoke({"count": 0}, config=cfg)
|
||||
|
||||
# Checkpoint should exist at step 1 (analyst completed)
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date))
|
||||
step = checkpoint_step(self.tmpdir, self.ticker, self.date)
|
||||
self.assertEqual(step, 1)
|
||||
|
||||
# Run 2: resume — trader succeeds this time
|
||||
_should_crash = False
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
result = graph.invoke(None, config=cfg)
|
||||
|
||||
# analyst added 1, trader added 10 → 11
|
||||
self.assertEqual(result["count"], 11)
|
||||
|
||||
def test_clear_checkpoint_allows_fresh_start(self):
|
||||
"""After clearing, the graph starts from scratch."""
|
||||
global _should_crash
|
||||
builder = _build_graph()
|
||||
tid = thread_id(self.ticker, self.date)
|
||||
cfg = {"configurable": {"thread_id": tid}}
|
||||
|
||||
# Create a checkpoint by crashing
|
||||
_should_crash = True
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph.invoke({"count": 0}, config=cfg)
|
||||
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date))
|
||||
|
||||
# Clear it
|
||||
clear_checkpoint(self.tmpdir, self.ticker, self.date)
|
||||
self.assertIsNone(checkpoint_step(self.tmpdir, self.ticker, self.date))
|
||||
|
||||
# Fresh run succeeds from scratch
|
||||
_should_crash = False
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
result = graph.invoke({"count": 0}, config=cfg)
|
||||
|
||||
self.assertEqual(result["count"], 11)
|
||||
|
||||
def test_different_date_starts_fresh(self):
|
||||
"""A different date must NOT resume from an existing checkpoint."""
|
||||
global _should_crash
|
||||
builder = _build_graph()
|
||||
date2 = "2026-04-21"
|
||||
|
||||
# Run with date1 — crash to leave a checkpoint
|
||||
_should_crash = True
|
||||
tid1 = thread_id(self.ticker, self.date)
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph.invoke({"count": 0}, config={"configurable": {"thread_id": tid1}})
|
||||
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date))
|
||||
|
||||
# date2 should have no checkpoint
|
||||
self.assertIsNone(checkpoint_step(self.tmpdir, self.ticker, date2))
|
||||
|
||||
# Run with date2 — should start fresh and succeed
|
||||
_should_crash = False
|
||||
tid2 = thread_id(self.ticker, date2)
|
||||
self.assertNotEqual(tid1, tid2)
|
||||
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
result = graph.invoke({"count": 0}, config={"configurable": {"thread_id": tid2}})
|
||||
|
||||
# Fresh run: analyst +1, trader +10 = 11
|
||||
self.assertEqual(result["count"], 11)
|
||||
|
||||
# Original date checkpoint still exists (untouched)
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date))
|
||||
|
||||
|
||||
class TestCheckpointSignature(unittest.TestCase):
|
||||
"""A different graph shape (analyst selection / depth / asset mode) must not
|
||||
resume the previous run's checkpoint (#1089)."""
|
||||
|
||||
def setUp(self):
|
||||
self.tmpdir = tempfile.mkdtemp()
|
||||
self.ticker = "TEST"
|
||||
self.date = "2026-04-20"
|
||||
|
||||
def test_empty_signature_is_legacy_id(self):
|
||||
self.assertEqual(
|
||||
thread_id(self.ticker, self.date),
|
||||
thread_id(self.ticker, self.date, ""),
|
||||
)
|
||||
|
||||
def test_signature_changes_thread_id(self):
|
||||
legacy = thread_id(self.ticker, self.date)
|
||||
sig_a = thread_id(self.ticker, self.date, "analysts=market,news|asset=stock")
|
||||
sig_b = thread_id(self.ticker, self.date, "analysts=market|asset=stock")
|
||||
self.assertNotEqual(sig_a, sig_b) # different graph shapes differ
|
||||
self.assertNotEqual(legacy, sig_a) # signature-keyed differs from legacy
|
||||
self.assertEqual( # same inputs are stable
|
||||
sig_a, thread_id(self.ticker, self.date, "analysts=market,news|asset=stock")
|
||||
)
|
||||
|
||||
def test_different_signature_starts_fresh(self):
|
||||
global _should_crash
|
||||
builder = _build_graph()
|
||||
sig1 = "analysts=market,news,fundamentals|asset=stock"
|
||||
sig2 = "analysts=market|asset=stock" # dropped analysts -> different graph
|
||||
|
||||
_should_crash = True
|
||||
tid1 = thread_id(self.ticker, self.date, sig1)
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph.invoke({"count": 0}, config={"configurable": {"thread_id": tid1}})
|
||||
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date, sig1))
|
||||
# A different graph shape has no checkpoint to resume from.
|
||||
self.assertIsNone(checkpoint_step(self.tmpdir, self.ticker, self.date, sig2))
|
||||
|
||||
_should_crash = False
|
||||
tid2 = thread_id(self.ticker, self.date, sig2)
|
||||
self.assertNotEqual(tid1, tid2)
|
||||
with get_checkpointer(self.tmpdir, self.ticker) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
result = graph.invoke({"count": 0}, config={"configurable": {"thread_id": tid2}})
|
||||
self.assertEqual(result["count"], 11)
|
||||
# sig1's checkpoint remains untouched.
|
||||
self.assertIsNotNone(checkpoint_step(self.tmpdir, self.ticker, self.date, sig1))
|
||||
|
||||
def test_run_signature_captures_graph_shape(self):
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
# Build a bare instance to exercise the pure helper without heavy __init__.
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
g.selected_analysts = ("market", "news")
|
||||
g.config = {"max_debate_rounds": 1, "max_risk_discuss_rounds": 1}
|
||||
base = g._run_signature("stock")
|
||||
|
||||
self.assertNotEqual(base, g._run_signature("crypto")) # asset mode
|
||||
g.selected_analysts = ("market",)
|
||||
self.assertNotEqual(base, g._run_signature("stock")) # analyst selection
|
||||
g.selected_analysts = ("market", "news")
|
||||
g.config = {"max_debate_rounds": 3, "max_risk_discuss_rounds": 1}
|
||||
self.assertNotEqual(base, g._run_signature("stock")) # debate depth
|
||||
g.config = {"max_debate_rounds": 1, "max_risk_discuss_rounds": 5}
|
||||
self.assertNotEqual(base, g._run_signature("stock")) # risk depth
|
||||
# Stable for identical inputs.
|
||||
g.config = {"max_debate_rounds": 1, "max_risk_discuss_rounds": 1}
|
||||
self.assertEqual(base, g._run_signature("stock"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,136 @@
|
||||
"""The CLI keeps running an analysis with no arguments, and gains `backtest`.
|
||||
|
||||
Every documented invocation is bare (`tradingagents --checkpoint`), so analysis
|
||||
has to stay the default action while a second command exists alongside it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from typer.testing import CliRunner
|
||||
|
||||
import cli.main as m
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def runner(monkeypatch):
|
||||
monkeypatch.setattr(m, "run_analysis", lambda **kw: calls.append(("analysis", kw)))
|
||||
calls.clear()
|
||||
return CliRunner()
|
||||
|
||||
|
||||
calls: list = []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_no_arguments_still_runs_an_analysis(runner):
|
||||
assert runner.invoke(m.app, []).exit_code == 0
|
||||
assert calls == [("analysis", {"checkpoint": None, "portfolio": None})]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_options_still_parse_without_a_subcommand(runner):
|
||||
assert runner.invoke(m.app, ["--checkpoint"]).exit_code == 0
|
||||
assert calls[0][1]["checkpoint"] is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_backtest_does_not_also_run_an_analysis(runner, monkeypatch, tmp_path):
|
||||
swept = []
|
||||
monkeypatch.setattr(m, "run_backtest", lambda *a, **kw: swept.append((a, kw)) or _Result(tmp_path))
|
||||
monkeypatch.setattr(m, "summarize", lambda log: _Summary())
|
||||
|
||||
result = runner.invoke(m.app, ["backtest", "NVDA,AAPL", "--start", "2026-06-01",
|
||||
"--end", "2026-06-15", "--every", "7"])
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert calls == [] # the interactive analysis must not run
|
||||
(tickers, dates, _config), kwargs = swept[0]
|
||||
assert tickers == ["NVDA", "AAPL"]
|
||||
assert dates == ["2026-06-01", "2026-06-08", "2026-06-15"]
|
||||
assert "scored" in result.output
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_backtest_reports_a_bad_date_instead_of_a_traceback(runner):
|
||||
result = runner.invoke(m.app, ["backtest", "NVDA", "--start", "June", "--end", "2026-06-15"])
|
||||
assert result.exit_code == 1
|
||||
assert "YYYY-MM-DD" in result.output
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_help_lists_the_backtest_command(runner):
|
||||
assert "backtest" in runner.invoke(m.app, ["--help"]).output
|
||||
|
||||
|
||||
class _Result:
|
||||
def __init__(self, tmp_path):
|
||||
self.run_id = "20260916_000000"
|
||||
self.log_path = tmp_path / "trading_memory.md"
|
||||
self.cells_run = 2
|
||||
self.skipped = 0
|
||||
self.failures = []
|
||||
self.settlement_failures = []
|
||||
|
||||
|
||||
class _Summary:
|
||||
def render(self):
|
||||
return "scored 2 cells"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_every_command_is_registered_when_run_as_a_module():
|
||||
"""README documents `python -m cli.main`, which executes the file top to
|
||||
bottom, so a command defined after the __main__ block would not exist."""
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
out = subprocess.run([sys.executable, "-m", "cli.main", "backtest", "--help"],
|
||||
capture_output=True, text=True, timeout=120)
|
||||
assert out.returncode == 0, out.stderr[-400:]
|
||||
# Where the terminal takes colour, help styles each option and splits
|
||||
# "--start" across escape sequences, so read the text without them.
|
||||
plain = re.sub(r"\x1b\[[0-9;]*m", "", out.stdout)
|
||||
assert "--start" in plain
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_backtest_can_continue_an_interrupted_sweep(runner, monkeypatch, tmp_path):
|
||||
"""Resuming is what makes a long sweep practical, and the Python API has it."""
|
||||
swept = []
|
||||
monkeypatch.setattr(m, "run_backtest", lambda *a, **kw: swept.append(kw) or _Result(tmp_path))
|
||||
monkeypatch.setattr(m, "summarize", lambda log: _Summary())
|
||||
|
||||
result = runner.invoke(m.app, ["backtest", "NVDA", "--start", "2026-06-01",
|
||||
"--end", "2026-06-08", "--run-id", "20260617_120000"])
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert swept[0]["run_id"] == "20260617_120000"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("args, expected", [
|
||||
(["backtest", "NVDA", "--start", "2026-08-01", "--end", "2026-06-08"], "before"),
|
||||
(["backtest", ",,", "--start", "2026-06-01", "--end", "2026-06-08"], "ticker"),
|
||||
])
|
||||
def test_backtest_rejects_input_that_would_sweep_nothing(runner, args, expected):
|
||||
"""An inverted range or an empty ticker list reported a clean zero-cell run,
|
||||
which reads as 'nothing to find' rather than 'you asked for nothing'."""
|
||||
result = runner.invoke(m.app, args)
|
||||
assert result.exit_code == 1
|
||||
assert expected in result.output.lower()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_backtest_reports_a_setup_failure_in_one_line(runner, monkeypatch):
|
||||
"""A missing key or a bad analyst name produced a raw traceback."""
|
||||
def _explode(*a, **kw):
|
||||
raise ValueError("API key for provider 'openai' is not set")
|
||||
|
||||
monkeypatch.setattr(m, "run_backtest", _explode)
|
||||
result = runner.invoke(m.app, ["backtest", "NVDA", "--start", "2026-06-01", "--end", "2026-06-08"])
|
||||
|
||||
assert result.exit_code == 1
|
||||
assert "API key" in result.output
|
||||
assert "Traceback" not in result.output
|
||||
@@ -0,0 +1,107 @@
|
||||
"""CLI config precedence (#976, #977).
|
||||
|
||||
An explicit environment override for the debate/risk round counts, or the
|
||||
checkpoint flag, must win over the interactive research-depth selection — the CLI
|
||||
must not clobber an env-configured value back to a prompt/flag default.
|
||||
"""
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
import cli.main as m
|
||||
import cli.run as cli_run
|
||||
|
||||
# Minimal selections dict shaped like get_user_selections()'s return value.
|
||||
SELECTIONS = {
|
||||
"research_depth": 5,
|
||||
"quick_think_llm": "gpt-5.4-mini",
|
||||
"deep_think_llm": "gpt-5.5",
|
||||
"backend_url": None,
|
||||
"llm_provider": "openai",
|
||||
"google_thinking_level": None,
|
||||
"openai_reasoning_effort": None,
|
||||
"anthropic_effort": None,
|
||||
"output_language": "English",
|
||||
}
|
||||
|
||||
|
||||
def test_research_depth_sets_both_rounds_without_env(monkeypatch):
|
||||
for var in ("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "TRADINGAGENTS_MAX_RISK_ROUNDS"):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None)
|
||||
assert cfg["max_debate_rounds"] == 5
|
||||
assert cfg["max_risk_discuss_rounds"] == 5
|
||||
|
||||
|
||||
def test_env_round_counts_win_over_selection(monkeypatch):
|
||||
monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "2")
|
||||
monkeypatch.setenv("TRADINGAGENTS_MAX_RISK_ROUNDS", "4")
|
||||
# DEFAULT_CONFIG already reflects the env (applied at import); emulate that.
|
||||
patched = dict(cli_run.DEFAULT_CONFIG, max_debate_rounds=2, max_risk_discuss_rounds=4)
|
||||
with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched):
|
||||
cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None)
|
||||
assert cfg["max_debate_rounds"] == 2 # env value, not research_depth=5
|
||||
assert cfg["max_risk_discuss_rounds"] == 4
|
||||
|
||||
|
||||
def test_partial_env_only_overrides_that_count(monkeypatch):
|
||||
monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "2")
|
||||
monkeypatch.delenv("TRADINGAGENTS_MAX_RISK_ROUNDS", raising=False)
|
||||
patched = dict(cli_run.DEFAULT_CONFIG, max_debate_rounds=2)
|
||||
with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched):
|
||||
cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None)
|
||||
assert cfg["max_debate_rounds"] == 2 # env wins
|
||||
assert cfg["max_risk_discuss_rounds"] == 5 # falls through to research_depth
|
||||
|
||||
|
||||
def test_checkpoint_none_preserves_env_default():
|
||||
patched = dict(cli_run.DEFAULT_CONFIG, checkpoint_enabled=True) # e.g. env-enabled
|
||||
with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched):
|
||||
cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None)
|
||||
assert cfg["checkpoint_enabled"] is True # not clobbered back to False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("flag", [True, False])
|
||||
def test_checkpoint_flag_overrides_env(flag):
|
||||
patched = dict(cli_run.DEFAULT_CONFIG, checkpoint_enabled=not flag)
|
||||
with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched):
|
||||
cfg = cli_run._build_run_config(SELECTIONS, checkpoint=flag)
|
||||
assert cfg["checkpoint_enabled"] is flag
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_glm_resolves_to_the_endpoint_its_key_belongs_to():
|
||||
"""The provider table, the client registry and the key mapping must name the
|
||||
same platform: glm is Z.AI international (ZHIPU_API_KEY) and glm-cn is
|
||||
BigModel China. A mismatch sends the key to the other platform and every
|
||||
call fails auth."""
|
||||
from cli.prompts import resolve_backend_url
|
||||
from tradingagents.llm_clients.api_key_env import get_api_key_env
|
||||
from tradingagents.llm_clients.openai_client import OPENAI_COMPATIBLE_PROVIDERS
|
||||
|
||||
assert resolve_backend_url("glm", None, None) == OPENAI_COMPATIBLE_PROVIDERS["glm"].base_url
|
||||
assert get_api_key_env("glm") == "ZHIPU_API_KEY"
|
||||
assert "z.ai" in OPENAI_COMPATIBLE_PROVIDERS["glm"].base_url
|
||||
assert "bigmodel.cn" in OPENAI_COMPATIBLE_PROVIDERS["glm-cn"].base_url
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_half_set_round_count_says_which_value_won(capsys, monkeypatch):
|
||||
"""With only one of the two round-count variables set, the depth prompt is
|
||||
still shown but half the answer is discarded; the user was never told."""
|
||||
|
||||
monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "1")
|
||||
monkeypatch.delenv("TRADINGAGENTS_MAX_RISK_ROUNDS", raising=False)
|
||||
printed = []
|
||||
monkeypatch.setattr(m.console, "print", lambda *a, **k: printed.append(str(a[0]) if a else ""))
|
||||
|
||||
config = cli_run._build_run_config({
|
||||
"ticker": "NVDA", "analysis_date": "2026-09-01", "asset_type": "stock",
|
||||
"analysts": [], "research_depth": 5, "llm_provider": "openai",
|
||||
"quick_think_llm": "gpt-5.6-luna", "deep_think_llm": "gpt-5.6",
|
||||
"backend_url": None, "output_language": "English",
|
||||
}, None)
|
||||
|
||||
assert config["max_risk_discuss_rounds"] == 5
|
||||
assert any("TRADINGAGENTS_MAX_DEBATE_ROUNDS" in line for line in printed), printed
|
||||
@@ -0,0 +1,193 @@
|
||||
"""The CLI must use the decision log the same way propagate() does.
|
||||
|
||||
The CLI streams the graph itself instead of calling propagate(), so memory steps
|
||||
that lived only in propagate() never ran on the primary entry point: pending
|
||||
decisions were not settled, the Portfolio Manager got no past context, and the
|
||||
finished decision was not recorded. Both paths now build their initial state and
|
||||
record their decision through the same graph methods.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
import cli.run as cli_run
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
|
||||
def _bare_graph(tmp_path):
|
||||
"""A graph without __init__ (no LLM clients), wired to a temp log."""
|
||||
graph = object.__new__(TradingAgentsGraph)
|
||||
graph.config = {"memory_log_path": str(tmp_path / "trading_memory.md")}
|
||||
graph.memory_log = TradingMemoryLog(graph.config)
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_create_run_state_settles_pending_and_carries_context(tmp_path, monkeypatch):
|
||||
from tradingagents.graph.propagation import Propagator
|
||||
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.propagator = Propagator()
|
||||
settled = []
|
||||
monkeypatch.setattr(graph, "settle_pending", settled.append, raising=False)
|
||||
monkeypatch.setattr(graph, "resolve_instrument_context", lambda t, a="stock", d=None: f"id:{t}", raising=False)
|
||||
monkeypatch.setattr(graph, "_memory_as_of", lambda d: d, raising=False)
|
||||
graph.memory_log.store_decision("NVDA", "2026-01-05", "Rating: Buy\nold call")
|
||||
graph.memory_log.update_with_outcome("NVDA", "2026-01-05", 0.01, 0.005, 5, "great trade", "2026-01-12")
|
||||
|
||||
state = graph.create_run_state("NVDA", "2026-02-01")
|
||||
|
||||
assert settled == ["NVDA"]
|
||||
assert "great trade" in state["past_context"]
|
||||
assert state["instrument_context"] == "id:NVDA"
|
||||
assert state["company_of_interest"] == "NVDA"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_record_decision_appends_a_pending_entry(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.record_decision("NVDA", "2026-01-10", {"final_trade_decision": "Rating: Buy\n\nBuy NVDA."})
|
||||
entries = graph.memory_log.load_entries()
|
||||
assert [(e["ticker"], e["pending"], e["rating"]) for e in entries] == [("NVDA", True, "Buy")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_record_decision_skips_a_run_without_a_decision(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.record_decision("NVDA", "2026-01-10", {})
|
||||
assert graph.memory_log.load_entries() == []
|
||||
|
||||
|
||||
# --- the CLI path ----------------------------------------------------------------
|
||||
|
||||
class _FakeGraph:
|
||||
"""Records the lifecycle calls run_analysis makes."""
|
||||
|
||||
def __init__(self, resuming=None):
|
||||
self.calls = []
|
||||
self.graph = self
|
||||
self.propagator = self
|
||||
self.resuming = resuming # None: checkpointing off
|
||||
self._resuming = False
|
||||
|
||||
def create_run_state(self, ticker, trade_date, asset_type="stock", portfolio=None):
|
||||
self.calls.append(("create_run_state", ticker, trade_date))
|
||||
return {"messages": [], "company_of_interest": ticker}
|
||||
|
||||
def process_signal(self, text):
|
||||
from tradingagents.agents.rating import parse_rating
|
||||
return parse_rating(text)
|
||||
|
||||
def record_decision(self, ticker, trade_date, final_state):
|
||||
self.calls.append(("record_decision", ticker, trade_date, final_state.get("final_trade_decision")))
|
||||
|
||||
def get_graph_args(self, callbacks=None):
|
||||
return {}
|
||||
|
||||
def begin_checkpoint(self, *a, **k):
|
||||
self._resuming = bool(self.resuming)
|
||||
return None if self.resuming is None else "thread"
|
||||
|
||||
def checkpoint_input(self, state):
|
||||
return state
|
||||
|
||||
def clear_checkpoint_on_success(self, *a, **k):
|
||||
self.calls.append(("clear_checkpoint",))
|
||||
|
||||
def end_checkpoint(self):
|
||||
pass
|
||||
|
||||
def stream(self, graph_input, **kwargs):
|
||||
yield {"messages": [], "market_report": "M"}
|
||||
yield {"messages": [], "final_trade_decision": "Rating: Buy\n\nBuy NVDA."}
|
||||
|
||||
|
||||
class _NullLive:
|
||||
def __init__(self, *a, **k):
|
||||
pass
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
|
||||
class _FakeBuffer:
|
||||
def __init__(self):
|
||||
self.messages = []
|
||||
self.tool_calls = []
|
||||
self.report_sections = {}
|
||||
self.agent_status = {}
|
||||
self.selected_analysts = []
|
||||
self._processed_message_ids = set()
|
||||
|
||||
def init_for_analysis(self, selected_analysts):
|
||||
self.selected_analysts = [a.lower() for a in selected_analysts]
|
||||
|
||||
def add_message(self, kind, content):
|
||||
self.messages.append((0.0, kind, content))
|
||||
|
||||
def add_tool_call(self, name, args):
|
||||
self.tool_calls.append((0.0, name, args))
|
||||
|
||||
def update_report_section(self, *a):
|
||||
pass
|
||||
|
||||
def update_agent_status(self, agent, status):
|
||||
self.agent_status[agent] = status
|
||||
|
||||
|
||||
def _run_cli(monkeypatch, tmp_path, fake):
|
||||
"""Drive run_analysis against ``fake``; returns the message buffer."""
|
||||
import cli.main as m
|
||||
from cli.models import AnalystType
|
||||
|
||||
buffer = _FakeBuffer()
|
||||
monkeypatch.setattr(cli_run, "TradingAgentsGraph", lambda *a, **k: fake)
|
||||
monkeypatch.setattr(cli_run, "message_buffer", buffer)
|
||||
monkeypatch.setattr(cli_run, "create_layout", lambda: None)
|
||||
monkeypatch.setattr(cli_run, "update_display", lambda *a, **k: None)
|
||||
monkeypatch.setattr(cli_run, "Live", _NullLive)
|
||||
monkeypatch.setattr(cli_run, "get_user_selections", lambda: {
|
||||
"ticker": "NVDA", "analysis_date": "2026-01-10",
|
||||
"analysts": [AnalystType.MARKET], "asset_type": "stock",
|
||||
})
|
||||
monkeypatch.setattr(cli_run, "_build_run_config", lambda selections, checkpoint: {
|
||||
"data_cache_dir": str(tmp_path / "cache"), "results_dir": str(tmp_path / "results"),
|
||||
})
|
||||
monkeypatch.setattr(m.typer, "prompt", lambda *a, **k: "N")
|
||||
cli_run.run_analysis()
|
||||
return buffer
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cli_run_uses_the_decision_log_like_propagate(tmp_path, monkeypatch):
|
||||
fake = _FakeGraph()
|
||||
_run_cli(monkeypatch, tmp_path, fake)
|
||||
|
||||
assert fake.calls == [
|
||||
("create_run_state", "NVDA", "2026-01-10"),
|
||||
# The decision is recorded from the merged stream, before the checkpoint
|
||||
# is cleared, matching propagate().
|
||||
("record_decision", "NVDA", "2026-01-10", "Rating: Buy\n\nBuy NVDA."),
|
||||
("clear_checkpoint",),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("resuming, said", [(True, "resuming"), (False, "starting fresh")])
|
||||
def test_the_cli_run_says_whether_it_resumed(tmp_path, monkeypatch, resuming, said):
|
||||
"""The README promises the run view tells a resumed run from a fresh one."""
|
||||
buffer = _run_cli(monkeypatch, tmp_path, _FakeGraph(resuming=resuming))
|
||||
|
||||
assert any(said in text.lower() for _, kind, text in buffer.messages if kind == "System")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_run_without_checkpointing_says_nothing_about_resuming(tmp_path, monkeypatch):
|
||||
buffer = _run_cli(monkeypatch, tmp_path, _FakeGraph())
|
||||
|
||||
assert not any("resum" in text.lower() or "fresh" in text.lower() for _, _, text in buffer.messages)
|
||||
@@ -0,0 +1,134 @@
|
||||
"""What the live display shows, and what the run log keeps.
|
||||
|
||||
The display drops a message it judges empty, and the state log is written for a
|
||||
person to read afterwards. Both got that wrong in ways that hide real content.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import unittest
|
||||
|
||||
import pytest
|
||||
|
||||
from cli.display import (
|
||||
AnalystWallTimeTracker,
|
||||
extract_content_string,
|
||||
sync_analyst_tracker_from_chunk,
|
||||
)
|
||||
from tradingagents.graph.analyst_execution import build_analyst_execution_plan
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("text", ["0", "False", "None", "[]", "{}", "0.0"])
|
||||
def test_a_message_that_reads_like_a_python_value_is_still_text(text):
|
||||
"""These were parsed as Python and judged empty, so the message vanished."""
|
||||
assert extract_content_string(text) == text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("value, expected", [
|
||||
(" Hold ", "Hold"),
|
||||
("", None),
|
||||
(" ", None),
|
||||
(None, None),
|
||||
([], None),
|
||||
({}, None),
|
||||
({"text": "from a dict"}, "from a dict"),
|
||||
([{"type": "text", "text": "part one"}, {"type": "text", "text": "part two"}], "part one part two"),
|
||||
])
|
||||
def test_the_other_shapes_are_unchanged(value, expected):
|
||||
assert extract_content_string(value) == expected
|
||||
|
||||
|
||||
def _state(ticker, final="评级: 买入"):
|
||||
return {
|
||||
"company_of_interest": ticker, "trade_date": "2026-09-01",
|
||||
"market_report": "市场", "sentiment_report": "情绪", "news_report": "新闻",
|
||||
"fundamentals_report": "基本面", "investment_plan": "计划",
|
||||
"trader_investment_plan": "交易计划", "final_trade_decision": final,
|
||||
"investment_debate_state": {"bull_history": "", "bear_history": "", "history": "",
|
||||
"current_response": "", "judge_decision": "", "count": 0},
|
||||
"risk_debate_state": {"aggressive_history": "", "conservative_history": "",
|
||||
"neutral_history": "", "history": "", "judge_decision": "",
|
||||
"latest_speaker": "", "current_aggressive_response": "",
|
||||
"current_conservative_response": "", "current_neutral_response": "",
|
||||
"count": 0},
|
||||
}
|
||||
|
||||
|
||||
def _bare_graph(tmp_path):
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
graph = object.__new__(TradingAgentsGraph)
|
||||
graph.config = {"results_dir": str(tmp_path)}
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_state_log_keeps_non_ascii_readable(tmp_path):
|
||||
"""Reports can be in any language; the log is read by a person."""
|
||||
_bare_graph(tmp_path)._log_state("2026-09-01", _state("600519.SS"))
|
||||
|
||||
written = next(tmp_path.rglob("full_states_log*.json")).read_text(encoding="utf-8")
|
||||
assert "买入" in written
|
||||
assert "\\u" not in written
|
||||
assert json.loads(written) # still valid JSON
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_live_display_does_not_scroll_the_terminal():
|
||||
"""A layout taller than the window makes rich redraw by scrolling, which
|
||||
reads as flicker; the alternate screen holds it in place (#784). The final
|
||||
report prints after the live view ends, so nothing is lost when it closes."""
|
||||
import inspect
|
||||
|
||||
import cli.main as m
|
||||
|
||||
assert "screen=True" in inspect.getsource(m.run_analysis)
|
||||
|
||||
|
||||
class AnalystWallTimeTrackerTests(unittest.TestCase):
|
||||
def test_records_wall_time_when_analyst_completes(self):
|
||||
plan = build_analyst_execution_plan(["market", "news"])
|
||||
tracker = AnalystWallTimeTracker(plan)
|
||||
|
||||
tracker.mark_started("market", started_at=10.0)
|
||||
tracker.mark_completed("market", completed_at=13.5)
|
||||
|
||||
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.50s")
|
||||
|
||||
def test_formats_summary_in_plan_order(self):
|
||||
plan = build_analyst_execution_plan(["news", "market"])
|
||||
tracker = AnalystWallTimeTracker(plan)
|
||||
|
||||
tracker.mark_started("market", started_at=20.0)
|
||||
tracker.mark_completed("market", completed_at=22.25)
|
||||
tracker.mark_started("news", started_at=10.0)
|
||||
tracker.mark_completed("news", completed_at=14.0)
|
||||
|
||||
self.assertEqual(
|
||||
tracker.format_summary(),
|
||||
"Analyst wall time: News 4.00s | Market 2.25s",
|
||||
)
|
||||
|
||||
def test_syncs_wall_time_from_sequential_chunks(self):
|
||||
plan = build_analyst_execution_plan(["market", "news"])
|
||||
tracker = AnalystWallTimeTracker(plan)
|
||||
|
||||
sync_analyst_tracker_from_chunk(tracker, {}, now=10.0)
|
||||
self.assertEqual(tracker.format_summary(), "Analyst wall time: pending")
|
||||
|
||||
sync_analyst_tracker_from_chunk(
|
||||
tracker,
|
||||
{"market_report": "done"},
|
||||
now=13.0,
|
||||
)
|
||||
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s")
|
||||
|
||||
sync_analyst_tracker_from_chunk(
|
||||
tracker,
|
||||
{"market_report": "done", "news_report": "done"},
|
||||
now=18.0,
|
||||
)
|
||||
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s | News 5.00s")
|
||||
@@ -0,0 +1,148 @@
|
||||
"""Tests for env-driven CLI behavior (#897, #873).
|
||||
|
||||
The config-layer override (TRADINGAGENTS_* -> DEFAULT_CONFIG) is covered by
|
||||
test_env_overrides.py. These tests cover the CLI layer: an env-configured
|
||||
provider/model/language must skip its interactive prompt and use the value.
|
||||
"""
|
||||
|
||||
import os
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
import cli.selections as cli_selections
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestProviderDefaultUrl(unittest.TestCase):
|
||||
def test_known_providers_resolve(self):
|
||||
from cli.prompts import provider_default_url
|
||||
self.assertEqual(provider_default_url("openai"), "https://api.openai.com/v1")
|
||||
self.assertEqual(provider_default_url("DeepSeek"), "https://api.deepseek.com")
|
||||
self.assertIsNone(provider_default_url("google")) # uses SDK default
|
||||
|
||||
def test_unknown_provider_returns_none(self):
|
||||
from cli.prompts import provider_default_url
|
||||
self.assertIsNone(provider_default_url("not-a-provider"))
|
||||
|
||||
def test_ollama_honors_base_url_env(self):
|
||||
from cli.prompts import provider_default_url
|
||||
with mock.patch.dict(os.environ, {"OLLAMA_BASE_URL": "http://host:1234/v1"}):
|
||||
self.assertEqual(provider_default_url("ollama"), "http://host:1234/v1")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCliSkipsPromptsFromEnv(unittest.TestCase):
|
||||
def test_env_config_skips_llm_prompts(self):
|
||||
|
||||
env = {
|
||||
"TRADINGAGENTS_LLM_PROVIDER": "openai",
|
||||
"TRADINGAGENTS_DEEP_THINK_LLM": "kimi-k2.5",
|
||||
"TRADINGAGENTS_QUICK_THINK_LLM": "deepseek-v4-pro",
|
||||
"TRADINGAGENTS_LLM_BACKEND_URL": "https://opencode.ai/zen/go/v1",
|
||||
"TRADINGAGENTS_OUTPUT_LANGUAGE": "Japanese",
|
||||
}
|
||||
fake_cfg = dict(cli_selections.DEFAULT_CONFIG)
|
||||
fake_cfg.update({
|
||||
"llm_provider": "openai",
|
||||
"backend_url": "https://opencode.ai/zen/go/v1",
|
||||
"quick_think_llm": "deepseek-v4-pro",
|
||||
"deep_think_llm": "kimi-k2.5",
|
||||
"output_language": "Japanese",
|
||||
})
|
||||
|
||||
with mock.patch.dict(os.environ, env, clear=False), \
|
||||
mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \
|
||||
mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \
|
||||
mock.patch.object(cli_selections, "display_announcements"), \
|
||||
mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \
|
||||
mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \
|
||||
mock.patch.object(cli_selections, "select_analysts", return_value=[]), \
|
||||
mock.patch.object(cli_selections, "select_research_depth", return_value=1), \
|
||||
mock.patch.object(cli_selections, "ensure_api_key") as ensure_key, \
|
||||
mock.patch.object(cli_selections, "select_llm_provider") as prompt_provider, \
|
||||
mock.patch.object(cli_selections, "ask_output_language") as prompt_lang, \
|
||||
mock.patch.object(cli_selections, "select_shallow_thinking_agent") as prompt_quick, \
|
||||
mock.patch.object(cli_selections, "select_deep_thinking_agent") as prompt_deep:
|
||||
sel = cli_selections.get_user_selections()
|
||||
|
||||
# None of the LLM selection prompts should have been shown.
|
||||
prompt_provider.assert_not_called()
|
||||
prompt_lang.assert_not_called()
|
||||
prompt_quick.assert_not_called()
|
||||
prompt_deep.assert_not_called()
|
||||
# API key is still verified for the env-configured provider.
|
||||
ensure_key.assert_called_once()
|
||||
|
||||
# The env values flow into the returned selections.
|
||||
self.assertEqual(sel["llm_provider"], "openai")
|
||||
self.assertEqual(sel["backend_url"], "https://opencode.ai/zen/go/v1")
|
||||
self.assertEqual(sel["quick_think_llm"], "deepseek-v4-pro")
|
||||
self.assertEqual(sel["deep_think_llm"], "kimi-k2.5")
|
||||
self.assertEqual(sel["output_language"], "Japanese")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestResearchDepthSkippedFromEnv(unittest.TestCase):
|
||||
def test_both_round_envs_skip_depth_prompt(self):
|
||||
|
||||
env = {
|
||||
"TRADINGAGENTS_MAX_DEBATE_ROUNDS": "2",
|
||||
"TRADINGAGENTS_MAX_RISK_ROUNDS": "4",
|
||||
}
|
||||
fake_cfg = dict(cli_selections.DEFAULT_CONFIG)
|
||||
fake_cfg.update({"max_debate_rounds": 2, "max_risk_discuss_rounds": 4})
|
||||
|
||||
with mock.patch.dict(os.environ, env, clear=False), \
|
||||
mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \
|
||||
mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \
|
||||
mock.patch.object(cli_selections, "display_announcements"), \
|
||||
mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \
|
||||
mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \
|
||||
mock.patch.object(cli_selections, "select_analysts", return_value=[]), \
|
||||
mock.patch.object(cli_selections, "select_research_depth") as prompt_depth, \
|
||||
mock.patch.object(cli_selections, "ensure_api_key"), \
|
||||
mock.patch.object(cli_selections, "select_llm_provider", return_value=("openai", None)), \
|
||||
mock.patch.object(cli_selections, "ask_output_language", return_value="English"), \
|
||||
mock.patch.object(cli_selections, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \
|
||||
mock.patch.object(cli_selections, "select_deep_thinking_agent", return_value="gpt-5.5"), \
|
||||
mock.patch.object(cli_selections, "ask_openai_reasoning_effort", return_value=None):
|
||||
sel = cli_selections.get_user_selections()
|
||||
|
||||
# The research-depth prompt is skipped; the value comes from the env config.
|
||||
prompt_depth.assert_not_called()
|
||||
self.assertEqual(sel["research_depth"], 2)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestReasoningEffortSkippedFromEnv(unittest.TestCase):
|
||||
def test_effort_env_skips_step8_prompt(self):
|
||||
|
||||
env = {"TRADINGAGENTS_OPENAI_REASONING_EFFORT": "high"}
|
||||
fake_cfg = dict(cli_selections.DEFAULT_CONFIG)
|
||||
fake_cfg.update({"openai_reasoning_effort": "high"})
|
||||
|
||||
with mock.patch.dict(os.environ, env, clear=False), \
|
||||
mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \
|
||||
mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \
|
||||
mock.patch.object(cli_selections, "display_announcements"), \
|
||||
mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \
|
||||
mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \
|
||||
mock.patch.object(cli_selections, "select_analysts", return_value=[]), \
|
||||
mock.patch.object(cli_selections, "select_research_depth", return_value=1), \
|
||||
mock.patch.object(cli_selections, "ensure_api_key"), \
|
||||
mock.patch.object(cli_selections, "select_llm_provider", return_value=("openai", None)), \
|
||||
mock.patch.object(cli_selections, "ask_output_language", return_value="English"), \
|
||||
mock.patch.object(cli_selections, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \
|
||||
mock.patch.object(cli_selections, "select_deep_thinking_agent", return_value="gpt-5.5"), \
|
||||
mock.patch.object(cli_selections, "ask_openai_reasoning_effort") as prompt_effort:
|
||||
sel = cli_selections.get_user_selections()
|
||||
|
||||
# The reasoning-effort prompt is skipped; the value comes from env config.
|
||||
prompt_effort.assert_not_called()
|
||||
self.assertEqual(sel["openai_reasoning_effort"], "high")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,57 @@
|
||||
"""A terminal without a console buffer must fail with one actionable line (#1138).
|
||||
|
||||
prompt_toolkit raises NoConsoleScreenBufferError before the first prompt in
|
||||
non-interactive Windows terminals; the CLI should not surface that traceback.
|
||||
The Windows-only exception import must also stay inert on other platforms.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
|
||||
from typer.testing import CliRunner
|
||||
|
||||
import cli.main as m
|
||||
|
||||
|
||||
def test_no_console_error_tuple_matches_platform():
|
||||
# Off Windows the win32 module is never imported (it asserts the platform),
|
||||
# so the tuple is empty — which `except` accepts and never matches. On
|
||||
# Windows it holds the real exception type, and a broken prompt_toolkit
|
||||
# would raise at import rather than silently disabling the handler.
|
||||
assert isinstance(m._NO_CONSOLE_ERRORS, tuple)
|
||||
assert all(issubclass(e, BaseException) for e in m._NO_CONSOLE_ERRORS)
|
||||
if sys.platform == "win32":
|
||||
assert m._NO_CONSOLE_ERRORS, "Windows must resolve the console error type"
|
||||
else:
|
||||
assert m._NO_CONSOLE_ERRORS == ()
|
||||
|
||||
|
||||
def test_missing_console_prints_actionable_message(monkeypatch):
|
||||
class _NoConsole(Exception):
|
||||
pass
|
||||
|
||||
# Simulate the Windows failure on any platform by registering a stand-in.
|
||||
monkeypatch.setattr(m, "_NO_CONSOLE_ERRORS", (_NoConsole,))
|
||||
|
||||
def _boom(*a, **k):
|
||||
raise _NoConsole("No Windows console found. Are you running cmd.exe?")
|
||||
|
||||
monkeypatch.setattr(m, "run_analysis", _boom)
|
||||
|
||||
result = CliRunner().invoke(m.app, [])
|
||||
assert result.exit_code == 1
|
||||
assert "no Windows console available" in result.output
|
||||
# The raw prompt_toolkit traceback must not reach the user.
|
||||
assert "Traceback" not in result.output
|
||||
|
||||
|
||||
def test_unrelated_errors_still_propagate(monkeypatch):
|
||||
# The handler must stay narrow: only the console error is translated.
|
||||
monkeypatch.setattr(m, "_NO_CONSOLE_ERRORS", (RuntimeError,))
|
||||
|
||||
def _boom(*a, **k):
|
||||
raise ValueError("unrelated")
|
||||
|
||||
monkeypatch.setattr(m, "run_analysis", _boom)
|
||||
result = CliRunner().invoke(m.app, [])
|
||||
assert isinstance(result.exception, ValueError)
|
||||
@@ -0,0 +1,169 @@
|
||||
"""The CLI remembers what you chose last time and offers it back.
|
||||
|
||||
Prefill only: every prompt still appears, so a run never starts on a choice the
|
||||
user did not see. Environment variables keep skipping their step outright and
|
||||
win over anything remembered. Remembered values are validated against the
|
||||
current choices each time, since models and providers come and go between
|
||||
versions and a stale one must not be offered.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
import cli.selections as cli_selections
|
||||
from cli.models import AnalystType
|
||||
from cli.prefs import load_last_run, sanitize, save_last_run
|
||||
|
||||
SAVED = {
|
||||
"output_language": "English",
|
||||
"analysts": ["market", "fundamentals"],
|
||||
"research_depth": 3,
|
||||
"llm_provider": "openai",
|
||||
"quick_think_llm": "gpt-5.6-mini",
|
||||
"deep_think_llm": "gpt-5.6",
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _home(tmp_path, monkeypatch):
|
||||
monkeypatch.setattr("cli.prefs._PREFS_PATH", tmp_path / "cli_prefs.json")
|
||||
return tmp_path
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_round_trip():
|
||||
save_last_run(SAVED)
|
||||
assert load_last_run() == SAVED
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_missing_file_is_not_an_error():
|
||||
assert load_last_run() == {}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_corrupt_file_degrades_to_no_memory(_home):
|
||||
(_home / "cli_prefs.json").write_text("{not json")
|
||||
assert load_last_run() == {}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_half_written_file_cannot_be_observed(_home):
|
||||
"""Two runs finishing together must never leave a torn file behind."""
|
||||
save_last_run(SAVED)
|
||||
save_last_run({**SAVED, "research_depth": 5})
|
||||
assert load_last_run()["research_depth"] == 5
|
||||
assert list((_home).glob("*.tmp*")) == []
|
||||
|
||||
|
||||
# --- validation against the current choices ---------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_model_that_no_longer_exists_is_dropped():
|
||||
# gpt-5.4 is still accepted by config, but is no longer in the picker's list.
|
||||
kept = sanitize({**SAVED, "quick_think_llm": "gpt-5.4"}, "stock")
|
||||
assert "quick_think_llm" not in kept
|
||||
assert kept["deep_think_llm"] == "gpt-5.6" # the valid sibling survives
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_unknown_provider_drops_itself_and_its_models():
|
||||
kept = sanitize({**SAVED, "llm_provider": "no-such-provider"}, "stock")
|
||||
assert "llm_provider" not in kept
|
||||
assert "quick_think_llm" not in kept and "deep_think_llm" not in kept
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_analysts_are_narrowed_to_the_asset_type():
|
||||
kept = sanitize(SAVED, "crypto")
|
||||
assert AnalystType.FUNDAMENTALS.value not in kept["analysts"]
|
||||
assert AnalystType.MARKET.value in kept["analysts"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_junk_values_are_dropped_rather_than_offered():
|
||||
kept = sanitize({"research_depth": 99, "analysts": ["astrology"], "output_language": 5}, "stock")
|
||||
assert kept == {}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_region_specific_provider_survives():
|
||||
kept = sanitize({**SAVED, "llm_provider": "qwen-cn", "quick_think_llm": None}, "stock")
|
||||
assert kept["llm_provider"] == "qwen-cn"
|
||||
|
||||
|
||||
# --- wiring ------------------------------------------------------------------
|
||||
|
||||
def _answer_every_prompt(monkeypatch):
|
||||
"""Drive the real selection flow, answering each prompt with a fixed value."""
|
||||
import cli.main as m
|
||||
|
||||
monkeypatch.setattr(cli_selections, "fetch_announcements", lambda: [])
|
||||
monkeypatch.setattr(cli_selections, "display_announcements", lambda *a: None)
|
||||
monkeypatch.setattr(cli_selections, "get_ticker", lambda: "NVDA")
|
||||
monkeypatch.setattr(cli_selections, "get_analysis_date", lambda: "2026-09-01")
|
||||
monkeypatch.setattr(cli_selections, "ask_output_language", lambda default=None: "English")
|
||||
monkeypatch.setattr(cli_selections, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET])
|
||||
monkeypatch.setattr(cli_selections, "select_research_depth", lambda default=None: 3)
|
||||
monkeypatch.setattr(cli_selections, "select_llm_provider", lambda default=None: ("openai", None))
|
||||
monkeypatch.setattr(cli_selections, "select_shallow_thinking_agent", lambda p, default=None: "gpt-5.6-mini")
|
||||
monkeypatch.setattr(cli_selections, "select_deep_thinking_agent", lambda p, default=None: "gpt-5.6")
|
||||
monkeypatch.setattr(cli_selections, "ask_openai_reasoning_effort", lambda: "medium")
|
||||
return m
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_selections_are_remembered_after_a_run(monkeypatch):
|
||||
"""Drives the real flow: a stubbed selections dict would hide a key mismatch."""
|
||||
_answer_every_prompt(monkeypatch)
|
||||
|
||||
cli_selections.get_user_selections()
|
||||
|
||||
remembered = load_last_run()
|
||||
assert remembered["analysts"] == ["market"]
|
||||
assert remembered["quick_think_llm"] == "gpt-5.6-mini"
|
||||
assert remembered["deep_think_llm"] == "gpt-5.6"
|
||||
assert remembered["llm_provider"] == "openai"
|
||||
assert "ticker" not in remembered # changes every run; never remembered
|
||||
assert "analysis_date" not in remembered # a stale date must not be offered
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_custom_language_is_remembered_without_breaking_the_next_run():
|
||||
"""A free-text answer is not one of the menu's choices, and questionary
|
||||
rejects a default it cannot find, so offering it back would crash startup."""
|
||||
from cli.prompts import ask_output_language
|
||||
|
||||
save_last_run({"output_language": "Turkish"})
|
||||
with mock.patch("cli.prompts.questionary.select") as select:
|
||||
select.return_value.ask.return_value = "English"
|
||||
ask_output_language(load_last_run()["output_language"])
|
||||
assert select.call_args.kwargs["default"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_remembered_endpoint_is_offered_back(monkeypatch):
|
||||
"""Users of a local or custom endpoint retyped the URL every run: it was
|
||||
remembered and validated, then never read."""
|
||||
|
||||
save_last_run({"llm_provider": "openai_compatible", "backend_url": "http://localhost:1234/v1"})
|
||||
offered = {}
|
||||
monkeypatch.setattr(cli_selections, "select_llm_provider", lambda default=None: ("openai_compatible", None))
|
||||
monkeypatch.setattr(cli_selections, "prompt_openai_compatible_url",
|
||||
lambda default=None: offered.setdefault("default", default) or "http://x/v1")
|
||||
monkeypatch.setattr(cli_selections, "fetch_announcements", lambda: [])
|
||||
monkeypatch.setattr(cli_selections, "display_announcements", lambda *a: None)
|
||||
monkeypatch.setattr(cli_selections, "get_ticker", lambda: "NVDA")
|
||||
monkeypatch.setattr(cli_selections, "get_analysis_date", lambda: "2026-09-01")
|
||||
monkeypatch.setattr(cli_selections, "ask_output_language", lambda default=None: "English")
|
||||
monkeypatch.setattr(cli_selections, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET])
|
||||
monkeypatch.setattr(cli_selections, "select_research_depth", lambda default=None: 1)
|
||||
monkeypatch.setattr(cli_selections, "select_shallow_thinking_agent", lambda p, default=None: "local-model")
|
||||
monkeypatch.setattr(cli_selections, "select_deep_thinking_agent", lambda p, default=None: "local-model")
|
||||
|
||||
cli_selections.get_user_selections()
|
||||
|
||||
assert offered["default"] == "http://localhost:1234/v1"
|
||||
@@ -0,0 +1,75 @@
|
||||
"""CLI symbol validation/classification must agree with the data path.
|
||||
|
||||
Regressions for #980 (validation rejected GC=F), #981 (BTCUSD misclassified as
|
||||
stock), #982 (BTC-USDT accepted but unpriceable on Yahoo).
|
||||
"""
|
||||
import pytest
|
||||
|
||||
import cli.run as cli_run
|
||||
from cli.models import AssetType
|
||||
from cli.prompts import detect_asset_type, is_valid_ticker_input, normalize_ticker_symbol
|
||||
from tradingagents.dataflows.symbols import normalize_symbol
|
||||
|
||||
|
||||
# --- #982: stablecoin-quoted crypto normalizes to Yahoo's -USD pair ---
|
||||
@pytest.mark.parametrize("raw,expected", [
|
||||
("BTCUSD", "BTC-USD"),
|
||||
("BTCUSDT", "BTC-USD"),
|
||||
("BTC-USDT", "BTC-USD"),
|
||||
("BTC-USDC", "BTC-USD"),
|
||||
("ethusdt", "ETH-USD"),
|
||||
# non-crypto must be untouched
|
||||
("AAPL", "AAPL"),
|
||||
("GC=F", "GC=F"),
|
||||
("600519.SS", "600519.SS"),
|
||||
("EURUSD", "EURUSD=X"),
|
||||
])
|
||||
def test_normalize_symbol_crypto_and_passthrough(raw, expected):
|
||||
assert normalize_symbol(raw) == expected
|
||||
|
||||
|
||||
# --- #980: validation accepts Yahoo futures/forex symbols ---
|
||||
@pytest.mark.parametrize("value,ok", [
|
||||
("GC=F", True),
|
||||
("EURUSD=X", True),
|
||||
("AAPL", True),
|
||||
("0700.HK", True),
|
||||
("^GSPC", True),
|
||||
("", True), # empty -> defaults to SPY downstream
|
||||
("bad symbol!", False), # space + '!' rejected
|
||||
("A" * 40, False), # too long
|
||||
])
|
||||
def test_ticker_input_validation(value, ok):
|
||||
assert is_valid_ticker_input(value) is ok
|
||||
|
||||
|
||||
# --- #981/#982: asset-type classified on the canonical symbol ---
|
||||
@pytest.mark.parametrize("raw,expected", [
|
||||
("BTCUSD", AssetType.CRYPTO),
|
||||
("BTC-USDT", AssetType.CRYPTO),
|
||||
("BTC-USD", AssetType.CRYPTO),
|
||||
("ETHUSD", AssetType.CRYPTO),
|
||||
("AAPL", AssetType.STOCK),
|
||||
("GC=F", AssetType.STOCK),
|
||||
("600519.SS", AssetType.STOCK),
|
||||
])
|
||||
def test_detect_asset_type(raw, expected):
|
||||
assert detect_asset_type(raw) == expected
|
||||
|
||||
|
||||
def test_cli_normalize_delegates_to_data_layer():
|
||||
# CLI must produce the same canonical symbol the data path will price.
|
||||
for raw in ("XAUUSD", "BTCUSD", "btc-usdt", "AAPL"):
|
||||
assert normalize_ticker_symbol(raw) == normalize_symbol(raw)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_run_directory_cannot_escape_the_results_directory(tmp_path, monkeypatch):
|
||||
"""Every other path that interpolates a ticker validates it first; the CLI's
|
||||
own results tree did not, so a ticker of '..' wrote a level up."""
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
cli_run._run_directory({"results_dir": str(tmp_path)}, "..", "2026-09-01")
|
||||
|
||||
ok = cli_run._run_directory({"results_dir": str(tmp_path)}, "NVDA", "2026-09-01")
|
||||
assert str(ok).startswith(str(tmp_path))
|
||||
@@ -0,0 +1,56 @@
|
||||
import unittest
|
||||
|
||||
from cli.models import AnalystType, AssetType
|
||||
from cli.prompts import detect_asset_type, filter_analysts_for_asset_type
|
||||
from tradingagents.graph.propagation import Propagator
|
||||
|
||||
|
||||
class CryptoAssetModeTests(unittest.TestCase):
|
||||
def test_detects_crypto_pair_symbols(self):
|
||||
self.assertEqual(detect_asset_type("BTC-USD"), AssetType.CRYPTO)
|
||||
self.assertEqual(detect_asset_type("eth-usd"), AssetType.CRYPTO)
|
||||
|
||||
def test_defaults_non_crypto_symbols_to_stock(self):
|
||||
self.assertEqual(detect_asset_type("AAPL"), AssetType.STOCK)
|
||||
self.assertEqual(detect_asset_type("SPY"), AssetType.STOCK)
|
||||
|
||||
def test_filters_out_fundamentals_analyst_for_crypto(self):
|
||||
analysts = [
|
||||
AnalystType.MARKET,
|
||||
AnalystType.SOCIAL,
|
||||
AnalystType.NEWS,
|
||||
AnalystType.FUNDAMENTALS,
|
||||
]
|
||||
|
||||
self.assertEqual(
|
||||
filter_analysts_for_asset_type(analysts, AssetType.CRYPTO),
|
||||
[
|
||||
AnalystType.MARKET,
|
||||
AnalystType.SOCIAL,
|
||||
AnalystType.NEWS,
|
||||
],
|
||||
)
|
||||
|
||||
def test_keeps_all_analysts_for_stock(self):
|
||||
analysts = [
|
||||
AnalystType.MARKET,
|
||||
AnalystType.SOCIAL,
|
||||
AnalystType.NEWS,
|
||||
AnalystType.FUNDAMENTALS,
|
||||
]
|
||||
|
||||
self.assertEqual(
|
||||
filter_analysts_for_asset_type(analysts, AssetType.STOCK),
|
||||
analysts,
|
||||
)
|
||||
|
||||
def test_propagator_includes_asset_type_in_initial_state(self):
|
||||
state = Propagator().create_initial_state(
|
||||
"BTC-USD", "2026-04-18", asset_type=AssetType.CRYPTO.value
|
||||
)
|
||||
|
||||
self.assertEqual(state["asset_type"], AssetType.CRYPTO.value)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,203 @@
|
||||
"""Config isolation: get/set must not leak nested-dict references."""
|
||||
|
||||
import copy
|
||||
import unittest
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config
|
||||
from tradingagents.dataflows.config import get_config, set_config
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class DataflowsConfigIsolationTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
set_config(copy.deepcopy(default_config.DEFAULT_CONFIG))
|
||||
|
||||
def test_get_config_returns_deep_copy(self):
|
||||
cfg = get_config()
|
||||
cfg["data_vendors"]["core_stock_apis"] = "alpha_vantage"
|
||||
cfg["tool_vendors"]["get_stock_data"] = "alpha_vantage"
|
||||
|
||||
fresh = get_config()
|
||||
self.assertEqual(fresh["data_vendors"]["core_stock_apis"], "yfinance")
|
||||
self.assertNotIn("get_stock_data", fresh["tool_vendors"])
|
||||
|
||||
def test_set_config_does_not_alias_caller_nested_dicts(self):
|
||||
custom = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
custom["data_vendors"]["core_stock_apis"] = "alpha_vantage"
|
||||
custom["tool_vendors"]["get_stock_data"] = "alpha_vantage"
|
||||
|
||||
set_config(custom)
|
||||
|
||||
custom["data_vendors"]["core_stock_apis"] = "yfinance"
|
||||
custom["tool_vendors"]["get_stock_data"] = "yfinance"
|
||||
|
||||
fresh = get_config()
|
||||
self.assertEqual(fresh["data_vendors"]["core_stock_apis"], "alpha_vantage")
|
||||
self.assertEqual(fresh["tool_vendors"]["get_stock_data"], "alpha_vantage")
|
||||
|
||||
def test_partial_nested_update_preserves_existing_defaults(self):
|
||||
set_config(
|
||||
{
|
||||
"data_vendors": {
|
||||
"core_stock_apis": "alpha_vantage",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
fresh = get_config()
|
||||
self.assertEqual(fresh["data_vendors"]["core_stock_apis"], "alpha_vantage")
|
||||
self.assertEqual(fresh["data_vendors"]["technical_indicators"], "yfinance")
|
||||
self.assertEqual(fresh["data_vendors"]["fundamental_data"], "yfinance")
|
||||
self.assertEqual(fresh["data_vendors"]["news_data"], "yfinance")
|
||||
|
||||
def test_nested_dict_updates_merge_one_level_deep(self):
|
||||
set_config({"tool_vendors": {"get_stock_data": "alpha_vantage"}})
|
||||
set_config({"tool_vendors": {"get_news": "alpha_vantage"}})
|
||||
|
||||
fresh = get_config()
|
||||
self.assertEqual(fresh["tool_vendors"]["get_stock_data"], "alpha_vantage")
|
||||
self.assertEqual(fresh["tool_vendors"]["get_news"], "alpha_vantage")
|
||||
|
||||
|
||||
# --- the config of the run in progress (#1369) --------------------------------
|
||||
|
||||
def _graph(config):
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
g.config = config
|
||||
g._checkpointer_ctx = None
|
||||
return g
|
||||
|
||||
|
||||
def _vendors_seen_by_a_run(graph, ticker="AAPL"):
|
||||
from tradingagents.dataflows.router import get_vendor
|
||||
|
||||
seen = []
|
||||
|
||||
def _run(*a, **k):
|
||||
seen.append(get_vendor("fundamental_data", "get_balance_sheet"))
|
||||
return {}, "Hold"
|
||||
|
||||
graph._run_graph = _run
|
||||
graph.propagate(ticker, "2026-09-01")
|
||||
return seen
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_run_reads_its_own_graphs_vendors_not_the_last_graph_built():
|
||||
"""Building a graph sets the process-wide config, and set_config merges, so
|
||||
a second graph built with the defaults was served the first one's vendors."""
|
||||
first = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
first["tool_vendors"] = {"get_balance_sheet": "sec_edgar,yfinance"}
|
||||
set_config(first) # graph A is built
|
||||
second = _graph(copy.deepcopy(default_config.DEFAULT_CONFIG))
|
||||
|
||||
assert _vendors_seen_by_a_run(second) == ["yfinance"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_graph_built_earlier_still_runs_with_its_own_config():
|
||||
"""Scoping at construction would hand graph A graph B's config if B was built
|
||||
after A; the config must be bound when the run starts."""
|
||||
a_config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
a_config["tool_vendors"] = {"get_balance_sheet": "sec_edgar,yfinance"}
|
||||
a = _graph(a_config)
|
||||
set_config(copy.deepcopy(default_config.DEFAULT_CONFIG)) # graph B is built
|
||||
|
||||
assert _vendors_seen_by_a_run(a) == ["sec_edgar,yfinance"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_concurrent_runs_each_read_their_own_config():
|
||||
import threading
|
||||
|
||||
barrier = threading.Barrier(2)
|
||||
results = {}
|
||||
|
||||
def run(name, vendor):
|
||||
config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
config["tool_vendors"] = {"get_balance_sheet": vendor}
|
||||
graph = _graph(config)
|
||||
from tradingagents.dataflows.router import get_vendor
|
||||
|
||||
def _run(*a, **k):
|
||||
barrier.wait(timeout=5) # both runs are in flight
|
||||
results[name] = get_vendor("fundamental_data", "get_balance_sheet")
|
||||
return {}, "Hold"
|
||||
|
||||
graph._run_graph = _run
|
||||
graph.propagate("AAPL", "2026-09-01")
|
||||
|
||||
threads = [threading.Thread(target=run, args=("a", "alpha_vantage")),
|
||||
threading.Thread(target=run, args=("b", "sec_edgar,yfinance"))]
|
||||
[t.start() for t in threads]
|
||||
[t.join() for t in threads]
|
||||
|
||||
assert results == {"a": "alpha_vantage", "b": "sec_edgar,yfinance"}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_settling_reads_the_graphs_own_config(monkeypatch):
|
||||
from tradingagents.dataflows.router import get_vendor
|
||||
|
||||
config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
config["tool_vendors"] = {"get_stock_data": "alpha_vantage"}
|
||||
graph = _graph(config)
|
||||
graph.memory_log = graph.reflector = None # the settlement below is a stand-in
|
||||
seen = []
|
||||
from tradingagents.graph import settlement
|
||||
|
||||
monkeypatch.setattr(settlement, "settle_pending",
|
||||
lambda *a: seen.append(get_vendor("core_stock_apis", "get_stock_data")))
|
||||
|
||||
graph.settle_pending("AAPL")
|
||||
|
||||
assert seen == ["alpha_vantage"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_tools_inside_a_langgraph_run_see_the_run_config():
|
||||
"""The fix rests on LangGraph carrying the caller's context into tool calls."""
|
||||
from langchain_core.messages import AIMessage
|
||||
from langchain_core.tools import tool
|
||||
from langgraph.graph import END, START, MessagesState, StateGraph
|
||||
from langgraph.prebuilt import ToolNode
|
||||
|
||||
from tradingagents.dataflows.config import run_config
|
||||
from tradingagents.dataflows.router import get_vendor
|
||||
|
||||
@tool
|
||||
def probe() -> str:
|
||||
"""Report the vendor the run would use."""
|
||||
return get_vendor("fundamental_data", "get_balance_sheet")
|
||||
|
||||
def call(state):
|
||||
return {"messages": [AIMessage("", tool_calls=[{"name": "probe", "args": {}, "id": "1"}])]}
|
||||
|
||||
g = StateGraph(MessagesState)
|
||||
g.add_node("call", call)
|
||||
g.add_node("tools", ToolNode([probe]))
|
||||
g.add_edge(START, "call")
|
||||
g.add_edge("call", "tools")
|
||||
g.add_edge("tools", END)
|
||||
|
||||
config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
config["tool_vendors"] = {"get_balance_sheet": "sec_edgar,yfinance"}
|
||||
with run_config(config):
|
||||
out = g.compile().invoke({"messages": [("user", "go")]})
|
||||
|
||||
assert out["messages"][-1].content == "sec_edgar,yfinance"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_run_config_missing_a_newer_key_still_reads_the_default():
|
||||
"""A config saved before a key existed must not fail inside a run."""
|
||||
from tradingagents.dataflows.config import get_config, run_config
|
||||
|
||||
config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
del config["news_article_limit"]
|
||||
with run_config(config):
|
||||
assert get_config()["news_article_limit"] == default_config.DEFAULT_CONFIG["news_article_limit"]
|
||||
@@ -0,0 +1,61 @@
|
||||
"""yfinance treats ``end`` as exclusive; we must request one extra day so the
|
||||
requested end_date (and the current day) is actually included.
|
||||
|
||||
Regressions for #986 (current-day OHLCV excluded) and #987 (requested end_date
|
||||
row omitted).
|
||||
"""
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
import tradingagents.dataflows.vendors.yahoo.market as yfin
|
||||
from tradingagents.dataflows.config import set_config
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_get_yfin_requests_inclusive_end(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
class FakeTicker:
|
||||
def __init__(self, symbol):
|
||||
pass
|
||||
|
||||
def history(self, start, end):
|
||||
captured["start"] = start
|
||||
captured["end"] = end
|
||||
idx = pd.to_datetime(["2025-05-08", "2025-05-09"])
|
||||
return pd.DataFrame(
|
||||
{"Open": [1.0, 2.0], "High": [1.0, 2.0], "Low": [1.0, 2.0],
|
||||
"Close": [1.0, 2.0], "Volume": [1, 2]},
|
||||
index=idx,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(yfin.yf, "Ticker", FakeTicker)
|
||||
out = yfin.get_YFin_data_online("AAPL", "2025-05-01", "2025-05-09")
|
||||
|
||||
# end is requested one day past end_date so 2025-05-09 is included (#987).
|
||||
assert captured["end"] == "2025-05-10"
|
||||
# Header still reflects the requested range, not the internal +1 day.
|
||||
assert "to 2025-05-09" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_load_ohlcv_requests_inclusive_end(monkeypatch, tmp_path):
|
||||
set_config({"data_cache_dir": str(tmp_path)})
|
||||
captured = {}
|
||||
|
||||
def fake_download(symbol, start, end, **kwargs):
|
||||
captured["end"] = end
|
||||
idx = pd.to_datetime([pd.Timestamp.today().normalize()])
|
||||
return pd.DataFrame(
|
||||
{"Open": [100.0], "High": [100.0], "Low": [100.0],
|
||||
"Close": [100.0], "Volume": [1]},
|
||||
index=idx,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(ohlcv.yf, "download", fake_download)
|
||||
today = pd.Timestamp.today().strftime("%Y-%m-%d")
|
||||
ohlcv.load_ohlcv("AAPL", today)
|
||||
|
||||
expected_end = (pd.Timestamp.today() + pd.Timedelta(days=1)).strftime("%Y-%m-%d")
|
||||
assert captured["end"] == expected_end # tomorrow -> today's row included (#986)
|
||||
@@ -0,0 +1,111 @@
|
||||
"""The first speaker in each debate must not rebut a nonexistent argument (#1176).
|
||||
|
||||
Each debate round's opening speaker receives an empty opponent response; the
|
||||
prompt used to interpolate it into a "refute the opponent" instruction, so models
|
||||
fabricated the other side's position. All five debators (bull, bear, and the
|
||||
three risk analysts) now substitute an explicit opening marker when the opponent
|
||||
has not spoken, and pass a real argument through unchanged.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.agents.context import opponent_argument_or_opening
|
||||
from tradingagents.agents.researchers.bear_researcher import create_bear_researcher
|
||||
from tradingagents.agents.researchers.bull_researcher import create_bull_researcher
|
||||
from tradingagents.agents.risk_mgmt.aggressive_debator import create_aggressive_debator
|
||||
from tradingagents.agents.risk_mgmt.conservative_debator import create_conservative_debator
|
||||
from tradingagents.agents.risk_mgmt.neutral_debator import create_neutral_debator
|
||||
|
||||
_REPORTS = {
|
||||
"company_of_interest": "AAPL", "asset_type": "stock",
|
||||
"market_report": "m", "sentiment_report": "s",
|
||||
"news_report": "n", "fundamentals_report": "f",
|
||||
}
|
||||
|
||||
|
||||
def _capturing_llm(captured: dict):
|
||||
llm = MagicMock()
|
||||
llm.invoke.side_effect = lambda prompt: (
|
||||
captured.__setitem__("prompt", prompt) or MagicMock(content="argument")
|
||||
)
|
||||
return llm
|
||||
|
||||
|
||||
def _investment_state(current_response):
|
||||
return {
|
||||
**_REPORTS,
|
||||
"count": 0,
|
||||
"investment_debate_state": {
|
||||
"history": "", "bull_history": "", "bear_history": "",
|
||||
"current_response": current_response, "count": 0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _risk_state(**responses):
|
||||
base = {
|
||||
"current_aggressive_response": "", "current_conservative_response": "",
|
||||
"current_neutral_response": "", "history": "", "aggressive_history": "",
|
||||
"conservative_history": "", "neutral_history": "", "count": 0,
|
||||
}
|
||||
base.update(responses)
|
||||
return {**_REPORTS, "trader_investment_plan": "plan", "risk_debate_state": base}
|
||||
|
||||
|
||||
# --- shared helper ----------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_helper_marks_empty_and_passes_through():
|
||||
assert "has not spoken yet" in opponent_argument_or_opening("", "bear analyst")
|
||||
assert opponent_argument_or_opening(" real point ", "bear") == "real point"
|
||||
|
||||
|
||||
# --- researchers ------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize(
|
||||
"factory,opponent",
|
||||
[(create_bull_researcher, "bear"), (create_bear_researcher, "bull")],
|
||||
)
|
||||
def test_researcher_opening_has_no_phantom_opponent(factory, opponent):
|
||||
captured = {}
|
||||
factory(_capturing_llm(captured))(_investment_state(""))
|
||||
assert "has not spoken yet" in captured["prompt"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_researcher_passes_real_opponent_argument():
|
||||
captured = {}
|
||||
state = _investment_state("Bear Analyst: valuation is stretched")
|
||||
create_bull_researcher(_capturing_llm(captured))(state)
|
||||
assert "valuation is stretched" in captured["prompt"]
|
||||
assert "has not spoken yet" not in captured["prompt"]
|
||||
|
||||
|
||||
# --- risk debators ----------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize(
|
||||
"factory", [create_aggressive_debator, create_conservative_debator, create_neutral_debator]
|
||||
)
|
||||
def test_risk_opening_has_no_phantom_opponent(factory):
|
||||
captured = {}
|
||||
factory(_capturing_llm(captured))(_risk_state())
|
||||
# Both opponent slots were empty -> two opening markers, no fabricated args.
|
||||
assert captured["prompt"].count("has not spoken yet") == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_risk_passes_real_opponent_arguments():
|
||||
captured = {}
|
||||
state = _risk_state(
|
||||
current_conservative_response="Conservative Analyst: trim risk",
|
||||
current_neutral_response="Neutral Analyst: hold steady",
|
||||
)
|
||||
create_aggressive_debator(_capturing_llm(captured))(state)
|
||||
assert "trim risk" in captured["prompt"]
|
||||
assert "hold steady" in captured["prompt"]
|
||||
assert "has not spoken yet" not in captured["prompt"]
|
||||
@@ -0,0 +1,239 @@
|
||||
"""Tests for DeepSeekChatOpenAI thinking-mode behaviour.
|
||||
|
||||
Two pieces verified:
|
||||
|
||||
1. ``reasoning_content`` is captured on receive into the AIMessage's
|
||||
``additional_kwargs`` and re-attached on send so DeepSeek's API
|
||||
sees the same value across turns.
|
||||
2. ``with_structured_output`` consults the capability table and
|
||||
suppresses ``tool_choice`` for models that reject it (V4 + reasoner),
|
||||
matching DeepSeek's official tool-calling pattern at
|
||||
https://api-docs.deepseek.com/guides/tool_calls.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langchain_core.prompt_values import ChatPromptValue
|
||||
from pydantic import BaseModel
|
||||
|
||||
from tradingagents.llm_clients.openai_client import (
|
||||
DeepSeekChatOpenAI,
|
||||
NormalizedChatOpenAI,
|
||||
_input_to_messages,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _input_to_messages — the helper that handles list / ChatPromptValue / other
|
||||
# (Gemini bot review note: non-list inputs must also work)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestInputToMessages:
|
||||
def test_list_input_returned_as_is(self):
|
||||
msgs = [HumanMessage(content="hi")]
|
||||
assert _input_to_messages(msgs) is msgs
|
||||
|
||||
def test_chat_prompt_value_unwrapped(self):
|
||||
msgs = [HumanMessage(content="hi")]
|
||||
prompt_value = ChatPromptValue(messages=msgs)
|
||||
assert _input_to_messages(prompt_value) == msgs
|
||||
|
||||
def test_string_input_yields_empty_list(self):
|
||||
# A bare string isn't a message-bearing input; the caller's normal
|
||||
# langchain conversion happens upstream of _get_request_payload.
|
||||
assert _input_to_messages("hello") == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reasoning content propagation across turns
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestDeepSeekReasoningContent:
|
||||
def _client(self):
|
||||
os.environ.setdefault("DEEPSEEK_API_KEY", "placeholder")
|
||||
return DeepSeekChatOpenAI(
|
||||
model="deepseek-v4-flash",
|
||||
api_key="placeholder",
|
||||
base_url="https://api.deepseek.com",
|
||||
)
|
||||
|
||||
def test_capture_on_receive(self):
|
||||
"""When the response carries reasoning_content, it lands on the
|
||||
AIMessage's additional_kwargs so the next turn can echo it back."""
|
||||
client = self._client()
|
||||
result = client._create_chat_result(
|
||||
{
|
||||
"model": "deepseek-v4-flash",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Plan: buy NVDA.",
|
||||
"reasoning_content": "Step 1: trend is up. Step 2: ...",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
)
|
||||
ai = result.generations[0].message
|
||||
assert ai.additional_kwargs["reasoning_content"] == "Step 1: trend is up. Step 2: ..."
|
||||
|
||||
def test_propagate_on_send(self):
|
||||
"""When an outgoing AIMessage carries reasoning_content, the request
|
||||
payload echoes it on the corresponding message dict."""
|
||||
client = self._client()
|
||||
prior = AIMessage(
|
||||
content="Plan",
|
||||
additional_kwargs={"reasoning_content": "weighed bull case"},
|
||||
)
|
||||
new_user = HumanMessage(content="Refine.")
|
||||
payload = client._get_request_payload([prior, new_user])
|
||||
# Find the assistant message in the payload
|
||||
assistant_dicts = [m for m in payload["messages"] if m.get("role") == "assistant"]
|
||||
assert assistant_dicts, "assistant message missing from outgoing payload"
|
||||
assert assistant_dicts[0]["reasoning_content"] == "weighed bull case"
|
||||
|
||||
def test_propagate_through_chat_prompt_value(self):
|
||||
"""Gemini bot review note: non-list inputs (ChatPromptValue) must
|
||||
also propagate reasoning_content."""
|
||||
client = self._client()
|
||||
prior = AIMessage(
|
||||
content="Plan",
|
||||
additional_kwargs={"reasoning_content": "weighed bull case"},
|
||||
)
|
||||
prompt_value = ChatPromptValue(messages=[prior, HumanMessage(content="Refine.")])
|
||||
payload = client._get_request_payload(prompt_value)
|
||||
assistant_dicts = [m for m in payload["messages"] if m.get("role") == "assistant"]
|
||||
assert assistant_dicts[0]["reasoning_content"] == "weighed bull case"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Capability-driven structured output: tool_choice suppressed for V4 + reasoner
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _bound_kwargs(runnable):
|
||||
"""Extract bind() kwargs from a with_structured_output result."""
|
||||
first = runnable.steps[0] if hasattr(runnable, "steps") else runnable
|
||||
return getattr(first, "kwargs", {})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStructuredOutputCapabilityDispatch:
|
||||
"""DeepSeek V4 and reasoner reject the tool_choice parameter
|
||||
(official guide: api-docs.deepseek.com/guides/tool_calls passes
|
||||
tools=[...] without tool_choice). Verify the capability dispatch
|
||||
suppresses tool_choice for those models and sends it for chat."""
|
||||
|
||||
class _Sample(BaseModel):
|
||||
answer: str
|
||||
|
||||
def _client(self, model):
|
||||
return DeepSeekChatOpenAI(
|
||||
model=model, api_key="placeholder", base_url="https://api.deepseek.com",
|
||||
)
|
||||
|
||||
def test_chat_sends_tool_choice(self):
|
||||
bound = self._client("deepseek-chat").with_structured_output(self._Sample)
|
||||
assert _bound_kwargs(bound).get("tool_choice") is not None
|
||||
|
||||
def test_reasoner_suppresses_tool_choice(self):
|
||||
bound = self._client("deepseek-reasoner").with_structured_output(self._Sample)
|
||||
# tool_choice is either absent or explicitly None — both are valid
|
||||
# signals that langchain's bind_tools will skip the parameter.
|
||||
assert _bound_kwargs(bound).get("tool_choice") in (None, ...) or \
|
||||
"tool_choice" not in _bound_kwargs(bound)
|
||||
|
||||
def test_v4_flash_suppresses_tool_choice(self):
|
||||
bound = self._client("deepseek-v4-flash").with_structured_output(self._Sample)
|
||||
assert _bound_kwargs(bound).get("tool_choice") is None or \
|
||||
"tool_choice" not in _bound_kwargs(bound)
|
||||
|
||||
def test_v4_pro_suppresses_tool_choice(self):
|
||||
bound = self._client("deepseek-v4-pro").with_structured_output(self._Sample)
|
||||
assert _bound_kwargs(bound).get("tool_choice") is None or \
|
||||
"tool_choice" not in _bound_kwargs(bound)
|
||||
|
||||
def test_future_v_variant_via_regex(self):
|
||||
"""Forward-compat: unknown deepseek-v\\d-* IDs inherit V4 quirks."""
|
||||
bound = self._client("deepseek-v5-hypothetical").with_structured_output(self._Sample)
|
||||
assert _bound_kwargs(bound).get("tool_choice") is None or \
|
||||
"tool_choice" not in _bound_kwargs(bound)
|
||||
|
||||
def test_schema_is_still_bound_as_tool(self):
|
||||
"""tool_choice is suppressed, but the schema is still bound as a tool —
|
||||
exactly matching DeepSeek's official tool-calling examples."""
|
||||
bound = self._client("deepseek-reasoner").with_structured_output(self._Sample)
|
||||
kwargs = _bound_kwargs(bound)
|
||||
tools = kwargs.get("tools", [])
|
||||
assert any(
|
||||
t.get("function", {}).get("name") == "_Sample" for t in tools
|
||||
), f"schema not bound as a tool: {tools}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Live API: structured output round-trips against the real DeepSeek backend
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _has_real_deepseek_key():
|
||||
key = os.environ.get("DEEPSEEK_API_KEY", "")
|
||||
return bool(key) and key != "placeholder"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.skipif(
|
||||
not _has_real_deepseek_key(),
|
||||
reason="DEEPSEEK_API_KEY not set (or placeholder); skipping live API call",
|
||||
)
|
||||
class TestDeepSeekLiveStructuredOutput:
|
||||
"""End-to-end: a real DeepSeek V4-flash call returns a typed instance.
|
||||
|
||||
Verifies the no-tool_choice path doesn't trigger the 400 reported in
|
||||
issue #678 and that the structured-output binding still parses to a
|
||||
Pydantic instance.
|
||||
"""
|
||||
|
||||
class _Pick(BaseModel):
|
||||
action: str
|
||||
confidence: float
|
||||
|
||||
def test_v4_flash_returns_structured_output(self):
|
||||
client = DeepSeekChatOpenAI(
|
||||
model="deepseek-v4-flash",
|
||||
api_key=os.environ["DEEPSEEK_API_KEY"],
|
||||
base_url="https://api.deepseek.com",
|
||||
timeout=60,
|
||||
)
|
||||
bound = client.with_structured_output(self._Pick)
|
||||
result = bound.invoke(
|
||||
"Pick BUY or SELL or HOLD for a tech stock with strong earnings. "
|
||||
"Confidence is a float between 0 and 1."
|
||||
)
|
||||
assert isinstance(result, self._Pick)
|
||||
assert result.action in {"BUY", "SELL", "HOLD"}
|
||||
assert 0.0 <= result.confidence <= 1.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Base class isolation: NormalizedChatOpenAI does NOT have DeepSeek behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestBaseClassIsolation:
|
||||
def test_normalized_does_not_propagate_reasoning_content(self):
|
||||
"""The general-purpose NormalizedChatOpenAI must not carry
|
||||
DeepSeek-specific behaviour. Only the subclass does."""
|
||||
assert not hasattr(NormalizedChatOpenAI, "_get_request_payload") or (
|
||||
NormalizedChatOpenAI._get_request_payload
|
||||
is NormalizedChatOpenAI.__bases__[0]._get_request_payload
|
||||
)
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Tests for TRADINGAGENTS_* env-var overlay onto DEFAULT_CONFIG."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config_module
|
||||
|
||||
|
||||
def _reload_with_env(monkeypatch, **overrides):
|
||||
"""Set/clear env vars then reload default_config to re-evaluate DEFAULT_CONFIG."""
|
||||
for key in list(default_config_module._ENV_OVERRIDES):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
for key, val in overrides.items():
|
||||
monkeypatch.setenv(key, val)
|
||||
return importlib.reload(default_config_module)
|
||||
|
||||
|
||||
def test_no_env_uses_built_in_defaults(monkeypatch):
|
||||
dc = _reload_with_env(monkeypatch)
|
||||
assert dc.DEFAULT_CONFIG["llm_provider"] == "openai"
|
||||
assert dc.DEFAULT_CONFIG["deep_think_llm"] == "gpt-6-sol"
|
||||
assert dc.DEFAULT_CONFIG["quick_think_llm"] == "gpt-6-luna"
|
||||
assert dc.DEFAULT_CONFIG["backend_url"] is None
|
||||
assert dc.DEFAULT_CONFIG["max_debate_rounds"] == 1
|
||||
assert dc.DEFAULT_CONFIG["checkpoint_enabled"] is False
|
||||
|
||||
|
||||
def test_string_overrides(monkeypatch):
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_LLM_PROVIDER="google",
|
||||
TRADINGAGENTS_DEEP_THINK_LLM="gemini-3-pro-preview",
|
||||
TRADINGAGENTS_QUICK_THINK_LLM="gemini-3-flash-preview",
|
||||
TRADINGAGENTS_LLM_BACKEND_URL="https://example.invalid/v1",
|
||||
TRADINGAGENTS_OUTPUT_LANGUAGE="Chinese",
|
||||
)
|
||||
assert dc.DEFAULT_CONFIG["llm_provider"] == "google"
|
||||
assert dc.DEFAULT_CONFIG["deep_think_llm"] == "gemini-3-pro-preview"
|
||||
assert dc.DEFAULT_CONFIG["quick_think_llm"] == "gemini-3-flash-preview"
|
||||
assert dc.DEFAULT_CONFIG["backend_url"] == "https://example.invalid/v1"
|
||||
assert dc.DEFAULT_CONFIG["output_language"] == "Chinese"
|
||||
|
||||
|
||||
def test_int_coercion(monkeypatch):
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_MAX_DEBATE_ROUNDS="3",
|
||||
TRADINGAGENTS_MAX_RISK_ROUNDS="2",
|
||||
)
|
||||
assert dc.DEFAULT_CONFIG["max_debate_rounds"] == 3
|
||||
assert isinstance(dc.DEFAULT_CONFIG["max_debate_rounds"], int)
|
||||
assert dc.DEFAULT_CONFIG["max_risk_discuss_rounds"] == 2
|
||||
assert isinstance(dc.DEFAULT_CONFIG["max_risk_discuss_rounds"], int)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,expected",
|
||||
[
|
||||
("true", True), ("True", True), ("1", True), ("yes", True), ("on", True),
|
||||
("false", False), ("False", False), ("0", False), ("no", False), ("off", False),
|
||||
],
|
||||
)
|
||||
def test_bool_coercion(monkeypatch, raw, expected):
|
||||
dc = _reload_with_env(monkeypatch, TRADINGAGENTS_CHECKPOINT_ENABLED=raw)
|
||||
assert dc.DEFAULT_CONFIG["checkpoint_enabled"] is expected
|
||||
|
||||
|
||||
def test_reasoning_thinking_overrides(monkeypatch):
|
||||
"""The provider reasoning/thinking knobs are env-configurable (non-interactive runs)."""
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_OPENAI_REASONING_EFFORT="high",
|
||||
TRADINGAGENTS_GOOGLE_THINKING_LEVEL="minimal",
|
||||
TRADINGAGENTS_ANTHROPIC_EFFORT="low",
|
||||
)
|
||||
assert dc.DEFAULT_CONFIG["openai_reasoning_effort"] == "high"
|
||||
assert dc.DEFAULT_CONFIG["google_thinking_level"] == "minimal"
|
||||
assert dc.DEFAULT_CONFIG["anthropic_effort"] == "low"
|
||||
|
||||
|
||||
def test_reasoning_effort_defaults_to_none(monkeypatch):
|
||||
"""Unset reasoning/thinking knobs stay None so each provider uses its own default."""
|
||||
dc = _reload_with_env(monkeypatch)
|
||||
assert dc.DEFAULT_CONFIG["openai_reasoning_effort"] is None
|
||||
assert dc.DEFAULT_CONFIG["google_thinking_level"] is None
|
||||
assert dc.DEFAULT_CONFIG["anthropic_effort"] is None
|
||||
|
||||
|
||||
def test_empty_env_value_is_passthrough(monkeypatch):
|
||||
"""Empty TRADINGAGENTS_* values must not clobber the built-in default."""
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_LLM_PROVIDER="",
|
||||
TRADINGAGENTS_MAX_DEBATE_ROUNDS="",
|
||||
)
|
||||
assert dc.DEFAULT_CONFIG["llm_provider"] == "openai"
|
||||
assert dc.DEFAULT_CONFIG["max_debate_rounds"] == 1
|
||||
|
||||
|
||||
def test_empty_path_value_keeps_the_default_path(monkeypatch):
|
||||
""".env.example lists the path variables blank; uncommenting one made the
|
||||
path empty, and the graph failed creating its directories."""
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_RESULTS_DIR="",
|
||||
TRADINGAGENTS_CACHE_DIR="",
|
||||
TRADINGAGENTS_MEMORY_LOG_PATH="",
|
||||
)
|
||||
home = dc._TRADINGAGENTS_HOME
|
||||
assert dc.DEFAULT_CONFIG["results_dir"] == os.path.join(home, "logs")
|
||||
assert dc.DEFAULT_CONFIG["data_cache_dir"] == os.path.join(home, "cache")
|
||||
assert dc.DEFAULT_CONFIG["memory_log_path"] == os.path.join(home, "memory", "trading_memory.md")
|
||||
|
||||
|
||||
def test_invalid_int_raises(monkeypatch):
|
||||
"""Garbage int values should surface a ValueError at import, not silently misconfigure."""
|
||||
monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "not-a-number")
|
||||
with pytest.raises(ValueError, match="TRADINGAGENTS_MAX_DEBATE_ROUNDS"):
|
||||
importlib.reload(default_config_module)
|
||||
# Restore module state for subsequent tests in this process
|
||||
monkeypatch.delenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", raising=False)
|
||||
importlib.reload(default_config_module)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad", ["treu", "flase", "maybe", "2", "enabled"])
|
||||
def test_invalid_bool_raises(monkeypatch, bad):
|
||||
"""A misspelled boolean must fail loudly (like ints) instead of silently False."""
|
||||
monkeypatch.setenv("TRADINGAGENTS_CHECKPOINT_ENABLED", bad)
|
||||
with pytest.raises(ValueError, match="TRADINGAGENTS_CHECKPOINT_ENABLED"):
|
||||
importlib.reload(default_config_module)
|
||||
monkeypatch.delenv("TRADINGAGENTS_CHECKPOINT_ENABLED", raising=False)
|
||||
importlib.reload(default_config_module)
|
||||
|
||||
|
||||
def test_unknown_env_var_is_ignored(monkeypatch):
|
||||
"""Env vars outside _ENV_OVERRIDES must not bleed into DEFAULT_CONFIG."""
|
||||
dc = _reload_with_env(
|
||||
monkeypatch,
|
||||
TRADINGAGENTS_NONEXISTENT_KEY="oops",
|
||||
)
|
||||
assert "nonexistent_key" not in dc.DEFAULT_CONFIG
|
||||
@@ -0,0 +1,287 @@
|
||||
"""FRED macro vendor: alias resolution, configuration errors, output formatting,
|
||||
missing-value handling, lookahead-safe windowing, and router integration.
|
||||
|
||||
All API access is mocked, so these run without a network connection or a key.
|
||||
"""
|
||||
import copy
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
import tradingagents.dataflows.config as config_module
|
||||
import tradingagents.default_config as default_config
|
||||
from tradingagents.dataflows import router
|
||||
from tradingagents.dataflows.config import set_config
|
||||
from tradingagents.dataflows.vendors import fred
|
||||
|
||||
# A small, stable set of observations to format against.
|
||||
_META = {
|
||||
"seriess": [
|
||||
{
|
||||
"title": "Unemployment Rate",
|
||||
"units_short": "%",
|
||||
"frequency": "Monthly",
|
||||
"seasonal_adjustment_short": "SA",
|
||||
}
|
||||
]
|
||||
}
|
||||
_OBS = {
|
||||
"observations": [
|
||||
{"date": "2025-06-01", "value": "4.1"},
|
||||
{"date": "2025-07-01", "value": "4.3"},
|
||||
{"date": "2025-08-01", "value": "."}, # missing -> skipped
|
||||
{"date": "2025-09-01", "value": "4.4"},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _request_stub(meta=_META, obs=_OBS):
|
||||
"""Build a _request replacement that dispatches on the endpoint path."""
|
||||
def _impl(path, params):
|
||||
if path == "series":
|
||||
return meta
|
||||
if path == "series/observations":
|
||||
return obs
|
||||
raise AssertionError(f"unexpected FRED path: {path}")
|
||||
return _impl
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class FredResolutionTests(unittest.TestCase):
|
||||
def test_alias_maps_to_series_id(self):
|
||||
self.assertEqual(fred._resolve_series_id("cpi"), "CPIAUCSL")
|
||||
self.assertEqual(fred._resolve_series_id("unemployment"), "UNRATE")
|
||||
|
||||
def test_alias_is_case_and_separator_insensitive(self):
|
||||
self.assertEqual(fred._resolve_series_id("Fed Funds Rate"), "FEDFUNDS")
|
||||
self.assertEqual(fred._resolve_series_id("10y-treasury"), "DGS10")
|
||||
|
||||
def test_unknown_alias_is_treated_as_raw_series_id(self):
|
||||
# Power users can pass any FRED series ID; we uppercase by convention.
|
||||
self.assertEqual(fred._resolve_series_id("dgs30"), "DGS30")
|
||||
self.assertEqual(fred._resolve_series_id("MyCustomSeries"), "MYCUSTOMSERIES")
|
||||
|
||||
def test_descriptive_phrase_is_rejected(self):
|
||||
# An LLM phrase (spaces / too long) is not a series ID — reject up front
|
||||
# with guidance rather than 400ing the API.
|
||||
for bad in ("bank of japan rate", "the unemployment number", "X" * 31):
|
||||
with self.assertRaises(ValueError):
|
||||
fred._resolve_series_id(bad)
|
||||
|
||||
def test_get_macro_data_returns_guidance_on_bad_indicator(self):
|
||||
# Invalid indicator -> actionable message, not a crash (no API call).
|
||||
out = fred.get_macro_data("bank of japan rate", "2026-01-01")
|
||||
self.assertIn("FRED", out)
|
||||
self.assertIn("not a known macro alias", out)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class FredConfigTests(unittest.TestCase):
|
||||
def test_missing_key_raises_not_configured(self):
|
||||
with mock.patch.dict("os.environ", {}, clear=True), \
|
||||
self.assertRaises(fred.FredNotConfiguredError):
|
||||
fred.get_api_key()
|
||||
|
||||
def test_not_configured_is_a_value_error(self):
|
||||
# Routing relies on this subclassing for "vendor unavailable" handling.
|
||||
self.assertTrue(issubclass(fred.FredNotConfiguredError, ValueError))
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class FredFormattingTests(unittest.TestCase):
|
||||
def test_report_has_header_latest_change_and_table(self):
|
||||
with mock.patch.object(fred, "_request", side_effect=_request_stub()):
|
||||
out = fred.get_macro_data("unemployment", "2025-09-30", 365)
|
||||
self.assertIn("## FRED: Unemployment Rate (UNRATE)", out)
|
||||
self.assertIn("Units: %", out)
|
||||
self.assertIn("Frequency: Monthly (SA)", out)
|
||||
self.assertIn("**Latest:** 4.4 (2025-09-01)", out)
|
||||
# change over the window: 4.4 - 4.1 = +0.30
|
||||
self.assertIn("+0.30", out)
|
||||
self.assertIn("| 2025-06-01 | 4.1 |", out)
|
||||
|
||||
def test_missing_value_is_skipped(self):
|
||||
with mock.patch.object(fred, "_request", side_effect=_request_stub()):
|
||||
out = fred.get_macro_data("unemployment", "2025-09-30", 365)
|
||||
# the "." observation must not appear as a row
|
||||
self.assertNotIn("2025-08-01", out)
|
||||
|
||||
def test_empty_window_reports_no_observations(self):
|
||||
empty = {"observations": []}
|
||||
with mock.patch.object(fred, "_request", side_effect=_request_stub(obs=empty)):
|
||||
out = fred.get_macro_data("unemployment", "2025-09-30", 30)
|
||||
self.assertIn("No observations", out)
|
||||
|
||||
def test_unknown_series_returns_not_found_message(self):
|
||||
# A well-formed but unknown series ID returns guidance, not a crash, so
|
||||
# the run is not aborted over an optional macro lookup.
|
||||
no_series = {"seriess": []}
|
||||
with mock.patch.object(fred, "_request", side_effect=_request_stub(meta=no_series)):
|
||||
out = fred.get_macro_data("totally_unknown_xyz", "2025-09-30", 30)
|
||||
self.assertIn("not found", out)
|
||||
|
||||
def test_long_series_is_truncated_but_change_uses_full_range(self):
|
||||
# Build > MAX_ROWS observations deterministically.
|
||||
obs = {
|
||||
"observations": [
|
||||
{"date": f"2025-01-{(i % 28) + 1:02d}", "value": str(i)}
|
||||
for i in range(fred.MAX_ROWS + 10)
|
||||
]
|
||||
}
|
||||
with mock.patch.object(fred, "_request", side_effect=_request_stub(obs=obs)):
|
||||
out = fred.get_macro_data("unemployment", "2025-12-31", 365)
|
||||
self.assertIn(f"most recent {fred.MAX_ROWS}", out)
|
||||
# change-over-window must reference the true first (0) and last value
|
||||
self.assertIn("from 0 ", out)
|
||||
body_rows = [ln for ln in out.splitlines() if ln.startswith("| 2025")]
|
||||
self.assertEqual(len(body_rows), fred.MAX_ROWS)
|
||||
|
||||
def test_window_is_lookahead_safe(self):
|
||||
# observation_end must equal curr_date so a past date never pulls future data.
|
||||
captured = {}
|
||||
|
||||
def _capture(path, params):
|
||||
captured[path] = params
|
||||
return _META if path == "series" else _OBS
|
||||
|
||||
with mock.patch.object(fred, "_request", side_effect=_capture):
|
||||
fred.get_macro_data("unemployment", "2025-09-30", 90)
|
||||
obs_params = captured["series/observations"]
|
||||
self.assertEqual(obs_params["observation_end"], "2025-09-30")
|
||||
self.assertEqual(obs_params["observation_start"], "2025-07-02") # 90d back
|
||||
|
||||
def test_requests_pin_the_data_vintage(self):
|
||||
# #1275: both the metadata and observations requests must pin the vintage
|
||||
# to curr_date (clamped to FRED's today), or FRED serves the latest
|
||||
# revision and revision-prone series leak future information. A past
|
||||
# curr_date sits below FRED's today, so it pins through unchanged.
|
||||
captured = {}
|
||||
|
||||
def _capture(path, params):
|
||||
captured[path] = params
|
||||
return _META if path == "series" else _OBS
|
||||
|
||||
with mock.patch.object(fred, "_fred_today", return_value="2026-01-01"), \
|
||||
mock.patch.object(fred, "_request", side_effect=_capture):
|
||||
fred.get_macro_data("cpi", "2025-09-30", 90)
|
||||
|
||||
for path in ("series", "series/observations"):
|
||||
self.assertEqual(captured[path]["realtime_start"], "2025-09-30", path)
|
||||
self.assertEqual(captured[path]["realtime_end"], "2025-09-30", path)
|
||||
|
||||
def test_future_curr_date_clamps_vintage_to_fred_today(self):
|
||||
# #1275 regression: on a live run curr_date is the caller's LOCAL date,
|
||||
# which can be a day ahead of FRED's US-Central clock. Pinning the vintage
|
||||
# to that future date 400s, and the routing layer then drops macro data
|
||||
# silently. The pin must clamp to FRED's today; the observation window
|
||||
# (future bars can't exist yet) stays at curr_date.
|
||||
captured = {}
|
||||
|
||||
def _capture(path, params):
|
||||
captured[path] = params
|
||||
return _META if path == "series" else _OBS
|
||||
|
||||
with mock.patch.object(fred, "_fred_today", return_value="2026-08-31"), \
|
||||
mock.patch.object(fred, "_request", side_effect=_capture):
|
||||
fred.get_macro_data("cpi", "2026-09-01", 90) # local a day ahead of Chicago
|
||||
|
||||
for path in ("series", "series/observations"):
|
||||
self.assertEqual(captured[path]["realtime_start"], "2026-08-31", path)
|
||||
self.assertEqual(captured[path]["realtime_end"], "2026-08-31", path)
|
||||
# the observation window still tracks curr_date, not the clamped vintage
|
||||
self.assertEqual(captured["series/observations"]["observation_end"], "2026-09-01")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class FredRoutingTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
|
||||
def tearDown(self):
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
|
||||
def test_macro_category_routes_to_fred(self):
|
||||
self.assertEqual(
|
||||
router.get_category_for_method("get_macro_indicators"), "macro_data"
|
||||
)
|
||||
set_config({"data_vendors": {"macro_data": "fred"}})
|
||||
with mock.patch.dict(
|
||||
router.VENDOR_METHODS,
|
||||
{"get_macro_indicators": {"fred": lambda *a, **k: "MACRO_OK"}},
|
||||
clear=False,
|
||||
):
|
||||
out = router.route_to_vendor("get_macro_indicators", "cpi", "2026-06-01", 365)
|
||||
self.assertEqual(out, "MACRO_OK")
|
||||
|
||||
def test_not_configured_degrades_gracefully(self):
|
||||
# macro_data is optional: with only fred and no key, the router degrades
|
||||
# to a sentinel instead of aborting the run — a missing optional key must
|
||||
# not crash an analysis.
|
||||
set_config({"data_vendors": {"macro_data": "fred"}})
|
||||
|
||||
def _unconfigured(*a, **k):
|
||||
raise fred.FredNotConfiguredError("FRED_API_KEY not set")
|
||||
|
||||
with mock.patch.dict(
|
||||
router.VENDOR_METHODS,
|
||||
{"get_macro_indicators": {"fred": _unconfigured}},
|
||||
clear=False,
|
||||
):
|
||||
out = router.route_to_vendor("get_macro_indicators", "cpi", "2026-06-01", 365)
|
||||
self.assertIn("DATA_UNAVAILABLE", out)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
|
||||
_KEY = "abcdef0123456789abcdef0123456789"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestKeyKeptOutOfErrors:
|
||||
"""The key travels as a query parameter, and requests quotes the full URL in
|
||||
its error messages, so any log or traceback would carry it (#1324)."""
|
||||
|
||||
def _raises(self, side_effect):
|
||||
with mock.patch.dict("os.environ", {"FRED_API_KEY": _KEY}), \
|
||||
mock.patch("tradingagents.dataflows.net.requests.get", side_effect=side_effect), \
|
||||
pytest.raises(requests.RequestException) as caught:
|
||||
fred._request("series", {"series_id": "DGS10"})
|
||||
return caught.value
|
||||
|
||||
def test_http_error_message_carries_no_key(self):
|
||||
response = mock.Mock(status_code=502)
|
||||
response.raise_for_status.side_effect = requests.HTTPError(
|
||||
f"502 Server Error for url: https://api.stlouisfed.org/fred/series?api_key={_KEY}",
|
||||
response=response,
|
||||
)
|
||||
with mock.patch.dict("os.environ", {"FRED_API_KEY": _KEY}), \
|
||||
mock.patch("tradingagents.dataflows.net.requests.get", return_value=response), \
|
||||
pytest.raises(requests.HTTPError) as caught:
|
||||
fred._request("series", {"series_id": "DGS10"})
|
||||
exc = caught.value
|
||||
assert _KEY not in str(exc) and _KEY not in repr(exc)
|
||||
# The response and request carry the full URL, so they are not attached.
|
||||
assert exc.response is None and exc.request is None
|
||||
assert exc.__cause__ is None and exc.__context__ is None # no chain holds the key
|
||||
|
||||
def test_connection_error_before_any_response_carries_no_key(self):
|
||||
exc = self._raises(requests.ConnectionError(
|
||||
f"Max retries exceeded with url: /fred/series?series_id=DGS10&api_key={_KEY}"))
|
||||
assert isinstance(exc, requests.ConnectionError)
|
||||
assert _KEY not in str(exc) and exc.__context__ is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_error_without_the_key_in_its_message_still_drops_the_request():
|
||||
# Some timeout messages omit the URL, but the attached request still has it.
|
||||
import requests as rq
|
||||
req = rq.Request("GET", f"https://api.stlouisfed.org/fred/series?api_key={_KEY}").prepare()
|
||||
with mock.patch.dict("os.environ", {"FRED_API_KEY": _KEY}), \
|
||||
mock.patch("tradingagents.dataflows.net.requests.get", side_effect=rq.Timeout("Read timed out.", request=req)), \
|
||||
pytest.raises(rq.Timeout) as caught:
|
||||
fred._request("series", {"series_id": "DGS10"})
|
||||
assert caught.value.request is None
|
||||
@@ -0,0 +1,134 @@
|
||||
"""Historical fundamentals must not leak a live company profile (#1300).
|
||||
|
||||
Vendor "company overview" endpoints (yfinance ``Ticker.info``, Alpha Vantage
|
||||
``OVERVIEW``) serve only present-day values: market cap, valuation multiples,
|
||||
the 52-week range and TTM income all move with today's quote, and even name,
|
||||
sector and industry shift when a company renames or is reclassified. None of it
|
||||
carries a historical vintage, so emitting it into a run dated in the past puts
|
||||
post-decision information into the analyst's context, in the same family as the
|
||||
FRED (#1275), social (#1220) and memory (#1251) leaks.
|
||||
|
||||
Both vendors withhold on one shared rule (``date_window.withhold_live_profile``)
|
||||
so switching ``fundamental_data`` between them cannot reintroduce the leak. The
|
||||
statement tools stay point-in-time by filtering on ``curr_date``, and a live run
|
||||
is unchanged. All API access is mocked.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows import date_window
|
||||
from tradingagents.dataflows.vendors.alpha_vantage import fundamentals as av
|
||||
from tradingagents.dataflows.vendors.yahoo import (
|
||||
fundamentals as yahoo_fundamentals,
|
||||
market as yahoo_market,
|
||||
)
|
||||
|
||||
_TODAY = "2026-09-07"
|
||||
_PAST = "2024-05-10"
|
||||
|
||||
# A profile payload mixing stable-looking identity fields with market-dependent ones.
|
||||
_INFO = {
|
||||
"longName": "Apple Inc.",
|
||||
"sector": "Technology",
|
||||
"industry": "Consumer Electronics",
|
||||
"marketCap": 3_500_000_000_000,
|
||||
"trailingPE": 34.2,
|
||||
"fiftyTwoWeekHigh": 260.1,
|
||||
"totalRevenue": 391_000_000_000,
|
||||
}
|
||||
|
||||
# Values that must never reach a historical run.
|
||||
_LEAKY = ("3500000000000", "34.2", "260.1", "391000000000",
|
||||
"Apple Inc.", "Technology", "Consumer Electronics")
|
||||
|
||||
|
||||
def _yf(curr_date, info=_INFO, today=_TODAY):
|
||||
with mock.patch.object(date_window, "get_current_date", return_value=today), \
|
||||
mock.patch.object(yahoo_fundamentals, "yf_retry", lambda fn: info), \
|
||||
mock.patch.object(yahoo_market.yf, "Ticker"):
|
||||
return yahoo_fundamentals.get_fundamentals("AAPL", curr_date)
|
||||
|
||||
|
||||
def _av(curr_date, today=_TODAY):
|
||||
"""Alpha Vantage path; the API call is mocked so a leak would be visible."""
|
||||
with mock.patch.object(date_window, "get_current_date", return_value=today), \
|
||||
mock.patch.object(av, "_make_api_request",
|
||||
return_value="MarketCapitalization: 3500000000000") as req:
|
||||
return av.get_fundamentals("AAPL", curr_date), req
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestYFinanceHistoricalRun:
|
||||
def test_no_profile_value_survives(self):
|
||||
out = _yf(_PAST)
|
||||
for leaked in _LEAKY:
|
||||
assert leaked not in out, f"leaked live-profile value {leaked!r}"
|
||||
|
||||
def test_states_the_as_of_date_and_explains_itself(self):
|
||||
# The analyst must be told why the figures are absent, so it does not
|
||||
# read the gap as a real signal or fabricate around it.
|
||||
out = _yf(_PAST)
|
||||
assert f"Point-in-time as of: {_PAST}" in out
|
||||
assert "withheld" in out
|
||||
assert _PAST in out and _TODAY not in out
|
||||
|
||||
def test_no_wall_clock_retrieval_stamp(self):
|
||||
# The old header stamped datetime.now(), which is what surfaced the leak.
|
||||
assert "Data retrieved on:" not in _yf(_PAST)
|
||||
|
||||
def test_the_request_is_not_even_made(self):
|
||||
# The response would only be discarded; skipping it also avoids burning
|
||||
# vendor quota on a call whose result cannot be used.
|
||||
with mock.patch.object(date_window, "get_current_date", return_value=_TODAY), \
|
||||
mock.patch.object(yahoo_market.yf, "Ticker") as tk:
|
||||
yahoo_fundamentals.get_fundamentals("AAPL", _PAST)
|
||||
tk.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestAlphaVantageHistoricalRun:
|
||||
"""The same rule must hold for the other fundamentals vendor, or switching
|
||||
data_vendors["fundamental_data"] would silently reintroduce the leak."""
|
||||
|
||||
def test_overview_is_withheld(self):
|
||||
out, _ = _av(_PAST)
|
||||
assert "3500000000000" not in out
|
||||
assert "withheld" in out
|
||||
assert f"Point-in-time as of: {_PAST}" in out
|
||||
|
||||
def test_the_api_call_is_not_made(self):
|
||||
_, req = _av(_PAST)
|
||||
req.assert_not_called()
|
||||
|
||||
def test_live_run_still_calls_the_api(self):
|
||||
out, req = _av(_TODAY)
|
||||
req.assert_called_once()
|
||||
assert "3500000000000" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLiveRunUnchanged:
|
||||
def test_yfinance_current_date_returns_the_full_profile(self):
|
||||
out = _yf(_TODAY)
|
||||
for value in _LEAKY:
|
||||
assert value in out
|
||||
assert "withheld" not in out
|
||||
|
||||
def test_yfinance_absent_curr_date_returns_the_full_profile(self):
|
||||
out = _yf(None)
|
||||
assert "Market Cap: 3500000000000" in out
|
||||
assert "withheld" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestNoUsableFieldsStillRaises:
|
||||
def test_stub_payload_raises_no_market_data(self):
|
||||
# yfinance returns {"trailingPegRatio": None} for unknown symbols; on a
|
||||
# live run that must stay a hard "no data", not a bare header.
|
||||
from tradingagents.dataflows.errors import NoMarketDataError
|
||||
|
||||
with pytest.raises(NoMarketDataError):
|
||||
_yf(_TODAY, info={"trailingPegRatio": None})
|
||||
@@ -0,0 +1,31 @@
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.google_client import GoogleClient
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGoogleApiKeyStandardization(unittest.TestCase):
|
||||
"""Verify GoogleClient accepts unified api_key parameter."""
|
||||
|
||||
@patch("tradingagents.llm_clients.google_client.NormalizedChatGoogleGenerativeAI")
|
||||
def test_api_key_handling(self, mock_chat):
|
||||
test_cases = [
|
||||
("unified api_key is mapped", {"api_key": "test-key-123"}, "test-key-123"),
|
||||
("legacy google_api_key still works", {"google_api_key": "legacy-key-456"}, "legacy-key-456"),
|
||||
("unified api_key takes precedence", {"api_key": "unified", "google_api_key": "legacy"}, "unified"),
|
||||
]
|
||||
|
||||
for msg, kwargs, expected_key in test_cases:
|
||||
with self.subTest(msg=msg):
|
||||
mock_chat.reset_mock()
|
||||
client = GoogleClient("gemini-3.5-flash", **kwargs)
|
||||
client.get_llm()
|
||||
call_kwargs = mock_chat.call_args[1]
|
||||
self.assertEqual(call_kwargs.get("google_api_key"), expected_key)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,62 @@
|
||||
"""Gemini thinking_level forwarding (Gemini 3.x).
|
||||
|
||||
The catalog is Gemini 3.x only, which takes the string ``thinking_level``
|
||||
directly. Pro, Gemini 3.8+ and the -latest aliases reject "minimal" with a 400,
|
||||
so it is mapped to "low" there; numbered Flash models before 3.8 accept it.
|
||||
"""
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.google_client import GoogleClient
|
||||
|
||||
|
||||
def _captured_kwargs(model, **kwargs):
|
||||
captured = {}
|
||||
with mock.patch.object(
|
||||
__import__("tradingagents.llm_clients.google_client", fromlist=["x"]),
|
||||
"NormalizedChatGoogleGenerativeAI",
|
||||
lambda **kw: captured.setdefault("kw", kw),
|
||||
):
|
||||
GoogleClient(model, api_key="x", **kwargs).get_llm()
|
||||
return captured["kw"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("level", ["minimal", "low", "medium", "high"])
|
||||
def test_flash_passes_thinking_level_through(level):
|
||||
kw = _captured_kwargs("gemini-3.5-flash", thinking_level=level)
|
||||
assert kw["thinking_level"] == level
|
||||
assert "thinking_budget" not in kw # the 2.5-era param is gone
|
||||
|
||||
|
||||
def test_pro_remaps_minimal_to_low():
|
||||
kw = _captured_kwargs("gemini-3.1-pro-preview", thinking_level="minimal")
|
||||
assert kw["thinking_level"] == "low" # Pro doesn't accept "minimal"
|
||||
|
||||
|
||||
def test_flash_38_remaps_minimal_to_low():
|
||||
kw = _captured_kwargs("gemini-3.8-flash", thinking_level="minimal")
|
||||
assert kw["thinking_level"] == "low" # 3.8 Flash 400s on "minimal"
|
||||
|
||||
|
||||
def test_flash_38_keeps_supported_levels():
|
||||
kw = _captured_kwargs("gemini-3.8-flash", thinking_level="high")
|
||||
assert kw["thinking_level"] == "high"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("alias", ["gemini-flash-latest", "gemini-pro-latest"])
|
||||
def test_latest_alias_remaps_minimal_to_low(alias):
|
||||
# Aliases move between generations; gemini-flash-latest 400s on "minimal".
|
||||
assert _captured_kwargs(alias, thinking_level="minimal")["thinking_level"] == "low"
|
||||
|
||||
|
||||
def test_pro_keeps_high():
|
||||
kw = _captured_kwargs("gemini-3.1-pro-preview", thinking_level="high")
|
||||
assert kw["thinking_level"] == "high"
|
||||
|
||||
|
||||
def test_no_thinking_level_is_omitted():
|
||||
kw = _captured_kwargs("gemini-3.5-flash")
|
||||
assert "thinking_level" not in kw
|
||||
assert "thinking_budget" not in kw
|
||||
@@ -0,0 +1,169 @@
|
||||
"""The whole graph, end to end, with scripted models and no network.
|
||||
|
||||
Every tool-using analyst calls each of its tools once through the vendor router,
|
||||
the debates and managers run, and the decision is parsed and logged. This pins
|
||||
the wiring: a restructure that drops a node, a tool or an edge fails here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import AIMessage, ToolMessage
|
||||
from langchain_core.outputs import ChatGeneration, ChatResult
|
||||
from langchain_core.runnables import RunnableLambda
|
||||
from pydantic import Field
|
||||
|
||||
from tradingagents.agents import context, schemas
|
||||
from tradingagents.agents.analysts import sentiment_analyst
|
||||
from tradingagents.dataflows import router
|
||||
from tradingagents.dataflows.vendors.yahoo import market as yahoo_market, snapshot
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
from tradingagents.graph import trading_graph
|
||||
|
||||
TRADE_DATE = "2026-01-09"
|
||||
|
||||
# Enough for every free-text reader: the PM's labelled rating and the trader's
|
||||
# closing proposal line.
|
||||
TEXT = "Report.\n\n**Rating**: Overweight\n\nFINAL TRANSACTION PROPOSAL: **BUY**"
|
||||
|
||||
STRUCTURED = {
|
||||
schemas.ResearchPlan: schemas.ResearchPlan(
|
||||
recommendation=schemas.PortfolioRating.OVERWEIGHT, rationale="r", strategic_actions="a"),
|
||||
schemas.TraderProposal: schemas.TraderProposal(action=schemas.TraderAction.BUY, reasoning="r"),
|
||||
schemas.PortfolioDecision: schemas.PortfolioDecision(
|
||||
rating=schemas.PortfolioRating.OVERWEIGHT, executive_summary="s", investment_thesis="t"),
|
||||
schemas.SentimentReport: schemas.SentimentReport(
|
||||
overall_band=schemas.SentimentBand.NEUTRAL, overall_score=5.0, confidence="low", narrative="n"),
|
||||
}
|
||||
|
||||
ARGS = {"symbol": "NVDA", "ticker": "NVDA", "curr_date": TRADE_DATE, "start_date": "2026-01-02",
|
||||
"end_date": TRADE_DATE, "indicator": "rsi", "topic": "Fed rate cut", "freq": "quarterly"}
|
||||
|
||||
|
||||
class ScriptedModel(BaseChatModel):
|
||||
"""Calls every bound tool once, then answers with TEXT."""
|
||||
|
||||
structured: bool = False
|
||||
tools: tuple = ()
|
||||
calls: list = Field(default_factory=list) # shared across bound copies
|
||||
fail_at: int | None = None # raise on this call, once
|
||||
|
||||
@property
|
||||
def _llm_type(self) -> str:
|
||||
return "scripted"
|
||||
|
||||
def bind_tools(self, tools, **kwargs):
|
||||
return self.model_copy(update={"tools": tuple(tools)})
|
||||
|
||||
def with_structured_output(self, schema, **kwargs):
|
||||
if not self.structured:
|
||||
raise NotImplementedError
|
||||
return RunnableLambda(lambda _: self._count() or STRUCTURED[schema])
|
||||
|
||||
def _count(self) -> None:
|
||||
self.calls.append(1)
|
||||
if len(self.calls) == self.fail_at:
|
||||
raise RuntimeError("provider unavailable")
|
||||
|
||||
def _generate(self, messages, stop=None, run_manager=None, **kwargs) -> ChatResult:
|
||||
self._count()
|
||||
if self.tools and not isinstance(messages[-1], ToolMessage):
|
||||
calls = [{"name": t.name, "id": f"call_{i}",
|
||||
"args": {k: v for k, v in ARGS.items()
|
||||
if k in t.tool_call_schema.model_json_schema()["properties"]}}
|
||||
for i, t in enumerate(self.tools)]
|
||||
message = AIMessage(content="", tool_calls=calls)
|
||||
else:
|
||||
message = AIMessage(content=TEXT)
|
||||
return ChatResult(generations=[ChatGeneration(message=message)])
|
||||
|
||||
|
||||
class _Client:
|
||||
def __init__(self, model):
|
||||
self.model = model
|
||||
|
||||
def get_llm(self):
|
||||
return self.model
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def offline(monkeypatch, tmp_path):
|
||||
"""Every vendor answers offline; returns the set of router methods called."""
|
||||
called: set[str] = set()
|
||||
for method, vendors in router.VENDOR_METHODS.items():
|
||||
for vendor in vendors:
|
||||
monkeypatch.setitem(vendors, vendor,
|
||||
lambda *a, _m=method, **k: called.add(_m) or f"{_m} data")
|
||||
prices = pd.DataFrame({
|
||||
"Date": pd.bdate_range(end=TRADE_DATE, periods=60),
|
||||
"Open": 100.0, "High": 101.0, "Low": 99.0, "Close": 100.5, "Volume": 1_000_000,
|
||||
})
|
||||
monkeypatch.setattr(snapshot, "load_ohlcv",
|
||||
lambda *a, **k: called.add("ohlcv") or prices.copy())
|
||||
monkeypatch.setattr(sentiment_analyst, "fetch_stocktwits_messages", lambda *a, **k: "no posts")
|
||||
monkeypatch.setattr(sentiment_analyst, "fetch_reddit_posts", lambda *a, **k: "no posts")
|
||||
monkeypatch.setattr(yahoo_market.yf, "Ticker", lambda s: type("T", (), {"info": {"longName": "NVIDIA"}})())
|
||||
context.resolve_instrument_identity.cache_clear()
|
||||
return called
|
||||
|
||||
|
||||
def _graph(tmp_path, monkeypatch, model, **config):
|
||||
cfg = copy.deepcopy(DEFAULT_CONFIG)
|
||||
cfg.update(results_dir=str(tmp_path / "results"), data_cache_dir=str(tmp_path / "cache"),
|
||||
memory_log_path=str(tmp_path / "log.md"), **config)
|
||||
monkeypatch.setattr(trading_graph, "create_llm_client", lambda **k: _Client(model))
|
||||
return trading_graph.TradingAgentsGraph(config=cfg)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("structured", [False, True], ids=["free-text", "structured"])
|
||||
def test_a_full_run_reaches_a_logged_decision(tmp_path, monkeypatch, offline, structured):
|
||||
graph = _graph(tmp_path, monkeypatch, ScriptedModel(structured=structured))
|
||||
|
||||
state, signal = graph.propagate("NVDA", TRADE_DATE)
|
||||
|
||||
assert signal == "Overweight"
|
||||
for key in ("market_report", "sentiment_report", "news_report", "fundamentals_report",
|
||||
"investment_plan", "trader_investment_plan", "final_trade_decision"):
|
||||
assert state[key].strip(), key
|
||||
tool_methods = {"get_stock_data", "get_indicators", "get_news", "get_global_news",
|
||||
"get_macro_indicators", "get_prediction_markets", "get_fundamentals",
|
||||
"get_balance_sheet", "get_cashflow", "get_income_statement",
|
||||
"get_insider_transactions", "ohlcv"}
|
||||
assert offline == tool_methods
|
||||
assert [e["rating"] for e in graph.memory_log.load_entries()] == ["Overweight"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_interrupted_run_resumes_from_its_checkpoint(tmp_path, monkeypatch, offline):
|
||||
model = ScriptedModel(fail_at=12) # past the analysts, before the decision
|
||||
graph = _graph(tmp_path, monkeypatch, model, checkpoint_enabled=True)
|
||||
with pytest.raises(RuntimeError, match="provider unavailable"):
|
||||
graph.propagate("NVDA", TRADE_DATE)
|
||||
calls_before = len(model.calls)
|
||||
|
||||
_, signal = graph.propagate("NVDA", TRADE_DATE)
|
||||
|
||||
assert signal == "Overweight"
|
||||
resumed_calls = len(model.calls) - calls_before
|
||||
full_run = ScriptedModel()
|
||||
_graph(tmp_path / "fresh", monkeypatch, full_run).propagate("NVDA", TRADE_DATE)
|
||||
# The resumed run makes only the calls the interrupted one had not completed.
|
||||
assert resumed_calls == len(full_run.calls) - (model.fail_at - 1)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_graph_reused_across_runs_keeps_no_run_state(tmp_path, monkeypatch, offline):
|
||||
"""A backtest reuses one graph over its whole grid; holding every run's full
|
||||
state would grow without bound."""
|
||||
graph = _graph(tmp_path, monkeypatch, ScriptedModel())
|
||||
for trade_date in ("2026-01-08", TRADE_DATE):
|
||||
graph.propagate("NVDA", trade_date)
|
||||
|
||||
held = [v for v in vars(graph).values() if isinstance(v, dict) and TRADE_DATE in v]
|
||||
assert held == []
|
||||
assert len(list(tmp_path.glob("results/NVDA/TradingAgentsStrategy_logs/*.json"))) == 2
|
||||
@@ -0,0 +1,59 @@
|
||||
"""Every report-producing agent must apply the configured output language
|
||||
(#740/#801).
|
||||
|
||||
A non-English run should produce a fully localized report, not a mix of
|
||||
languages. The bug originally happened because several agents silently omitted
|
||||
the instruction (fixed in 6b384f7); this test codifies the invariant so a future
|
||||
refactor can't quietly drop it again.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.agents.context import get_language_instruction
|
||||
|
||||
_AGENTS_DIR = Path(__file__).resolve().parents[1] / "tradingagents" / "agents"
|
||||
|
||||
# Every node whose text reaches the saved report. If you add a report-producing
|
||||
# agent, add it here — and make it call get_language_instruction().
|
||||
REPORT_AGENTS = [
|
||||
"analysts/market_analyst.py",
|
||||
"analysts/news_analyst.py",
|
||||
"analysts/fundamentals_analyst.py",
|
||||
"analysts/sentiment_analyst.py",
|
||||
"researchers/bull_researcher.py",
|
||||
"researchers/bear_researcher.py",
|
||||
"managers/research_manager.py",
|
||||
"managers/portfolio_manager.py",
|
||||
"risk_mgmt/aggressive_debator.py",
|
||||
"risk_mgmt/conservative_debator.py",
|
||||
"risk_mgmt/neutral_debator.py",
|
||||
"trader/trader.py",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanguageInstruction:
|
||||
def test_english_adds_no_tokens(self, monkeypatch):
|
||||
from tradingagents.dataflows.config import set_config
|
||||
set_config({"output_language": "English"})
|
||||
assert get_language_instruction() == ""
|
||||
|
||||
def test_non_english_emits_directive(self):
|
||||
from tradingagents.dataflows.config import set_config
|
||||
set_config({"output_language": "中文"})
|
||||
out = get_language_instruction()
|
||||
assert "中文" in out
|
||||
assert "entire response" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("rel", REPORT_AGENTS)
|
||||
def test_report_agent_applies_language_instruction(rel):
|
||||
path = _AGENTS_DIR / rel
|
||||
assert path.exists(), f"missing agent module: {rel}"
|
||||
src = path.read_text(encoding="utf-8")
|
||||
assert "get_language_instruction()" in src, (
|
||||
f"{rel} does not apply get_language_instruction(); its output would "
|
||||
f"ignore the configured output_language (#740/#801)."
|
||||
)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Tests for deterministic instrument-identity resolution (#814) and the
|
||||
context-anchored message placeholder (#888)."""
|
||||
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
|
||||
|
||||
from tradingagents.agents.context import (
|
||||
build_instrument_context,
|
||||
create_msg_delete,
|
||||
get_instrument_context_from_state,
|
||||
resolve_instrument_identity,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class ResolveInstrumentIdentityTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
resolve_instrument_identity.cache_clear()
|
||||
|
||||
def test_resolves_company_metadata_from_yfinance(self):
|
||||
with patch("tradingagents.dataflows.vendors.yahoo.market.yf.Ticker") as mock:
|
||||
mock.return_value.info = {
|
||||
"longName": "TOTO LTD.",
|
||||
"shortName": "TOTO",
|
||||
"sector": "Industrials",
|
||||
"industry": "Building Products & Equipment",
|
||||
"exchange": "PNK",
|
||||
"quoteType": "EQUITY",
|
||||
}
|
||||
identity = resolve_instrument_identity("totdy")
|
||||
mock.assert_called_once_with("TOTDY")
|
||||
self.assertEqual(identity["company_name"], "TOTO LTD.")
|
||||
self.assertEqual(identity["sector"], "Industrials")
|
||||
self.assertEqual(identity["industry"], "Building Products & Equipment")
|
||||
self.assertEqual(identity["exchange"], "PNK")
|
||||
|
||||
def test_falls_back_to_short_name(self):
|
||||
with patch("tradingagents.dataflows.vendors.yahoo.market.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"shortName": "TOTO", "sector": "Industrials"}
|
||||
identity = resolve_instrument_identity("TOTDY")
|
||||
self.assertEqual(identity["company_name"], "TOTO")
|
||||
|
||||
def test_skips_placeholder_values(self):
|
||||
with patch("tradingagents.dataflows.vendors.yahoo.market.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"longName": " ", "sector": "None", "industry": "n/a"}
|
||||
identity = resolve_instrument_identity("TOTDY")
|
||||
self.assertEqual(identity, {})
|
||||
|
||||
def test_fails_open_on_exception(self):
|
||||
with patch(
|
||||
"tradingagents.dataflows.vendors.yahoo.market.yf.Ticker",
|
||||
side_effect=RuntimeError("rate limited"),
|
||||
):
|
||||
self.assertEqual(resolve_instrument_identity("TOTDY"), {})
|
||||
|
||||
def test_result_is_cached(self):
|
||||
with patch("tradingagents.dataflows.vendors.yahoo.market.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"longName": "TOTO LTD."}
|
||||
first = resolve_instrument_identity("TOTDY")
|
||||
second = resolve_instrument_identity("TOTDY")
|
||||
mock.assert_called_once() # second call served from cache
|
||||
self.assertEqual(first, second)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class BuildInstrumentContextTests(unittest.TestCase):
|
||||
def test_mentions_exact_symbol_without_identity(self):
|
||||
context = build_instrument_context("7203.T")
|
||||
self.assertIn("7203.T", context)
|
||||
self.assertIn("exchange suffix", context)
|
||||
self.assertNotIn("Resolved identity", context)
|
||||
|
||||
def test_injects_resolved_identity(self):
|
||||
context = build_instrument_context(
|
||||
"TOTDY", "stock",
|
||||
{
|
||||
"company_name": "TOTO LTD.",
|
||||
"sector": "Industrials",
|
||||
"industry": "Building Products & Equipment",
|
||||
"exchange": "PNK",
|
||||
},
|
||||
)
|
||||
self.assertIn("Company: TOTO LTD.", context)
|
||||
self.assertIn("Industrials / Building Products & Equipment", context)
|
||||
self.assertIn("Exchange: PNK", context)
|
||||
self.assertIn("Do not substitute a different company", context)
|
||||
|
||||
def test_crypto_uses_name_label_and_keeps_hint(self):
|
||||
context = build_instrument_context(
|
||||
"BTC-USD", "crypto", {"company_name": "Bitcoin USD"}
|
||||
)
|
||||
self.assertIn("Name: Bitcoin USD", context)
|
||||
self.assertIn("crypto asset rather than a company", context)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class GetInstrumentContextFromStateTests(unittest.TestCase):
|
||||
def test_prefers_precomputed_context(self):
|
||||
state = {"company_of_interest": "TOTDY", "instrument_context": "PRECOMPUTED"}
|
||||
self.assertEqual(get_instrument_context_from_state(state), "PRECOMPUTED")
|
||||
|
||||
def test_fallback_is_network_free_ticker_only(self):
|
||||
# No instrument_context and no yfinance call — must not hit the network.
|
||||
with patch("tradingagents.dataflows.vendors.yahoo.market.yf.Ticker") as mock:
|
||||
context = get_instrument_context_from_state(
|
||||
{"company_of_interest": "NVDA", "asset_type": "stock"}
|
||||
)
|
||||
mock.assert_not_called()
|
||||
self.assertIn("NVDA", context)
|
||||
|
||||
def test_fallback_respects_asset_type(self):
|
||||
context = get_instrument_context_from_state(
|
||||
{"company_of_interest": "BTC-USD", "asset_type": "crypto"}
|
||||
)
|
||||
self.assertIn("crypto asset", context)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class ContextAnchoredPlaceholderTests(unittest.TestCase):
|
||||
"""#888 — the message-clear placeholder must not be a bare 'Continue'."""
|
||||
|
||||
def _run(self, state_extra):
|
||||
state = {
|
||||
"messages": [
|
||||
HumanMessage(content="old", id="h1"),
|
||||
AIMessage(content="reply", id="a1"),
|
||||
],
|
||||
**state_extra,
|
||||
}
|
||||
return create_msg_delete()(state)
|
||||
|
||||
def test_placeholder_is_not_bare_continue(self):
|
||||
result = self._run(
|
||||
{"company_of_interest": "EC", "asset_type": "stock", "trade_date": "2026-05-28"}
|
||||
)
|
||||
placeholder = result["messages"][-1]
|
||||
self.assertIsInstance(placeholder, HumanMessage)
|
||||
self.assertNotEqual(placeholder.content.strip(), "Continue")
|
||||
|
||||
def test_placeholder_carries_resolved_identity(self):
|
||||
result = self._run(
|
||||
{
|
||||
"company_of_interest": "EC",
|
||||
"instrument_context": "The instrument to analyze is `EC`. Resolved identity: Company: Ecopetrol.",
|
||||
"trade_date": "2026-05-28",
|
||||
}
|
||||
)
|
||||
content = result["messages"][-1].content
|
||||
self.assertIn("Ecopetrol", content)
|
||||
self.assertIn("2026-05-28", content)
|
||||
|
||||
def test_old_messages_are_removed(self):
|
||||
result = self._run({"company_of_interest": "EC", "trade_date": "2026-05-28"})
|
||||
removals = [m for m in result["messages"] if isinstance(m, RemoveMessage)]
|
||||
humans = [m for m in result["messages"] if isinstance(m, HumanMessage)]
|
||||
self.assertEqual(len(removals), 2)
|
||||
self.assertEqual(len(humans), 1)
|
||||
|
||||
def test_safe_defaults_when_state_minimal(self):
|
||||
result = create_msg_delete()({"messages": [], "company_of_interest": "EC"})
|
||||
placeholder = result["messages"][-1]
|
||||
self.assertNotEqual(placeholder.content.strip(), "Continue")
|
||||
self.assertIn("EC", placeholder.content)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Only the data layer imports vendor libraries.
|
||||
|
||||
Vendor calls belong in dataflows, where failures are raised as VendorError
|
||||
subclasses; a call made elsewhere can report an outage as a fact about the market.
|
||||
"""
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
VENDOR_LIBRARIES = {"yfinance"}
|
||||
|
||||
|
||||
def _imports(path: Path) -> set[str]:
|
||||
names = set()
|
||||
for node in ast.walk(ast.parse(path.read_text(encoding="utf-8"))):
|
||||
if isinstance(node, ast.Import):
|
||||
names |= {a.name.split(".")[0] for a in node.names}
|
||||
elif isinstance(node, ast.ImportFrom) and node.module and not node.level:
|
||||
names.add(node.module.split(".")[0])
|
||||
return names
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_vendor_libraries_are_imported_only_by_the_data_layer():
|
||||
data_layer = ROOT / "tradingagents" / "dataflows"
|
||||
offenders = sorted(
|
||||
str(path.relative_to(ROOT))
|
||||
for package in ("tradingagents", "cli")
|
||||
for path in (ROOT / package).rglob("*.py")
|
||||
if data_layer not in path.parents and _imports(path) & VENDOR_LIBRARIES
|
||||
)
|
||||
assert offenders == []
|
||||
@@ -0,0 +1,96 @@
|
||||
"""Configurable LLM SDK retry budget (#1090/#1091).
|
||||
|
||||
A single transient 429 burst used to kill an otherwise-healthy multi-agent run
|
||||
because each provider SDK's max_retries (default 2) was not exposed. This adds an
|
||||
opt-in llm_max_retries knob forwarded to every provider chat client.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config_module
|
||||
from tradingagents.llm_clients.factory import _coerce_max_retries, build_llm_kwargs
|
||||
|
||||
# --- coercion / validation -------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("value,expected", [(0, 0), (2, 2), (10, 10), ("6", 6)])
|
||||
def test_coerce_accepts_non_negative_ints_and_numeric_strings(value, expected):
|
||||
assert _coerce_max_retries(value) == expected
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", [-1, "-3"])
|
||||
def test_coerce_rejects_negative(bad):
|
||||
with pytest.raises(ValueError, match=">= 0"):
|
||||
_coerce_max_retries(bad)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", [True, False])
|
||||
def test_coerce_rejects_booleans(bad):
|
||||
with pytest.raises(ValueError, match="boolean"):
|
||||
_coerce_max_retries(bad)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", ["abc", "1.5", None])
|
||||
def test_coerce_rejects_non_integers(bad):
|
||||
with pytest.raises(ValueError, match="integer"):
|
||||
_coerce_max_retries(bad)
|
||||
|
||||
|
||||
# --- forwarding into provider kwargs --------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_not_forwarded_when_unset():
|
||||
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 = 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 = 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):
|
||||
build_llm_kwargs({"llm_provider": "openai", "llm_max_retries": -1})
|
||||
|
||||
|
||||
# --- env overlay -----------------------------------------------------------
|
||||
|
||||
def _reload_with_env(monkeypatch, **overrides):
|
||||
for key in list(default_config_module._ENV_OVERRIDES):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
for key, val in overrides.items():
|
||||
monkeypatch.setenv(key, val)
|
||||
return importlib.reload(default_config_module)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_default_is_none(monkeypatch):
|
||||
dc = _reload_with_env(monkeypatch)
|
||||
assert dc.DEFAULT_CONFIG["llm_max_retries"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_env_override_sets_config(monkeypatch):
|
||||
dc = _reload_with_env(monkeypatch, TRADINGAGENTS_LLM_MAX_RETRIES="8")
|
||||
# None-default key: env value arrives as a string and is coerced downstream.
|
||||
assert dc.DEFAULT_CONFIG["llm_max_retries"] == "8"
|
||||
assert _coerce_max_retries(dc.DEFAULT_CONFIG["llm_max_retries"]) == 8
|
||||
@@ -0,0 +1,118 @@
|
||||
"""Configurable output-token cap (#1204).
|
||||
|
||||
Some model/gateway combinations (e.g. deepseek-v4-flash deployments) emit
|
||||
unbounded reasoning/output and hang or trip an idle timeout. An opt-in
|
||||
``max_tokens`` config knob is forwarded to every provider so a run can bound it;
|
||||
Gemini names the parameter ``max_output_tokens``, so it is forwarded under the
|
||||
right key per provider.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.default_config as default_config_module
|
||||
from tradingagents.llm_clients.factory import _coerce_max_tokens, build_llm_kwargs
|
||||
|
||||
# --- coercion / validation -------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("value,expected", [(1, 1), (8192, 8192), ("4096", 4096)])
|
||||
def test_coerce_accepts_positive_ints_and_numeric_strings(value, expected):
|
||||
assert _coerce_max_tokens(value) == expected
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", [0, -1, "0", "-5"])
|
||||
def test_coerce_rejects_non_positive(bad):
|
||||
with pytest.raises(ValueError, match="> 0"):
|
||||
_coerce_max_tokens(bad)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", [True, False])
|
||||
def test_coerce_rejects_booleans(bad):
|
||||
with pytest.raises(ValueError, match="boolean"):
|
||||
_coerce_max_tokens(bad)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("bad", ["abc", "1.5", None])
|
||||
def test_coerce_rejects_non_integers(bad):
|
||||
with pytest.raises(ValueError, match="integer"):
|
||||
_coerce_max_tokens(bad)
|
||||
|
||||
|
||||
# --- forwarding into provider kwargs (right key per provider) --------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_not_forwarded_when_unset():
|
||||
kwargs = build_llm_kwargs({"llm_provider": "openai", "max_tokens": None})
|
||||
assert "max_tokens" not in kwargs
|
||||
assert "max_output_tokens" not in kwargs
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider", ["openai", "anthropic", "deepseek", "openai_compatible"])
|
||||
def test_forwarded_as_max_tokens_for_non_google(provider):
|
||||
kwargs = build_llm_kwargs({"llm_provider": provider, "max_tokens": 8192})
|
||||
assert kwargs["max_tokens"] == 8192
|
||||
assert "max_output_tokens" not in kwargs
|
||||
|
||||
|
||||
@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 = 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 = 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):
|
||||
build_llm_kwargs({"llm_provider": "openai", "max_tokens": 0})
|
||||
|
||||
|
||||
# --- client-side allowlists carry the kwarg --------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_openai_and_google_clients_accept_the_kwarg():
|
||||
from tradingagents.llm_clients import openai_client
|
||||
from tradingagents.llm_clients.google_client import GoogleClient # noqa: F401
|
||||
assert "max_tokens" in openai_client._PASSTHROUGH_KWARGS
|
||||
# Google client forwards max_output_tokens through construction.
|
||||
llm = GoogleClient("gemini-3.5-flash", api_key="x", max_output_tokens=8192).get_llm()
|
||||
assert getattr(llm, "max_output_tokens", None) == 8192
|
||||
|
||||
|
||||
# --- env overlay -----------------------------------------------------------
|
||||
|
||||
def _reload_with_env(monkeypatch, **overrides):
|
||||
for key in list(default_config_module._ENV_OVERRIDES):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
for key, val in overrides.items():
|
||||
monkeypatch.setenv(key, val)
|
||||
return importlib.reload(default_config_module)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_default_is_none(monkeypatch):
|
||||
dc = _reload_with_env(monkeypatch)
|
||||
assert dc.DEFAULT_CONFIG["max_tokens"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_env_override_sets_config(monkeypatch):
|
||||
dc = _reload_with_env(monkeypatch, TRADINGAGENTS_MAX_TOKENS="8192")
|
||||
assert dc.DEFAULT_CONFIG["max_tokens"] == "8192"
|
||||
assert _coerce_max_tokens(dc.DEFAULT_CONFIG["max_tokens"]) == 8192
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,95 @@
|
||||
"""Memory-log lessons must be point-in-time safe in a backtest (#1251).
|
||||
|
||||
get_past_context previously returned every resolved lesson regardless of the run
|
||||
date, so a historical run could learn from an outcome that had not happened yet.
|
||||
Resolved entries now record the date their outcome became known (``resolved:``),
|
||||
and get_past_context(as_of=...) filters on it. Legacy entries without a
|
||||
resolution date are excluded from a point-in-time query (conservative migration).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
|
||||
|
||||
def _log(tmp_path):
|
||||
return TradingMemoryLog({"memory_log_path": str(tmp_path / "mem.md")})
|
||||
|
||||
|
||||
def _resolve(log, ticker, date, resolution_date, reflection):
|
||||
log.store_decision(ticker, date, f"Rating: Buy\n{reflection}")
|
||||
log.update_with_outcome(
|
||||
ticker, date, 0.05, 0.02, 5, reflection, resolution_date=resolution_date,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_resolution_date_is_stored_and_parsed(tmp_path):
|
||||
log = _log(tmp_path)
|
||||
_resolve(log, "NVDA", "2026-01-05", "2026-01-10", "outcome known 01-10")
|
||||
entry = log.load_entries()[0]
|
||||
assert entry["resolved"] == "2026-01-10"
|
||||
assert "resolved:2026-01-10" in (tmp_path / "mem.md").read_text()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_as_of_excludes_lessons_resolved_after_the_run_date(tmp_path):
|
||||
log = _log(tmp_path)
|
||||
# Decision on 01-05, outcome only known on 01-10.
|
||||
_resolve(log, "NVDA", "2026-01-05", "2026-01-10", "great trade")
|
||||
|
||||
# A run as-of 01-07 must NOT see it (the outcome was still in the future).
|
||||
assert log.get_past_context("NVDA", as_of="2026-01-07") == ""
|
||||
# A run as-of 01-10 (and later) sees it.
|
||||
assert "great trade" in log.get_past_context("NVDA", as_of="2026-01-10")
|
||||
assert "great trade" in log.get_past_context("NVDA", as_of="2026-02-01")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_no_as_of_is_unfiltered_live_behavior(tmp_path):
|
||||
log = _log(tmp_path)
|
||||
_resolve(log, "NVDA", "2026-01-05", "2026-01-10", "great trade")
|
||||
# Live run (no as_of): unchanged behavior, lesson is shown.
|
||||
assert "great trade" in log.get_past_context("NVDA")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_legacy_entry_without_resolution_date_excluded_in_backtest(tmp_path):
|
||||
log = _log(tmp_path)
|
||||
# Simulate a pre-migration resolved entry: no resolution_date recorded.
|
||||
log.store_decision("NVDA", "2026-01-05", "Rating: Buy\nlegacy lesson")
|
||||
log.update_with_outcome("NVDA", "2026-01-05", 0.05, 0.02, 5, "legacy lesson")
|
||||
entry = log.load_entries()[0]
|
||||
assert entry["resolved"] is None
|
||||
|
||||
# Conservative: excluded from a point-in-time query (can't prove it was known)...
|
||||
assert log.get_past_context("NVDA", as_of="2026-06-01") == ""
|
||||
# ...but still available on a live (unfiltered) run.
|
||||
assert "legacy lesson" in log.get_past_context("NVDA")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cross_ticker_lessons_are_also_gated(tmp_path):
|
||||
log = _log(tmp_path)
|
||||
_resolve(log, "AAPL", "2026-01-05", "2026-01-10", "cross lesson")
|
||||
# Querying a different ticker as-of before resolution: no cross lesson leaks.
|
||||
assert log.get_past_context("NVDA", as_of="2026-01-07") == ""
|
||||
assert "cross lesson" in log.get_past_context("NVDA", as_of="2026-01-10")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_memory_as_of_gates_historical_but_not_live():
|
||||
# The graph filters only for a past trade date; a current-date run passes
|
||||
# None so live behavior and legacy entries are unaffected (#1251).
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
past = "2024-01-01"
|
||||
today = datetime.now().strftime("%Y-%m-%d")
|
||||
future = (datetime.now() + timedelta(days=30)).strftime("%Y-%m-%d")
|
||||
assert g._memory_as_of(past) == past # backtest -> filter on the trade date
|
||||
assert g._memory_as_of(today) is None # live -> no filter
|
||||
assert g._memory_as_of(future) is None # future-dated run -> no filter
|
||||
@@ -0,0 +1,73 @@
|
||||
"""Tests for MinimaxChatOpenAI quirks.
|
||||
|
||||
Verifies the subclass injects ``reasoning_split=True`` into outgoing
|
||||
requests so M2.x reasoning models put their <think> block into
|
||||
``reasoning_details`` instead of polluting ``message.content``.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage
|
||||
from pydantic import BaseModel
|
||||
|
||||
from tradingagents.llm_clients.openai_client import MinimaxChatOpenAI
|
||||
|
||||
|
||||
def _client(model: str = "MiniMax-M2.7"):
|
||||
os.environ.setdefault("MINIMAX_API_KEY", "placeholder")
|
||||
return MinimaxChatOpenAI(
|
||||
model=model,
|
||||
api_key="placeholder",
|
||||
base_url="https://api.minimax.io/v1",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMinimaxReasoningSplit:
|
||||
def test_reasoning_split_sent_via_extra_body_not_top_level(self):
|
||||
# Must be in extra_body, not top-level: the openai SDK validates
|
||||
# top-level params and rejects unknown ones like reasoning_split (#826).
|
||||
payload = _client()._get_request_payload([HumanMessage(content="hi")])
|
||||
assert payload.get("extra_body", {}).get("reasoning_split") is True
|
||||
assert "reasoning_split" not in payload # never top-level
|
||||
|
||||
def test_non_reasoning_minimax_does_not_inject_reasoning_split(self):
|
||||
"""Coding Plan / MiniMax-Text-01 / any non-M2-prefixed model must NOT
|
||||
receive reasoning_split at all (top-level or extra_body) (#826)."""
|
||||
for model in ("minimax-text-01", "MiniMax-Coding-Plan"):
|
||||
payload = _client(model)._get_request_payload(
|
||||
[HumanMessage(content="hi")]
|
||||
)
|
||||
assert "reasoning_split" not in payload
|
||||
assert "reasoning_split" not in payload.get("extra_body", {})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMinimaxStructuredOutputDispatch:
|
||||
"""M2.x models route through the capability table — tool_choice is
|
||||
suppressed but the schema is still bound as a tool."""
|
||||
|
||||
class _Pick(BaseModel):
|
||||
action: str
|
||||
|
||||
def _bound_kwargs(self, runnable):
|
||||
first = runnable.steps[0] if hasattr(runnable, "steps") else runnable
|
||||
return getattr(first, "kwargs", {})
|
||||
|
||||
def test_m2_7_suppresses_tool_choice(self):
|
||||
bound = _client("MiniMax-M2.7").with_structured_output(self._Pick)
|
||||
kwargs = self._bound_kwargs(bound)
|
||||
assert kwargs.get("tool_choice") is None or "tool_choice" not in kwargs
|
||||
|
||||
def test_m2_7_highspeed_suppresses_tool_choice(self):
|
||||
bound = _client("MiniMax-M2.7-highspeed").with_structured_output(self._Pick)
|
||||
kwargs = self._bound_kwargs(bound)
|
||||
assert kwargs.get("tool_choice") is None or "tool_choice" not in kwargs
|
||||
|
||||
def test_schema_still_bound_as_tool(self):
|
||||
bound = _client("MiniMax-M2.7").with_structured_output(self._Pick)
|
||||
tools = self._bound_kwargs(bound).get("tools", [])
|
||||
assert any(
|
||||
t.get("function", {}).get("name") == "_Pick" for t in tools
|
||||
), f"schema not bound: {tools}"
|
||||
@@ -0,0 +1,100 @@
|
||||
import unittest
|
||||
import warnings
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.base_client import BaseLLMClient
|
||||
from tradingagents.llm_clients.model_catalog import get_known_models
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
|
||||
class DummyLLMClient(BaseLLMClient):
|
||||
def __init__(self, provider: str, model: str):
|
||||
self.provider = provider
|
||||
super().__init__(model)
|
||||
|
||||
def get_llm(self):
|
||||
self.warn_if_unknown_model()
|
||||
return object()
|
||||
|
||||
def validate_model(self) -> bool:
|
||||
return validate_model(self.provider, self.model)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class ModelValidationTests(unittest.TestCase):
|
||||
def test_cli_catalog_models_are_all_validator_approved(self):
|
||||
for provider, models in get_known_models().items():
|
||||
if provider in ("ollama", "openrouter"):
|
||||
continue
|
||||
|
||||
for model in models:
|
||||
with self.subTest(provider=provider, model=model):
|
||||
self.assertTrue(validate_model(provider, model))
|
||||
|
||||
def test_unknown_model_emits_warning_for_strict_provider(self):
|
||||
client = DummyLLMClient("openai", "not-a-real-openai-model")
|
||||
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
client.get_llm()
|
||||
|
||||
self.assertEqual(len(caught), 1)
|
||||
self.assertIn("not-a-real-openai-model", str(caught[0].message))
|
||||
self.assertIn("openai", str(caught[0].message))
|
||||
|
||||
def test_openrouter_and_ollama_accept_custom_models_without_warning(self):
|
||||
for provider in ("openrouter", "ollama"):
|
||||
client = DummyLLMClient(provider, "custom-model-name")
|
||||
|
||||
with self.subTest(provider=provider):
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
client.get_llm()
|
||||
|
||||
self.assertEqual(caught, [])
|
||||
|
||||
|
||||
def test_legacy_ids_stay_valid_without_being_offered():
|
||||
from tradingagents.llm_clients.model_catalog import LEGACY_MODELS, MODEL_OPTIONS
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
for provider, ids in LEGACY_MODELS.items():
|
||||
offered = {v for opts in MODEL_OPTIONS[provider].values() for _, v in opts}
|
||||
for model in ids:
|
||||
assert validate_model(provider, model), model
|
||||
assert model not in offered, f"{model} is legacy but still in the picker"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_explicit_alias_of_a_listed_model_is_known():
|
||||
"""gpt-5.6 is served under its own name and as gpt-5.6-sol; naming the
|
||||
explicit one should not warn that the model is unknown."""
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
assert validate_model("openai", "gpt-5.6-sol")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider", ["openai", "anthropic", "google", "xai"])
|
||||
@pytest.mark.parametrize("mode", ["quick", "deep"])
|
||||
def test_every_provider_lets_you_name_your_own_model(provider, mode):
|
||||
"""The docs tell users to name any model their provider serves; the picker
|
||||
has to offer that too, or a new model is unreachable until we ship a list."""
|
||||
from tradingagents.llm_clients.model_catalog import get_model_options
|
||||
|
||||
assert "custom" in [value for _, value in get_model_options(provider, mode)]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider, model", [
|
||||
("xai", "grok-4.20-0309-reasoning"),
|
||||
("deepseek", "deepseek-v4-flash"),
|
||||
("qwen", "qwen3.7-max"),
|
||||
])
|
||||
def test_a_retired_model_id_still_runs_without_a_warning(provider, model):
|
||||
"""A config written against an earlier release keeps working: the provider
|
||||
still serves these, they are just no longer offered in the picker."""
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
assert validate_model(provider, model)
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Guard the news analyst prompt against tool-signature drift (#1116).
|
||||
|
||||
The prompt used to advertise ``get_news(query, ...)`` while the tool takes a
|
||||
``ticker``, tricking the LLM into hallucinating free-text query calls.
|
||||
"""
|
||||
import inspect
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.agents.analysts.news_analyst as na
|
||||
from tradingagents.agents.tools import get_news
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_get_news_takes_ticker_not_query():
|
||||
arg_names = set(get_news.args.keys())
|
||||
assert "ticker" in arg_names
|
||||
assert "query" not in arg_names
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_news_prompt_matches_get_news_signature():
|
||||
src = inspect.getsource(na)
|
||||
assert "get_news(ticker, start_date, end_date)" in src
|
||||
assert "get_news(query" not in src
|
||||
@@ -0,0 +1,257 @@
|
||||
"""yfinance news must not leak future-dated (or undated, in a backtest) articles
|
||||
into a historical window.
|
||||
|
||||
Regressions for #992 (flat articles bypassed the date filter), #1007 (global
|
||||
news injected future articles), #993 (empty-after-filter returned a blank body),
|
||||
and #1126 (inclusive upper bound leaked the midnight-after article; host-local
|
||||
timestamp parsing made filtering machine-dependent).
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.dataflows.vendors.yahoo.news as ynews
|
||||
from tradingagents.dataflows.date_window import in_window
|
||||
|
||||
|
||||
def _epoch(date_str):
|
||||
"""Epoch seconds for UTC midnight of ``date_str`` (host-timezone independent)."""
|
||||
return int(datetime.strptime(date_str, "%Y-%m-%d").replace(tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_flat_article_publish_time_is_parsed():
|
||||
# #992: flat articles now carry a pub_date (was always None -> unfilterable).
|
||||
# #1126: parsed as UTC-aware, so the date can't shift with the host timezone.
|
||||
data = ynews._extract_article_data(
|
||||
{"title": "X", "publisher": "P", "link": "l", "providerPublishTime": _epoch("2025-05-09")}
|
||||
)
|
||||
assert data["pub_date"] is not None
|
||||
assert data["pub_date"].tzinfo is not None
|
||||
assert data["pub_date"] == datetime(2025, 5, 9, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_window_excludes_future_and_undated_in_backtest():
|
||||
start = datetime(2025, 5, 1)
|
||||
end = datetime(2025, 5, 9) # historical window (well in the past)
|
||||
inside = datetime(2025, 5, 5)
|
||||
future = datetime(2025, 6, 1)
|
||||
assert in_window(inside, start, end) is True
|
||||
assert in_window(future, start, end) is False # look-ahead blocked
|
||||
assert in_window(None, start, end) is False # undated -> excluded in backtest
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_window_keeps_undated_in_live_window():
|
||||
# Live window (reaches today): undated articles can't be "future", so keep them.
|
||||
now = datetime.now(timezone.utc)
|
||||
assert in_window(None, now, now) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_upper_bound_is_exclusive():
|
||||
# #1126: an article stamped exactly midnight AFTER end_date leaked in under
|
||||
# the old inclusive bound; the whole of end_date itself must still be kept.
|
||||
start = datetime(2025, 5, 1)
|
||||
end = datetime(2025, 5, 9)
|
||||
midnight_after = datetime(2025, 5, 10, 0, 0, 0, tzinfo=timezone.utc)
|
||||
last_moment = datetime(2025, 5, 9, 23, 59, 59, tzinfo=timezone.utc)
|
||||
assert in_window(midnight_after, start, end) is False
|
||||
assert in_window(last_moment, start, end) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_offset_aware_timestamp_is_converted_not_truncated():
|
||||
# #1126: 2025-05-10T01:00+05:00 is really 2025-05-09T20:00Z -> inside the
|
||||
# window. Stripping tzinfo (old behavior) misread it as 05-10 and dropped it.
|
||||
start = datetime(2025, 5, 1)
|
||||
end = datetime(2025, 5, 9)
|
||||
aware = datetime.fromisoformat("2025-05-10T01:00:00+05:00")
|
||||
assert in_window(aware, start, end) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_global_news_future_flat_article_excluded(monkeypatch):
|
||||
# #1007: a flat, future-dated global article must not appear in a historical run.
|
||||
future_article = {"title": "FUTURE EVENT", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-06-01")}
|
||||
past_article = {"title": "PAST EVENT", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-05-05")}
|
||||
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = [future_article, past_article]
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
out = ynews.get_global_news_yfinance("2025-05-09", look_back_days=7, limit=10)
|
||||
assert "PAST EVENT" in out
|
||||
assert "FUTURE EVENT" not in out # #1007
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_global_news_empty_after_filter_is_informative(monkeypatch):
|
||||
# #993: everything filtered out -> a clear message, not a blank-bodied report.
|
||||
only_future = {"title": "FUTURE", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-06-01")}
|
||||
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = [only_future]
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
out = ynews.get_global_news_yfinance("2025-05-09", look_back_days=7, limit=10)
|
||||
assert "###" not in out # no empty article body
|
||||
# Only a later article came back, so the feed does not reach this window.
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
def _ticker_with(articles, monkeypatch):
|
||||
class FakeTicker:
|
||||
def __init__(self, *a, **k):
|
||||
pass
|
||||
|
||||
def get_news(self, count=20):
|
||||
return articles
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Ticker", FakeTicker)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_ticker_news_window_before_feed_coverage_is_unavailable(monkeypatch):
|
||||
# Yahoo serves only recent articles: a historical window gets none of them,
|
||||
# which must read as "cannot answer", not "no news happened".
|
||||
recent = [{"title": "RECENT", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2026-09-10")}]
|
||||
_ticker_with(recent, monkeypatch)
|
||||
out = ynews.get_news_yfinance("AAPL", "2026-08-07", "2026-08-14")
|
||||
assert "RECENT" not in out
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
assert "2026-09-10" not in out # an article after the window
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_ticker_news_covered_but_empty_window_is_a_real_absence(monkeypatch):
|
||||
articles = [{"title": "RECENT", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2026-09-10")},
|
||||
{"title": "OLDER", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2026-07-01")}]
|
||||
_ticker_with(articles, monkeypatch)
|
||||
out = ynews.get_news_yfinance("AAPL", "2026-08-07", "2026-08-14")
|
||||
assert "No news found" in out
|
||||
assert "unavailable" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("dates, expect_gap", [
|
||||
([], True), # empty feed: covers at most now
|
||||
([None], True), # undated only: same
|
||||
([datetime(2026, 5, 20, tzinfo=timezone.utc)], True), # all after the window
|
||||
([datetime(2026, 5, 4, tzinfo=timezone.utc)], True), # starts mid-window: partial
|
||||
([datetime(2026, 5, 1, 18, tzinfo=timezone.utc)], False), # reaches the first day
|
||||
([datetime(2026, 5, 20, tzinfo=timezone.utc),
|
||||
datetime(2026, 4, 1, tzinfo=timezone.utc)], False), # coverage reaches back
|
||||
])
|
||||
def test_coverage_gap_boundaries(dates, expect_gap):
|
||||
from tradingagents.dataflows.date_window import coverage_gap
|
||||
|
||||
out = coverage_gap(dates, "2026-05-01", "2026-05-08", "Feed", "items")
|
||||
assert (out is not None) is expect_gap
|
||||
if expect_gap:
|
||||
assert "unavailable for 2026-05-01..2026-05-08" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_ticker_news_empty_feed_for_a_past_window_is_unavailable(monkeypatch):
|
||||
_ticker_with([], monkeypatch)
|
||||
out = ynews.get_news_yfinance("AAPL", "2026-08-07", "2026-08-14")
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_ticker_news_null_feed_is_handled(monkeypatch):
|
||||
# Yahoo can return None instead of a list; that is unavailability, not an error.
|
||||
_ticker_with(None, monkeypatch)
|
||||
out = ynews.get_news_yfinance("AAPL", "2026-08-07", "2026-08-14")
|
||||
assert "unavailable" in out and "Error" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_global_news_empty_feed_for_a_past_window_is_unavailable(monkeypatch):
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = []
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
out = ynews.get_global_news_yfinance("2025-05-09", look_back_days=7, limit=10)
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_global_news_does_not_infer_coverage_from_a_stale_search_hit(monkeypatch):
|
||||
# Global news merges fuzzy searches; one old hit before the window says
|
||||
# nothing about the days in between, so the window stays unavailable.
|
||||
stale = {"title": "STALE", "publisher": "P", "link": "l", "providerPublishTime": _epoch("2025-01-01")}
|
||||
fresh = {"title": "FRESH", "publisher": "P", "link": "l", "providerPublishTime": _epoch("2025-06-01")}
|
||||
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = [fresh, stale]
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
out = ynews.get_global_news_yfinance("2025-05-09", look_back_days=7, limit=10)
|
||||
assert "unavailable" in out and "No global news found" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_coverage_gap_future_window_is_unavailable():
|
||||
from datetime import timedelta
|
||||
|
||||
from tradingagents.dataflows.date_window import coverage_gap
|
||||
today = datetime.now(timezone.utc).date()
|
||||
out = coverage_gap([], str(today), str(today + timedelta(days=3)), "Feed", "items")
|
||||
assert out is not None and "past today" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_out_of_window_articles_do_not_consume_the_article_budget(monkeypatch):
|
||||
"""The limit counts articles the run may see, not candidates fetched (#1356).
|
||||
|
||||
Out-of-window items were counted first, so they filled the budget, stopped
|
||||
the remaining searches, and the in-window news was reported as absent.
|
||||
"""
|
||||
stale = [{"title": f"OLD {i}", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-01-01")} for i in range(2)]
|
||||
wanted = {"title": "IN WINDOW", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-05-08")}
|
||||
pages = [stale, [wanted]]
|
||||
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = pages.pop(0) if pages else []
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
monkeypatch.setattr(ynews, "get_config", lambda: {
|
||||
"global_news_lookback_days": 7, "global_news_article_limit": 2,
|
||||
"global_news_queries": ["markets", "economy"],
|
||||
})
|
||||
|
||||
out = ynews.get_global_news_yfinance("2025-05-09")
|
||||
|
||||
assert "IN WINDOW" in out
|
||||
assert "OLD 0" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_article_limit_still_caps_what_is_returned(monkeypatch):
|
||||
articles = [{"title": f"NEWS {i}", "publisher": "P", "link": "l",
|
||||
"providerPublishTime": _epoch("2025-05-08")} for i in range(5)]
|
||||
|
||||
class FakeSearch:
|
||||
def __init__(self, *a, **k):
|
||||
self.news = articles
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Search", FakeSearch)
|
||||
out = ynews.get_global_news_yfinance("2025-05-09", look_back_days=7, limit=3)
|
||||
|
||||
assert out.count("### ") == 3
|
||||
@@ -0,0 +1,105 @@
|
||||
"""Tests that empty vendor results never become fabricated data.
|
||||
|
||||
Covers two systematic fixes:
|
||||
- load_ohlcv must not cache an empty download (cache poisoning), and must
|
||||
raise NoMarketDataError instead of returning an empty frame.
|
||||
- route_to_vendor must convert NoMarketDataError into a single explicit
|
||||
"NO_DATA_AVAILABLE" sentinel after all vendors are exhausted.
|
||||
"""
|
||||
|
||||
import os
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows import router
|
||||
from tradingagents.dataflows.config import set_config
|
||||
from tradingagents.dataflows.errors import NoMarketDataError
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLoadOhlcvNoPoison(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self._tmp = os.path.join(os.path.dirname(__file__), "_tmp_cache")
|
||||
os.makedirs(self._tmp, exist_ok=True)
|
||||
set_config({"data_cache_dir": self._tmp})
|
||||
|
||||
def tearDown(self):
|
||||
for f in os.listdir(self._tmp):
|
||||
os.remove(os.path.join(self._tmp, f))
|
||||
os.rmdir(self._tmp)
|
||||
|
||||
def test_empty_download_raises_and_does_not_cache(self):
|
||||
empty = pd.DataFrame()
|
||||
# Yahoo answers, so an empty download means the symbol has no data.
|
||||
reachable = mock.patch.object(ohlcv, "vendor_reachable", return_value=True)
|
||||
reachable.start()
|
||||
self.addCleanup(reachable.stop)
|
||||
with mock.patch.object(ohlcv.yf, "download", return_value=empty), \
|
||||
self.assertRaises(NoMarketDataError):
|
||||
ohlcv.load_ohlcv("FAKE", "2026-01-01")
|
||||
# Nothing should have been written to the cache.
|
||||
self.assertEqual(os.listdir(self._tmp), [])
|
||||
|
||||
# A second call must re-attempt the fetch (no poisoned cache served).
|
||||
with mock.patch.object(ohlcv.yf, "download", return_value=empty) as dl2:
|
||||
with self.assertRaises(NoMarketDataError):
|
||||
ohlcv.load_ohlcv("FAKE", "2026-01-01")
|
||||
self.assertTrue(dl2.called)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRouteToVendorSentinel(unittest.TestCase):
|
||||
def test_no_data_from_all_vendors_returns_sentinel(self):
|
||||
def raises_no_data(symbol, *a, **k):
|
||||
raise NoMarketDataError(symbol, "GC=F", "no rows")
|
||||
|
||||
patched = {"yfinance": raises_no_data, "alpha_vantage": raises_no_data}
|
||||
with mock.patch.dict(
|
||||
router.VENDOR_METHODS, {"get_stock_data": patched}, clear=False
|
||||
):
|
||||
result = router.route_to_vendor(
|
||||
"get_stock_data", "XAUUSD+", "2026-01-01", "2026-01-10"
|
||||
)
|
||||
self.assertIn("NO_DATA_AVAILABLE", result)
|
||||
self.assertIn("XAUUSD+", result)
|
||||
self.assertIn("GC=F", result)
|
||||
self.assertIn("Do not estimate", result)
|
||||
|
||||
def test_unconfigured_fallback_does_not_mask_no_data(self):
|
||||
# When the primary vendor reports no data and the fallback is simply
|
||||
# unavailable (e.g. missing API key -> raises), the no-data sentinel
|
||||
# must win rather than the fallback's incidental error crashing out.
|
||||
def raises_no_data(symbol, *a, **k):
|
||||
raise NoMarketDataError(symbol, symbol, "no rows")
|
||||
|
||||
def raises_unavailable(symbol, *a, **k):
|
||||
raise ValueError("ALPHA_VANTAGE_API_KEY environment variable is not set.")
|
||||
|
||||
patched = {"yfinance": raises_no_data, "alpha_vantage": raises_unavailable}
|
||||
with mock.patch.dict(
|
||||
router.VENDOR_METHODS, {"get_stock_data": patched}, clear=False
|
||||
):
|
||||
result = router.route_to_vendor(
|
||||
"get_stock_data", "FAKE", "2026-01-01", "2026-01-10"
|
||||
)
|
||||
self.assertIn("NO_DATA_AVAILABLE", result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_unreachable_yahoo_is_not_reported_as_a_symbol_without_insider_data():
|
||||
from tradingagents.dataflows.errors import VendorRateLimitError
|
||||
from tradingagents.dataflows.vendors.yahoo import fundamentals
|
||||
|
||||
ticker = type("T", (), {"insider_transactions": pd.DataFrame()})()
|
||||
with mock.patch.object(fundamentals.yf, "Ticker", return_value=ticker), \
|
||||
mock.patch.object(fundamentals, "vendor_reachable", return_value=False), \
|
||||
pytest.raises(VendorRateLimitError):
|
||||
fundamentals.get_insider_transactions("AAPL", curr_date="2026-09-21")
|
||||
@@ -0,0 +1,111 @@
|
||||
"""The OHLCV cache: one file per symbol, fresh only on the day it was written.
|
||||
|
||||
A current-day request also refetches past a TTL, so a run started before the
|
||||
day's bar was final is not served that snapshot all day (#1150). Keying the file
|
||||
by symbol rather than by day keeps the cache from growing a file per symbol per
|
||||
day (#1330).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv
|
||||
|
||||
NOW = pd.Timestamp("2026-07-18 12:00")
|
||||
STALE = ohlcv.OHLCV_CACHE_TTL_SECONDS + 60
|
||||
|
||||
|
||||
def _stamp(path, ts):
|
||||
"""Set ``path``'s mtime to the wall-clock ``ts``, read back in local time as
|
||||
the cache does. A naive ``pd.Timestamp.timestamp()`` would be taken as UTC."""
|
||||
t = ts.to_pydatetime().timestamp()
|
||||
os.utime(path, (t, t))
|
||||
|
||||
|
||||
def _write(tmp_path, name="AAPL-YFin-data.csv", age_seconds=0.0, last_date="2026-07-17"):
|
||||
f = tmp_path / name
|
||||
pd.DataFrame({"Date": [last_date], "Close": [100.0]}).to_csv(f, index=False)
|
||||
_stamp(f, NOW - pd.Timedelta(seconds=age_seconds))
|
||||
return f
|
||||
|
||||
|
||||
def _load(tmp_path, monkeypatch, curr_date, download):
|
||||
monkeypatch.setattr(ohlcv, "get_config", lambda: {"data_cache_dir": str(tmp_path)})
|
||||
monkeypatch.setattr(ohlcv.pd.Timestamp, "today", staticmethod(lambda: NOW))
|
||||
monkeypatch.setattr(ohlcv.yf, "download", download)
|
||||
return ohlcv.load_ohlcv("AAPL", curr_date)
|
||||
|
||||
|
||||
def _fail_download(*a, **k):
|
||||
raise AssertionError("fresh cache must not refetch")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_current_day_cache_past_ttl_is_not_fresh(tmp_path):
|
||||
# Today's bar missing or still in progress: row inspection can't tell, so the TTL governs.
|
||||
assert ohlcv._cache_is_fresh(_write(tmp_path, age_seconds=STALE), NOW.normalize(), NOW) is False
|
||||
f = _write(tmp_path, age_seconds=STALE, last_date="2026-07-18")
|
||||
assert ohlcv._cache_is_fresh(f, NOW.normalize(), NOW) is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_recent_cache_is_fresh(tmp_path):
|
||||
# Written moments ago: don't hammer the vendor (weekend/holiday guard).
|
||||
assert ohlcv._cache_is_fresh(_write(tmp_path), NOW.normalize(), NOW) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_historical_request_uses_todays_cache_past_the_ttl(tmp_path):
|
||||
f = _write(tmp_path, age_seconds=STALE, last_date="2026-04-30")
|
||||
assert ohlcv._cache_is_fresh(f, pd.Timestamp("2026-05-01"), NOW) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_download_from_an_earlier_day_is_not_fresh(tmp_path):
|
||||
f = _write(tmp_path, age_seconds=13 * 3600) # yesterday 23:00
|
||||
assert ohlcv._cache_is_fresh(f, pd.Timestamp("2026-05-01"), NOW) is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_load_ohlcv_refetches_stale_same_day_cache(tmp_path, monkeypatch):
|
||||
"""End-to-end: the freshness check is wired into load_ohlcv's cache branch."""
|
||||
_write(tmp_path, age_seconds=STALE)
|
||||
calls = []
|
||||
|
||||
def _fake_download(*a, **k):
|
||||
calls.append(1)
|
||||
return pd.DataFrame(
|
||||
{"Date": pd.to_datetime(["2026-07-17", "2026-07-18"]), "Close": [100.0, 222.0]}
|
||||
).set_index("Date")
|
||||
|
||||
out = _load(tmp_path, monkeypatch, "2026-07-18", _fake_download)
|
||||
assert calls, "stale same-day cache must trigger a refetch"
|
||||
assert 222.0 in out["Close"].values, "refreshed close must reach the caller"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_load_ohlcv_reuses_fresh_same_day_cache(tmp_path, monkeypatch):
|
||||
_write(tmp_path, last_date="2026-07-18")
|
||||
_load(tmp_path, monkeypatch, "2026-07-18", _fail_download)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_one_cache_file_per_symbol_across_days(tmp_path, monkeypatch):
|
||||
"""A later day's download replaces the symbol's file instead of adding one (#1330)."""
|
||||
monkeypatch.setattr(ohlcv, "get_config", lambda: {"data_cache_dir": str(tmp_path)})
|
||||
frame = pd.DataFrame({"Date": pd.to_datetime(["2026-07-16", "2026-07-17"]), "Close": [1.0, 2.0]})
|
||||
downloads = []
|
||||
monkeypatch.setattr(ohlcv.yf, "download", lambda *a, **k: downloads.append(1) or frame.set_index("Date"))
|
||||
|
||||
for day in ("2026-07-18 10:00", "2026-07-19 10:00", "2026-07-20 10:00"):
|
||||
now = pd.Timestamp(day)
|
||||
monkeypatch.setattr(ohlcv.pd.Timestamp, "today", staticmethod(lambda now=now: now))
|
||||
ohlcv.load_ohlcv("AAPL", "2026-07-17")
|
||||
written = list(tmp_path.glob("AAPL-*.csv"))
|
||||
_stamp(written[0], now)
|
||||
|
||||
assert len(downloads) == 3, "each new day refetches"
|
||||
assert [p.name for p in tmp_path.iterdir()] == ["AAPL-YFin-data.csv"]
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Tests for tolerating a non-`Date` index column in the Yahoo OHLCV loader (#890).
|
||||
|
||||
Guards against a download frame whose date column is `index` or `Datetime`
|
||||
instead of `Date`, which would otherwise silently drop every indicator.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import warnings
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv
|
||||
|
||||
|
||||
def _ohlcv(date_col: str) -> pd.DataFrame:
|
||||
"""OHLCV frame whose date column is named `date_col`."""
|
||||
dates = pd.bdate_range("2026-04-01", periods=10)
|
||||
return pd.DataFrame({
|
||||
date_col: dates,
|
||||
"Open": [100.0 + i for i in range(10)],
|
||||
"High": [101.0 + i for i in range(10)],
|
||||
"Low": [99.0 + i for i in range(10)],
|
||||
"Close": [100.5 + i for i in range(10)],
|
||||
"Volume": [1_000_000 + i for i in range(10)],
|
||||
})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEnsureDateColumn:
|
||||
def test_renames_index_column(self):
|
||||
out = ohlcv._ensure_date_column(_ohlcv("index"))
|
||||
assert "Date" in out.columns and "index" not in out.columns
|
||||
|
||||
def test_renames_datetime_and_date_variants(self):
|
||||
assert "Date" in ohlcv._ensure_date_column(_ohlcv("Datetime")).columns
|
||||
assert "Date" in ohlcv._ensure_date_column(_ohlcv("date")).columns
|
||||
|
||||
def test_leaves_existing_date_untouched(self):
|
||||
df = _ohlcv("Date")
|
||||
assert ohlcv._ensure_date_column(df) is df # no-op short-circuit
|
||||
|
||||
def test_no_datelike_column_is_left_alone(self):
|
||||
df = pd.DataFrame({"Close": [1, 2, 3]})
|
||||
out = ohlcv._ensure_date_column(df)
|
||||
assert "Date" not in out.columns # nothing to rename; caller handles
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCleanDataframeAcrossVersions:
|
||||
def test_clean_handles_index_column(self):
|
||||
"""A frame with `index` instead of `Date` must still clean to a
|
||||
usable, date-parsed frame (was KeyError: 'Date')."""
|
||||
cleaned = ohlcv._clean_dataframe(_ohlcv("index"))
|
||||
assert "Date" in cleaned.columns
|
||||
assert pd.api.types.is_datetime64_any_dtype(cleaned["Date"])
|
||||
assert len(cleaned) == 10
|
||||
|
||||
def test_clean_handles_legacy_date_column(self):
|
||||
cleaned = ohlcv._clean_dataframe(_ohlcv("Date"))
|
||||
assert len(cleaned) == 10
|
||||
|
||||
def test_indicators_compute_after_index_rename(self):
|
||||
"""stockstats must compute indicators on a frame whose date column
|
||||
arrived as `index`, instead of erroring per indicator."""
|
||||
from stockstats import wrap
|
||||
cleaned = ohlcv._clean_dataframe(_ohlcv("index"))
|
||||
df = wrap(cleaned)
|
||||
df["close_5_sma"] # triggers calculation
|
||||
assert "close_5_sma" in df.columns
|
||||
assert df["close_5_sma"].notna().any()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cleaning_a_frame_with_undated_rows_writes_to_its_own_copy():
|
||||
raw = pd.DataFrame({"Date": ["2026-01-08", None, "2026-01-09"], "Close": ["1", "2", "x"]})
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("error")
|
||||
cleaned = ohlcv._clean_dataframe(raw)
|
||||
assert cleaned["Close"].tolist()[0] == 1.0
|
||||
@@ -0,0 +1,208 @@
|
||||
"""The latest trading day's bar must not silently vanish (#1201).
|
||||
|
||||
yfinance can return the newest in-range bar with a NaN close (an unsettled or
|
||||
glitched session). The old path parsed dates without normalizing timezone and
|
||||
dropped every NaN-close row before applying the curr_date cutoff, so the latest
|
||||
bar disappeared and the previous trading day looked like the latest. Now dates
|
||||
are normalized before the cutoff, so the frame ends at the last settled bar
|
||||
instead of carrying a fabricated close.
|
||||
|
||||
Refusing the whole frame instead (the first attempt at #1201) reported a
|
||||
tradable symbol as invalid or delisted (#1289), so only a range with no close
|
||||
anywhere counts as no data and the staleness check judges the rest.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.errors import NoMarketDataError
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv
|
||||
|
||||
|
||||
def _stamp(path, ts):
|
||||
"""Set ``path``'s mtime to the wall-clock ``ts``, read back in local time as
|
||||
the cache does. A naive ``pd.Timestamp.timestamp()`` would be taken as UTC."""
|
||||
t = ts.to_pydatetime().timestamp()
|
||||
os.utime(path, (t, t))
|
||||
|
||||
# --- date normalization -----------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_normalize_dates_strips_tz_and_normalizes_to_midnight():
|
||||
aware = pd.Series(pd.to_datetime(
|
||||
["2026-05-08 09:30:00-04:00", "2026-05-09 16:00:00-04:00"]
|
||||
))
|
||||
out = ohlcv._normalize_dates(aware)
|
||||
assert out.dt.tz is None
|
||||
assert list(out) == [pd.Timestamp("2026-05-08"), pd.Timestamp("2026-05-09")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_normalize_dates_leaves_naive_dates_at_midnight():
|
||||
naive = pd.Series(pd.to_datetime(["2026-05-08 14:30:00", "2026-05-09 00:00:00"]))
|
||||
out = ohlcv._normalize_dates(naive)
|
||||
assert out.dt.tz is None
|
||||
assert list(out) == [pd.Timestamp("2026-05-08"), pd.Timestamp("2026-05-09")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_normalize_dates_handles_mixed_dst_offsets():
|
||||
# 5y of US bars span DST; via a cache CSV they arrive as mixed-offset
|
||||
# strings, which pd.to_datetime can't unify. Each keeps its own local date.
|
||||
mixed = pd.Series([
|
||||
"2026-01-08 00:00:00-05:00", # EST
|
||||
"2026-06-08 00:00:00-04:00", # EDT
|
||||
"not-a-date", # -> NaT
|
||||
])
|
||||
out = ohlcv._normalize_dates(mixed)
|
||||
assert out.iloc[0] == pd.Timestamp("2026-01-08")
|
||||
assert out.iloc[1] == pd.Timestamp("2026-06-08")
|
||||
assert pd.isna(out.iloc[2])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_normalize_dates_keeps_positive_offset_local_date():
|
||||
# A Tokyo bar at local midnight (+09:00) must stay on its own calendar day,
|
||||
# not shift to the previous UTC day (which utc=True parsing would cause).
|
||||
jst = pd.Series(["2026-05-08 00:00:00+09:00"])
|
||||
assert ohlcv._normalize_dates(jst).iloc[0] == pd.Timestamp("2026-05-08")
|
||||
|
||||
|
||||
# --- fill vs guard responsibilities ----------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_clean_dataframe_keeps_nan_close_for_the_caller_to_inspect():
|
||||
# _clean_dataframe normalizes but no longer drops the NaN close itself.
|
||||
df = pd.DataFrame({"Date": ["2026-05-08", "2026-05-09"], "Close": [100.0, float("nan")]})
|
||||
cleaned = ohlcv._clean_dataframe(df)
|
||||
assert len(cleaned) == 2
|
||||
assert pd.isna(cleaned["Close"].iloc[-1])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_fill_price_gaps_drops_nan_close_rows():
|
||||
df = pd.DataFrame({"Date": pd.to_datetime(["2026-05-07", "2026-05-08"]),
|
||||
"Close": [float("nan"), 100.0]})
|
||||
filled = ohlcv._fill_price_gaps(df)
|
||||
assert len(filled) == 1
|
||||
assert filled["Close"].iloc[0] == 100.0
|
||||
|
||||
|
||||
# --- load_ohlcv end-to-end (with a mocked cache read) -----------------------
|
||||
|
||||
def _run_load(monkeypatch, tmp_path, frame, curr_date):
|
||||
"""Drive load_ohlcv against a pre-seeded cache frame (no network)."""
|
||||
monkeypatch.setattr(ohlcv, "get_config", lambda: {"data_cache_dir": str(tmp_path)})
|
||||
today = pd.Timestamp(curr_date)
|
||||
monkeypatch.setattr(ohlcv.pd.Timestamp, "today", staticmethod(lambda: today))
|
||||
cache_file = tmp_path / "AAPL-YFin-data.csv"
|
||||
cache_file.write_text(frame.to_csv(index=False))
|
||||
_stamp(cache_file, today)
|
||||
|
||||
def _fail_download(*a, **k):
|
||||
raise AssertionError("should use the seeded cache, not download")
|
||||
monkeypatch.setattr(ohlcv.yf, "download", _fail_download)
|
||||
return ohlcv.load_ohlcv("AAPL", curr_date)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_unsettled_latest_bar_is_served_as_the_last_settled_bar(monkeypatch, tmp_path):
|
||||
# Newest bar (the curr_date) has no close: serve the last settled bar rather
|
||||
# than reporting the whole symbol as unavailable (#1289).
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-05-07", "2026-05-08"],
|
||||
"Open": [100.0, 101.0], "High": [101.0, 102.0], "Low": [99.0, 100.0],
|
||||
"Close": [100.5, float("nan")], "Volume": [1_000_000, 1_000_000],
|
||||
})
|
||||
out = _run_load(monkeypatch, tmp_path, frame, "2026-05-08")
|
||||
assert out["Date"].iloc[-1] == pd.Timestamp("2026-05-07")
|
||||
assert out["Close"].iloc[-1] == 100.5
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_no_settled_bar_at_all_is_still_no_data(monkeypatch, tmp_path):
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-05-07", "2026-05-08"],
|
||||
"Open": [100.0, 101.0], "High": [101.0, 102.0], "Low": [99.0, 100.0],
|
||||
"Close": [float("nan"), float("nan")], "Volume": [1_000_000, 1_000_000],
|
||||
})
|
||||
with pytest.raises(NoMarketDataError, match="no bar in range has a closing price"):
|
||||
_run_load(monkeypatch, tmp_path, frame, "2026-05-08")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_serving_the_last_settled_bar_does_not_bypass_the_staleness_check(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
# Falling back must not resurrect a long-dead series: once the closeless
|
||||
# tail is gone, the remaining bar is judged on its age like any other.
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-01-05", "2026-05-08"],
|
||||
"Open": [100.0, 101.0], "High": [101.0, 102.0], "Low": [99.0, 100.0],
|
||||
"Close": [100.5, float("nan")], "Volume": [1_000_000, 1_000_000],
|
||||
})
|
||||
with pytest.raises(NoMarketDataError, match="stale"):
|
||||
_run_load(monkeypatch, tmp_path, frame, "2026-05-08")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_older_nan_close_row_is_still_dropped(monkeypatch, tmp_path):
|
||||
# A stale gap mid-series is dropped; the valid latest bar is served.
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-05-06", "2026-05-07", "2026-05-08"],
|
||||
"Open": [100.0, 101.0, 102.0], "High": [101.0, 102.0, 103.0],
|
||||
"Low": [99.0, 100.0, 101.0],
|
||||
"Close": [100.5, float("nan"), 102.5], "Volume": [1_000_000, 1_000_000, 1_000_000],
|
||||
})
|
||||
out = _run_load(monkeypatch, tmp_path, frame, "2026-05-08")
|
||||
assert out["Close"].iloc[-1] == 102.5
|
||||
assert (out["Date"] == pd.Timestamp("2026-05-07")).sum() == 0 # the NaN row is gone
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_tz_aware_latest_bar_is_kept_at_the_cutoff(monkeypatch, tmp_path):
|
||||
# A tz-aware/intraday latest bar on the cutoff day must not be filtered out
|
||||
# by a naive-vs-aware comparison.
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-05-07 09:30:00-04:00", "2026-05-08 09:30:00-04:00"],
|
||||
"Open": [100.0, 101.0], "High": [101.0, 102.0], "Low": [99.0, 100.0],
|
||||
"Close": [100.5, 101.5], "Volume": [1_000_000, 1_000_000],
|
||||
})
|
||||
out = _run_load(monkeypatch, tmp_path, frame, "2026-05-08")
|
||||
assert out["Close"].iloc[-1] == 101.5
|
||||
assert out["Date"].iloc[-1] == pd.Timestamp("2026-05-08")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_snapshot_does_not_present_a_filled_price_as_reported(monkeypatch, tmp_path):
|
||||
"""Gap filling exists so indicators compute on a continuous series. The
|
||||
verification snapshot is the one place a number must be what the vendor
|
||||
reported, or the module built to stop invented prices supplies them."""
|
||||
from tradingagents.dataflows.vendors.yahoo import ohlcv, snapshot
|
||||
|
||||
frame = pd.DataFrame({
|
||||
"Date": ["2026-05-06", "2026-05-07", "2026-05-08"],
|
||||
"Open": [100.0, 104.5, ""], # the latest bar has not settled
|
||||
"High": [101.0, 105.5, ""],
|
||||
"Low": [99.0, 103.5, ""],
|
||||
"Close": [100.5, 105.0, 106.0],
|
||||
"Volume": [1000000, 1000000, ""],
|
||||
})
|
||||
today = pd.Timestamp("2026-05-08 12:00")
|
||||
monkeypatch.setattr(ohlcv, "get_config", lambda: {"data_cache_dir": str(tmp_path)})
|
||||
monkeypatch.setattr(ohlcv.pd.Timestamp, "today", staticmethod(lambda: today))
|
||||
cache = tmp_path / "AAPL-YFin-data.csv"
|
||||
cache.write_text(frame.to_csv(index=False))
|
||||
_stamp(cache, today)
|
||||
monkeypatch.setattr(ohlcv.yf, "download", lambda *a, **k: (_ for _ in ()).throw(
|
||||
AssertionError("should read the seeded cache")))
|
||||
|
||||
out = snapshot.build_verified_market_snapshot("AAPL", "2026-05-08", 3)
|
||||
|
||||
row = out.split("Latest verified OHLCV row")[1].split("###")[0]
|
||||
assert "104.50" not in row and "105.50" not in row # the previous session's numbers
|
||||
assert "106.00" in row # the close the vendor did report
|
||||
@@ -0,0 +1,225 @@
|
||||
"""Tests for OLLAMA_BASE_URL env-var override across CLI and client paths."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import re
|
||||
|
||||
import pytest
|
||||
|
||||
# Rich colorizes console output and highlights numbers and URLs, which splits
|
||||
# asserted substrings with escape codes ("port \x1b[1;33m11434"). Whether it
|
||||
# does so depends on the ambient terminal, so strip the codes to keep these
|
||||
# assertions independent of where the suite runs.
|
||||
_ANSI = re.compile(r"\x1b\[[0-9;]*m")
|
||||
|
||||
|
||||
def _console_out(capsys) -> str:
|
||||
return _ANSI.sub("", capsys.readouterr().out)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def _resync_reloaded_modules():
|
||||
"""Restore module state after this file's importlib.reload() calls.
|
||||
|
||||
Several tests below reload ``cli.prompts`` to re-evaluate OLLAMA_BASE_URL.
|
||||
That leaves the modules importing from it (cli.selections, then cli.run and
|
||||
cli.main) bound to the pre-reload functions, which breaks identity checks in
|
||||
unrelated tests that run afterward. Re-sync them in import order on teardown.
|
||||
"""
|
||||
yield
|
||||
import cli.main
|
||||
import cli.prompts
|
||||
import cli.run
|
||||
import cli.selections
|
||||
for module in (cli.prompts, cli.selections, cli.run, cli.main):
|
||||
importlib.reload(module)
|
||||
|
||||
|
||||
# ---- openai_client side: registry-driven base_url resolution --------------
|
||||
|
||||
|
||||
def _reload_client():
|
||||
import tradingagents.llm_clients.openai_client as mod
|
||||
return importlib.reload(mod)
|
||||
|
||||
|
||||
def _base_url(mod, provider, **kwargs):
|
||||
return str(mod.OpenAIClient(model="m", provider=provider, **kwargs).get_llm().openai_api_base)
|
||||
|
||||
|
||||
def test_resolver_returns_default_when_env_unset(monkeypatch):
|
||||
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
|
||||
mod = _reload_client()
|
||||
assert _base_url(mod, "ollama") == "http://localhost:11434/v1"
|
||||
|
||||
|
||||
def test_resolver_returns_env_when_set(monkeypatch):
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-ollama:11434/v1")
|
||||
mod = _reload_client()
|
||||
assert _base_url(mod, "ollama") == "http://remote-ollama:11434/v1"
|
||||
|
||||
|
||||
def test_resolver_evaluation_is_call_time(monkeypatch):
|
||||
"""Setting the env AFTER module import must still take effect."""
|
||||
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
|
||||
mod = _reload_client()
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://late-set:11434/v1")
|
||||
assert _base_url(mod, "ollama") == "http://late-set:11434/v1"
|
||||
|
||||
|
||||
def test_resolver_does_not_affect_other_providers(monkeypatch):
|
||||
"""OLLAMA_BASE_URL should NOT leak into xai/deepseek/etc."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://elsewhere/v1")
|
||||
mod = _reload_client()
|
||||
assert _base_url(mod, "xai") == "https://api.x.ai/v1"
|
||||
assert _base_url(mod, "deepseek") == "https://api.deepseek.com"
|
||||
|
||||
|
||||
def test_client_get_llm_picks_up_env(monkeypatch):
|
||||
"""End-to-end: OllamaClient.get_llm() respects OLLAMA_BASE_URL."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://my-ollama:11434/v1")
|
||||
mod = _reload_client()
|
||||
client = mod.OpenAIClient(model="llama3.1", provider="ollama")
|
||||
llm = client.get_llm()
|
||||
assert "my-ollama" in str(llm.openai_api_base)
|
||||
|
||||
|
||||
def test_explicit_base_url_overrides_env(monkeypatch):
|
||||
"""An explicit base_url passed to the client wins over the env var."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://env-set:11434/v1")
|
||||
mod = _reload_client()
|
||||
client = mod.OpenAIClient(
|
||||
model="llama3.1",
|
||||
provider="ollama",
|
||||
base_url="http://explicit:11434/v1",
|
||||
)
|
||||
llm = client.get_llm()
|
||||
assert "explicit" in str(llm.openai_api_base)
|
||||
assert "env-set" not in str(llm.openai_api_base)
|
||||
|
||||
|
||||
# ---- cli.prompts side: select_llm_provider dropdown -------------------------
|
||||
|
||||
|
||||
def test_cli_dropdown_uses_env(monkeypatch):
|
||||
"""The Ollama entry in the CLI dropdown must reflect OLLAMA_BASE_URL."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://cli-remote:11434/v1")
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
# Reach inside the function via the same env-read it does at call time
|
||||
ollama_url = (
|
||||
__import__("os").environ.get("OLLAMA_BASE_URL")
|
||||
or "http://localhost:11434/v1"
|
||||
)
|
||||
assert ollama_url == "http://cli-remote:11434/v1"
|
||||
|
||||
|
||||
def test_cli_dropdown_default_when_unset(monkeypatch):
|
||||
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
ollama_url = (
|
||||
__import__("os").environ.get("OLLAMA_BASE_URL")
|
||||
or "http://localhost:11434/v1"
|
||||
)
|
||||
assert ollama_url == "http://localhost:11434/v1"
|
||||
|
||||
|
||||
# ---- confirm_ollama_endpoint UX -------------------------------------------
|
||||
|
||||
|
||||
def test_confirm_endpoint_shows_default(monkeypatch, capsys):
|
||||
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
prompts.confirm_ollama_endpoint("http://localhost:11434/v1")
|
||||
out = _console_out(capsys)
|
||||
assert "http://localhost:11434/v1" in out
|
||||
assert "OLLAMA_BASE_URL" not in out # not from env
|
||||
assert "Note" not in out # no warnings for the canonical default
|
||||
|
||||
|
||||
def test_confirm_endpoint_marks_env_origin(monkeypatch, capsys):
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host:11434/v1")
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
prompts.confirm_ollama_endpoint("http://remote-host:11434/v1")
|
||||
out = _console_out(capsys)
|
||||
assert "http://remote-host:11434/v1" in out
|
||||
assert "OLLAMA_BASE_URL" in out
|
||||
|
||||
|
||||
def test_confirm_endpoint_warns_on_missing_scheme(monkeypatch, capsys):
|
||||
"""If user sets OLLAMA_BASE_URL=0.0.0.128, advise on the expected shape."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "0.0.0.128")
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
prompts.confirm_ollama_endpoint("0.0.0.128")
|
||||
out = _console_out(capsys)
|
||||
assert "missing a scheme" in out
|
||||
assert "http://<host>:11434/v1" in out
|
||||
|
||||
|
||||
def test_confirm_endpoint_warns_on_non_default_port_remote(monkeypatch, capsys):
|
||||
"""A remote host with no :11434 gets a soft hint about port mismatch."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host/v1")
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
prompts.confirm_ollama_endpoint("http://remote-host/v1")
|
||||
out = _console_out(capsys)
|
||||
assert "port 11434" in out
|
||||
|
||||
|
||||
def test_confirm_endpoint_quiet_on_local_no_port(monkeypatch, capsys):
|
||||
"""Local host without port shouldn't trigger the remote-port hint."""
|
||||
monkeypatch.setenv("OLLAMA_BASE_URL", "http://localhost/v1")
|
||||
from cli import prompts
|
||||
importlib.reload(prompts)
|
||||
prompts.confirm_ollama_endpoint("http://localhost/v1")
|
||||
out = _console_out(capsys)
|
||||
assert "Note" not in out # localhost is fine without explicit port
|
||||
|
||||
|
||||
def test_ollama_model_labels_no_local_suffix():
|
||||
"""Labels should no longer claim '(local)' since the endpoint is dynamic."""
|
||||
from tradingagents.llm_clients.model_catalog import get_model_options
|
||||
for mode in ("quick", "deep"):
|
||||
labels = [label for label, _ in get_model_options("ollama", mode)]
|
||||
assert all("local" not in label for label in labels), labels
|
||||
|
||||
|
||||
def test_ollama_offers_custom_model_id():
|
||||
"""Ollama users with custom-pulled models can pick 'Custom model ID'."""
|
||||
from tradingagents.llm_clients.model_catalog import get_model_options
|
||||
for mode in ("quick", "deep"):
|
||||
entries = get_model_options("ollama", mode)
|
||||
values = [v for _, v in entries]
|
||||
assert "custom" in values, f"Ollama {mode!r} missing 'custom' option: {entries}"
|
||||
# Custom option is last so it doesn't push the curated defaults off-screen
|
||||
assert values[-1] == "custom", f"'custom' should be last entry: {values}"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_structured_output_suppresses_object_tool_choice(monkeypatch):
|
||||
"""Ollama rejects the object-form tool_choice like other local servers
|
||||
(#1062), and a local model ID has no capability entry saying otherwise, so
|
||||
it takes the same client as the generic local endpoint."""
|
||||
from langchain_openai import ChatOpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
from tradingagents.llm_clients import create_llm_client
|
||||
|
||||
class Schema(BaseModel):
|
||||
x: int
|
||||
|
||||
captured = {}
|
||||
monkeypatch.setattr(
|
||||
ChatOpenAI,
|
||||
"with_structured_output",
|
||||
lambda self, schema, method=None, **kw: captured.update({"method": method, **kw}) or "BOUND",
|
||||
)
|
||||
|
||||
create_llm_client(provider="ollama", model="qwen3:30b").get_llm().with_structured_output(Schema)
|
||||
|
||||
assert captured["tool_choice"] is None
|
||||
@@ -0,0 +1,100 @@
|
||||
"""Generic OpenAI-compatible provider (vLLM / LM Studio / llama.cpp / relays).
|
||||
|
||||
Verifies the user-supplied base_url is required and honored, the key is optional
|
||||
(keyless local default), Chat Completions (not the Responses API) is used, any
|
||||
model name is accepted, and the env backend URL precedence (#978).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.api_key_env import get_api_key_env
|
||||
from tradingagents.llm_clients.factory import create_llm_client
|
||||
from tradingagents.llm_clients.validators import validate_model
|
||||
|
||||
# Note: assert by class NAME, not isinstance — other tests reload the
|
||||
# openai_client module, which would otherwise create a second class identity.
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_factory_routes_to_openai_client():
|
||||
client = create_llm_client(
|
||||
provider="openai_compatible", model="my-model", base_url="http://localhost:8000/v1"
|
||||
)
|
||||
assert type(client).__name__ == "OpenAIClient"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_base_url_required(monkeypatch):
|
||||
monkeypatch.delenv("OPENAI_COMPATIBLE_API_KEY", raising=False)
|
||||
with pytest.raises(ValueError, match="requires a base_url"):
|
||||
create_llm_client(provider="openai_compatible", model="m").get_llm()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_keyless_local_uses_placeholder_and_chat_completions(monkeypatch):
|
||||
monkeypatch.delenv("OPENAI_COMPATIBLE_API_KEY", raising=False)
|
||||
llm = create_llm_client(
|
||||
provider="openai_compatible", model="qwen2.5", base_url="http://localhost:8000/v1"
|
||||
).get_llm()
|
||||
assert type(llm).__name__ == "LocalCompatibleChatOpenAI"
|
||||
assert str(llm.openai_api_base) == "http://localhost:8000/v1"
|
||||
# keyless local servers: a placeholder key is sent
|
||||
key = llm.openai_api_key.get_secret_value() if hasattr(llm.openai_api_key, "get_secret_value") else llm.openai_api_key
|
||||
assert key == "EMPTY"
|
||||
# must use Chat Completions, not OpenAI's Responses API
|
||||
assert getattr(llm, "use_responses_api", False) in (False, None)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_optional_key_from_env(monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_COMPATIBLE_API_KEY", "sk-relay-123")
|
||||
llm = create_llm_client(
|
||||
provider="openai_compatible", model="m", base_url="https://relay.example/v1"
|
||||
).get_llm()
|
||||
key = llm.openai_api_key.get_secret_value() if hasattr(llm.openai_api_key, "get_secret_value") else llm.openai_api_key
|
||||
assert key == "sk-relay-123"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_any_model_accepted_no_forced_key():
|
||||
assert validate_model("openai_compatible", "literally-anything") is True
|
||||
# The key env exists (read for keyed relays) but the provider is marked
|
||||
# key-optional, so the CLI never forces a prompt and keyless servers work.
|
||||
assert get_api_key_env("openai_compatible") == "OPENAI_COMPATIBLE_API_KEY"
|
||||
from tradingagents.llm_clients.openai_client import OPENAI_COMPATIBLE_PROVIDERS
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["openai_compatible"].key_optional is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_env_backend_url_precedence():
|
||||
# #978: explicit env URL wins over the menu/default regardless of provider source.
|
||||
from cli.prompts import resolve_backend_url
|
||||
assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url="http://proxy/v1") == "http://proxy/v1"
|
||||
assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url=None) == "https://api.openai.com/v1"
|
||||
assert resolve_backend_url("deepseek", None, None) == "https://api.deepseek.com"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_structured_output_suppresses_object_tool_choice(monkeypatch):
|
||||
# LM Studio / vLLM reject the object-form tool_choice langchain sends for
|
||||
# function-calling structured output (#1057). The generic provider binds the
|
||||
# schema as a tool but must not force tool_choice.
|
||||
from langchain_openai import ChatOpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
class Schema(BaseModel):
|
||||
x: int
|
||||
|
||||
captured = {}
|
||||
monkeypatch.setattr(
|
||||
ChatOpenAI,
|
||||
"with_structured_output",
|
||||
lambda self, schema, method=None, **kw: captured.update({"method": method, **kw}) or "BOUND",
|
||||
)
|
||||
llm = create_llm_client(
|
||||
provider="openai_compatible", model="local-llm-30b", base_url="http://localhost:1234/v1"
|
||||
).get_llm()
|
||||
out = llm.with_structured_output(Schema)
|
||||
assert out == "BOUND"
|
||||
assert captured["method"] == "function_calling"
|
||||
assert captured["tool_choice"] is None # not the object form
|
||||
@@ -0,0 +1,43 @@
|
||||
"""OpenAI ``reasoning_effort`` is gated to reasoning models.
|
||||
|
||||
Non-reasoning OpenAI models (gpt-4.1, gpt-4o, ...) 400 with "Unsupported
|
||||
parameter: 'reasoning.effort'". The client must drop the kwarg for those rather
|
||||
than forward it and crash the run. The GPT-5 family and the o-series accept it.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.openai_client import (
|
||||
OpenAIClient,
|
||||
_supports_reasoning_effort,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected",
|
||||
[
|
||||
("gpt-5.5", True), ("gpt-5.4", True), ("gpt-5.4-mini", True),
|
||||
("gpt-5.5-pro", True), ("gpt-6-astra", True), ("o1", True), ("o3-mini", True),
|
||||
("gpt-4.1", False), ("gpt-4o", False), ("gpt-4o-mini", False),
|
||||
("gpt-3.5-turbo", False), ("gpt-10", True),
|
||||
("gpt-5foo", False), ("gpt-60x", False), ("o3rd-party", False),
|
||||
],
|
||||
)
|
||||
def test_supports_reasoning_effort(model, expected):
|
||||
assert _supports_reasoning_effort(model) is expected
|
||||
|
||||
|
||||
def _effort_on(model, monkeypatch):
|
||||
# A fake key lets get_llm() construct the client without a network call.
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
||||
llm = OpenAIClient(model, provider="openai", reasoning_effort="low").get_llm()
|
||||
return getattr(llm, "reasoning_effort", None)
|
||||
|
||||
|
||||
def test_reasoning_model_receives_effort(monkeypatch):
|
||||
assert _effort_on("gpt-5.4-mini", monkeypatch) == "low"
|
||||
|
||||
|
||||
def test_non_reasoning_model_drops_effort(monkeypatch):
|
||||
# gpt-4.1 would 400 with reasoning_effort — it must be dropped.
|
||||
assert _effort_on("gpt-4.1", monkeypatch) is None
|
||||
@@ -0,0 +1,43 @@
|
||||
"""The Responses API only exists on native OpenAI; a custom base_url on the
|
||||
openai provider must fall back to Chat Completions (#1024)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.openai_client import (
|
||||
OpenAIClient,
|
||||
_is_native_openai_base_url,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class NativeBaseUrlTests:
|
||||
def test_unset_is_native(self):
|
||||
assert _is_native_openai_base_url(None) is True
|
||||
assert _is_native_openai_base_url("") is True
|
||||
|
||||
def test_openai_hosts_are_native(self):
|
||||
assert _is_native_openai_base_url("https://api.openai.com/v1") is True
|
||||
assert _is_native_openai_base_url("api.openai.com/v1") is True
|
||||
|
||||
def test_custom_endpoints_are_not_native(self):
|
||||
assert _is_native_openai_base_url("http://localhost:1234/v1") is False
|
||||
assert _is_native_openai_base_url("https://my-gateway.example.com/v1") is False
|
||||
assert _is_native_openai_base_url("https://api.openai.com.evil.com/v1") is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class ResponsesApiSelectionTests:
|
||||
def test_native_openai_enables_responses_api(self, monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
|
||||
llm = OpenAIClient("gpt-5.5", provider="openai").get_llm()
|
||||
assert getattr(llm, "use_responses_api", False) is True
|
||||
|
||||
def test_custom_base_url_disables_responses_api(self, monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
|
||||
llm = OpenAIClient(
|
||||
"gpt-5.5", base_url="http://localhost:1234/v1", provider="openai"
|
||||
).get_llm()
|
||||
# use_responses_api should be absent/False so the client speaks Chat Completions.
|
||||
assert getattr(llm, "use_responses_api", False) is False
|
||||
@@ -0,0 +1,122 @@
|
||||
"""OpenRouter model selection: prompts are labeled by mode (#1000); required
|
||||
prompts exit cleanly on cancel; the output-language prompt defaults to English
|
||||
on cancel; and the OpenRouter list is newest-first."""
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
from cli import prompts
|
||||
|
||||
|
||||
def _asks(value):
|
||||
return mock.Mock(ask=mock.Mock(return_value=value))
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOpenRouterPromptLabel:
|
||||
@pytest.mark.parametrize("mode,label", [("quick", "Quick-Thinking"), ("deep", "Deep-Thinking")])
|
||||
def test_prompt_states_the_mode(self, mode, label):
|
||||
captured = {}
|
||||
|
||||
def fake_select(message, **kwargs):
|
||||
captured["message"] = message
|
||||
return _asks("openrouter/some-model")
|
||||
|
||||
with mock.patch.object(prompts, "_fetch_openrouter_models",
|
||||
return_value=[("Some Model", "openrouter/some-model")]), \
|
||||
mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
|
||||
out = prompts.select_openrouter_model(mode)
|
||||
|
||||
assert label in captured["message"]
|
||||
assert out == "openrouter/some-model"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOpenRouterLatestFirst:
|
||||
def test_models_sorted_newest_first(self):
|
||||
payload = {"data": [
|
||||
{"id": "old/model", "name": "Old", "created": 1000},
|
||||
{"id": "new/model", "name": "New", "created": 3000},
|
||||
{"id": "mid/model", "name": "Mid", "created": 2000},
|
||||
]}
|
||||
resp = mock.Mock()
|
||||
resp.json.return_value = payload
|
||||
resp.raise_for_status = mock.Mock()
|
||||
with mock.patch("requests.get", return_value=resp):
|
||||
out = prompts._fetch_openrouter_models()
|
||||
assert [mid for _, mid in out] == ["new/model", "mid/model", "old/model"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMainstreamFilter:
|
||||
def test_dropdown_prefers_mainstream_over_niche(self):
|
||||
# _fetch returns newest-first; the shortlist should drop niche namespaces.
|
||||
models = [
|
||||
("Fusion", "openrouter/fusion"),
|
||||
("Niche", "nex-agi/nex-n2-pro:free"),
|
||||
("Claude", "anthropic/claude-x"),
|
||||
("GPT", "openai/gpt-x"),
|
||||
]
|
||||
captured = {}
|
||||
|
||||
def fake_select(message, **kwargs):
|
||||
captured["values"] = [c.value for c in kwargs["choices"]]
|
||||
return _asks("anthropic/claude-x")
|
||||
|
||||
with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=models), \
|
||||
mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
|
||||
prompts.select_openrouter_model("quick")
|
||||
|
||||
assert "anthropic/claude-x" in captured["values"]
|
||||
assert "openai/gpt-x" in captured["values"]
|
||||
assert "openrouter/fusion" not in captured["values"]
|
||||
assert "nex-agi/nex-n2-pro:free" not in captured["values"]
|
||||
assert "custom" in captured["values"] # escape hatch preserved
|
||||
|
||||
def test_falls_back_to_all_when_no_mainstream(self):
|
||||
models = [("Niche", "nex-agi/x"), ("Other", "thedrummer/y")]
|
||||
captured = {}
|
||||
|
||||
def fake_select(message, **kwargs):
|
||||
captured["values"] = [c.value for c in kwargs["choices"]]
|
||||
return _asks("nex-agi/x")
|
||||
|
||||
with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=models), \
|
||||
mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
|
||||
prompts.select_openrouter_model("deep")
|
||||
|
||||
assert "nex-agi/x" in captured["values"] # fallback keeps the list usable
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCancelExitsCleanly:
|
||||
def test_dropdown_cancel_exits(self):
|
||||
with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=[]), \
|
||||
mock.patch.object(prompts.questionary, "select", return_value=_asks(None)), \
|
||||
pytest.raises(SystemExit):
|
||||
prompts.select_openrouter_model("quick")
|
||||
|
||||
def test_custom_id_cancel_exits(self):
|
||||
with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=[]), \
|
||||
mock.patch.object(prompts.questionary, "select", return_value=_asks("custom")), \
|
||||
mock.patch.object(prompts.questionary, "text", return_value=_asks(None)), \
|
||||
pytest.raises(SystemExit):
|
||||
prompts.select_openrouter_model("deep")
|
||||
|
||||
def test_prompt_custom_model_id_cancel_exits(self):
|
||||
with mock.patch.object(prompts.questionary, "text", return_value=_asks(None)), \
|
||||
pytest.raises(SystemExit):
|
||||
prompts._prompt_custom_model_id()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanguageDefaultsToEnglish:
|
||||
def test_select_cancel_defaults_english(self):
|
||||
with mock.patch.object(prompts.questionary, "select", return_value=_asks(None)):
|
||||
assert prompts.ask_output_language() == "English"
|
||||
|
||||
def test_custom_language_cancel_defaults_english(self):
|
||||
with mock.patch.object(prompts.questionary, "select", return_value=_asks("custom")), \
|
||||
mock.patch.object(prompts.questionary, "text", return_value=_asks(None)):
|
||||
assert prompts.ask_output_language() == "English"
|
||||
@@ -0,0 +1,130 @@
|
||||
"""Polymarket prediction-market vendor: forward-looking filtering, volume
|
||||
ranking, formatting, graceful degradation, and router integration.
|
||||
|
||||
All API access is mocked, so these run without a network connection.
|
||||
"""
|
||||
import copy
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
import tradingagents.dataflows.config as config_module
|
||||
import tradingagents.default_config as default_config
|
||||
from tradingagents.dataflows import router
|
||||
from tradingagents.dataflows.config import set_config
|
||||
from tradingagents.dataflows.vendors import polymarket
|
||||
|
||||
|
||||
def _market(question, prob, *, volume, end_date, closed=False, wk=None):
|
||||
return {
|
||||
"question": question,
|
||||
"outcomes": '["Yes", "No"]',
|
||||
"outcomePrices": f'["{prob}", "{round(1 - prob, 4)}"]',
|
||||
"volumeNum": volume,
|
||||
"endDate": end_date,
|
||||
"closed": closed,
|
||||
"oneWeekPriceChange": wk,
|
||||
}
|
||||
|
||||
|
||||
# One event with a mix: a high-volume open market, a closed one, a past-dated
|
||||
# one, and a lower-volume open one. Far-future / far-past dates keep the test
|
||||
# independent of the real clock.
|
||||
_SEARCH = {
|
||||
"events": [
|
||||
{
|
||||
"markets": [
|
||||
_market("Open big?", 0.76, volume=5_000_000, end_date="2030-12-31T00:00:00Z", wk=-0.045),
|
||||
_market("Resolved already?", 1.0, volume=9_000_000, end_date="2030-12-31T00:00:00Z", closed=True),
|
||||
_market("Past event?", 0.5, volume=8_000_000, end_date="2020-01-01T00:00:00Z"),
|
||||
_market("Open small?", 0.30, volume=1_000, end_date="2030-06-30T00:00:00Z"),
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class PolymarketFilterTests(unittest.TestCase):
|
||||
def test_closed_and_past_markets_are_excluded(self):
|
||||
with mock.patch.object(polymarket, "_request", return_value=_SEARCH):
|
||||
out = polymarket.get_prediction_markets("anything", limit=10)
|
||||
self.assertIn("Open big?", out)
|
||||
self.assertIn("Open small?", out)
|
||||
self.assertNotIn("Resolved already?", out) # closed
|
||||
self.assertNotIn("Past event?", out) # endDate in the past
|
||||
|
||||
def test_ranked_by_volume(self):
|
||||
with mock.patch.object(polymarket, "_request", return_value=_SEARCH):
|
||||
out = polymarket.get_prediction_markets("anything", limit=10)
|
||||
self.assertLess(out.index("Open big?"), out.index("Open small?"))
|
||||
|
||||
def test_limit_caps_results(self):
|
||||
with mock.patch.object(polymarket, "_request", return_value=_SEARCH):
|
||||
out = polymarket.get_prediction_markets("anything", limit=1)
|
||||
self.assertIn("Open big?", out)
|
||||
self.assertNotIn("Open small?", out)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class PolymarketFormatTests(unittest.TestCase):
|
||||
def test_probability_volume_and_weekly_change_render(self):
|
||||
with mock.patch.object(polymarket, "_request", return_value=_SEARCH):
|
||||
out = polymarket.get_prediction_markets("anything", limit=10)
|
||||
self.assertIn("Yes 76%", out)
|
||||
self.assertIn("$5,000,000 volume", out)
|
||||
self.assertIn("resolves 2030-12-31", out)
|
||||
self.assertIn("1-week -4.5pp", out) # -0.045 -> -4.5pp
|
||||
|
||||
def test_weekly_change_omitted_when_absent(self):
|
||||
# "Open small?" has wk=None -> no 1-week clause on its line.
|
||||
with mock.patch.object(polymarket, "_request", return_value=_SEARCH):
|
||||
out = polymarket.get_prediction_markets("anything", limit=10)
|
||||
small_line = next(ln for ln in out.splitlines() if "Open small?" in ln)
|
||||
self.assertNotIn("1-week", small_line)
|
||||
|
||||
def test_no_matches_reports_clearly(self):
|
||||
with mock.patch.object(polymarket, "_request", return_value={"events": []}):
|
||||
out = polymarket.get_prediction_markets("obscure ticker", limit=6)
|
||||
self.assertIn("No open prediction markets", out)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class PolymarketResilienceTests(unittest.TestCase):
|
||||
def test_network_error_degrades_gracefully(self):
|
||||
# An external-service hiccup must not raise into the analyst.
|
||||
with mock.patch.object(
|
||||
polymarket, "_request", side_effect=requests.RequestException("boom")
|
||||
):
|
||||
out = polymarket.get_prediction_markets("Fed rate cut")
|
||||
self.assertIn("unavailable", out.lower())
|
||||
self.assertIn("Fed rate cut", out)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class PolymarketRoutingTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
|
||||
def tearDown(self):
|
||||
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
|
||||
|
||||
def test_category_routes_to_polymarket(self):
|
||||
self.assertEqual(
|
||||
router.get_category_for_method("get_prediction_markets"),
|
||||
"prediction_markets",
|
||||
)
|
||||
set_config({"data_vendors": {"prediction_markets": "polymarket"}})
|
||||
with mock.patch.dict(
|
||||
router.VENDOR_METHODS,
|
||||
{"get_prediction_markets": {"polymarket": lambda *a, **k: "POLY_OK"}},
|
||||
clear=False,
|
||||
):
|
||||
out = router.route_to_vendor("get_prediction_markets", "fed", 5)
|
||||
self.assertEqual(out, "POLY_OK")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,240 @@
|
||||
"""Portfolio context: what the caller holds, threaded into the decision agents.
|
||||
|
||||
Decisions were made with no knowledge of the current book, so "add to a full
|
||||
position" and "open a new one" read alike. The context is optional and carries
|
||||
three distinct states: a position, a flat book, and no context at all. Nothing
|
||||
may present the third as the second. The research team stays blind so the bull
|
||||
and bear cases are not anchored by the caller's position.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.agents.context import get_portfolio_context_from_state
|
||||
from tradingagents.portfolio import PortfolioContext, load_portfolio
|
||||
|
||||
HOLDING = {
|
||||
"cash": 25000.0,
|
||||
"currency": "USD",
|
||||
"positions": [
|
||||
{"ticker": "AAPL", "quantity": 120, "average_price": 150.0},
|
||||
{"ticker": "MSFT", "quantity": 10},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_position_in_the_analyzed_instrument_leads_the_render():
|
||||
text = PortfolioContext.model_validate(HOLDING).render("AAPL")
|
||||
assert "120" in text and "150" in text
|
||||
assert "MSFT" in text and "25,000" in text and "USD" in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_flat_book_says_no_position_rather_than_omitting_it():
|
||||
text = PortfolioContext.model_validate({"cash": 1000.0, "positions": []}).render("AAPL")
|
||||
assert "No current position in AAPL" in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_ticker_held_under_another_spelling_is_matched():
|
||||
text = PortfolioContext.model_validate({"positions": [{"ticker": "aapl", "quantity": 5}]}).render("AAPL")
|
||||
assert "No current position" not in text and "5" in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_absent_context_is_reported_as_not_provided():
|
||||
notice = get_portfolio_context_from_state({"company_of_interest": "AAPL"})
|
||||
assert "not provided" in notice.lower()
|
||||
assert "no position" not in notice.lower() # missing must not read as flat
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_rendered_context_reaches_the_agents_from_state():
|
||||
block = get_portfolio_context_from_state({"portfolio_context": "Portfolio: flat", "company_of_interest": "AAPL"})
|
||||
assert block == "Portfolio: flat"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_load_rejects_a_malformed_file_with_a_clear_error(tmp_path):
|
||||
bad = tmp_path / "p.json"
|
||||
bad.write_text(json.dumps({"positions": [{"quantity": 5}]}))
|
||||
with pytest.raises(ValueError, match="portfolio"):
|
||||
load_portfolio(bad)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_load_reads_a_valid_file(tmp_path):
|
||||
good = tmp_path / "p.json"
|
||||
good.write_text(json.dumps(HOLDING))
|
||||
assert load_portfolio(good).positions[0].ticker == "AAPL"
|
||||
|
||||
|
||||
# --- threading through the graph --------------------------------------------
|
||||
|
||||
def _bare_graph(tmp_path):
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
from tradingagents.graph.propagation import Propagator
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
graph = object.__new__(TradingAgentsGraph)
|
||||
graph.config = {"memory_log_path": str(tmp_path / "m.md"), "max_debate_rounds": 1,
|
||||
"max_risk_discuss_rounds": 1}
|
||||
graph.memory_log = TradingMemoryLog(graph.config)
|
||||
graph.propagator = Propagator()
|
||||
graph.selected_analysts = ["market"]
|
||||
graph.settle_pending = lambda t: None
|
||||
graph.resolve_instrument_context = lambda t, a="stock", d=None: ""
|
||||
graph._memory_as_of = lambda d: None
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_create_run_state_renders_the_portfolio_once(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
state = graph.create_run_state("AAPL", "2026-08-14", portfolio=PortfolioContext.model_validate(HOLDING))
|
||||
assert "120" in state["portfolio_context"]
|
||||
assert graph.create_run_state("AAPL", "2026-08-14")["portfolio_context"] == ""
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_checkpoint_signature_changes_with_the_portfolio(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
none = graph._run_signature("stock")
|
||||
flat = graph._run_signature("stock", PortfolioContext())
|
||||
held = graph._run_signature("stock", PortfolioContext.model_validate(HOLDING))
|
||||
assert len({none, flat, held}) == 3
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("module, factory", [
|
||||
("tradingagents.agents.trader.trader", "create_trader"),
|
||||
("tradingagents.agents.managers.portfolio_manager", "create_portfolio_manager"),
|
||||
("tradingagents.agents.risk_mgmt.aggressive_debator", "create_aggressive_debator"),
|
||||
("tradingagents.agents.risk_mgmt.conservative_debator", "create_conservative_debator"),
|
||||
("tradingagents.agents.risk_mgmt.neutral_debator", "create_neutral_debator"),
|
||||
])
|
||||
def test_decision_agents_see_the_portfolio(module, factory, monkeypatch):
|
||||
"""The prompt each decision agent sends carries the portfolio block."""
|
||||
import importlib
|
||||
|
||||
mod = importlib.import_module(module)
|
||||
seen = []
|
||||
|
||||
class _LLM:
|
||||
def invoke(self, prompt, *a, **k):
|
||||
seen.append(prompt if isinstance(prompt, str) else json.dumps(str(prompt)))
|
||||
from langchain_core.messages import AIMessage
|
||||
return AIMessage("Rating: Hold\n\nnothing to do")
|
||||
|
||||
def with_structured_output(self, *a, **k):
|
||||
raise NotImplementedError # force the free-text path
|
||||
|
||||
state = {
|
||||
"company_of_interest": "AAPL", "trade_date": "2026-08-14", "asset_type": "stock",
|
||||
"instrument_context": "", "market_report": "M", "sentiment_report": "S",
|
||||
"news_report": "N", "fundamentals_report": "F", "investment_plan": "P",
|
||||
"trader_investment_plan": "T", "past_context": "",
|
||||
"portfolio_context": "PORTFOLIO_BLOCK_MARKER",
|
||||
"investment_debate_state": {"history": "", "judge_decision": "", "count": 0},
|
||||
"risk_debate_state": {"history": "", "latest_speaker": "", "count": 0,
|
||||
"aggressive_history": "", "conservative_history": "", "neutral_history": "",
|
||||
"current_aggressive_response": "", "current_conservative_response": "",
|
||||
"current_neutral_response": "", "judge_decision": ""},
|
||||
}
|
||||
node = getattr(mod, factory)(_LLM())
|
||||
node(state)
|
||||
assert any("PORTFOLIO_BLOCK_MARKER" in p for p in seen), f"{factory} prompt lacks the portfolio block"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_research_team_stays_blind_to_the_portfolio():
|
||||
import inspect
|
||||
|
||||
from tradingagents.agents.researchers import bear_researcher, bull_researcher
|
||||
for mod in (bull_researcher, bear_researcher):
|
||||
assert "portfolio_context" not in inspect.getsource(mod)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_completed_run_clears_the_checkpoint_it_wrote(tmp_path, monkeypatch):
|
||||
"""The clear must key on the same portfolio the run was checkpointed under.
|
||||
|
||||
Keyed on a different one it deletes nothing, and the next identical call
|
||||
resumes the finished thread and returns the old decision without running.
|
||||
"""
|
||||
import tradingagents.graph.trading_graph as tg
|
||||
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.config.update({"checkpoint_enabled": True, "data_cache_dir": str(tmp_path),
|
||||
"results_dir": str(tmp_path)})
|
||||
graph.debug = False
|
||||
graph._resuming = False
|
||||
graph.propagator.get_graph_args = lambda callbacks=None: {}
|
||||
graph.process_signal = lambda d: "Hold"
|
||||
graph._log_state = lambda *a, **k: None
|
||||
graph.graph = type("G", (), {"invoke": lambda self, i, **k: {"final_trade_decision": "Rating: Hold\n\nx"}})()
|
||||
book = PortfolioContext.model_validate(HOLDING)
|
||||
|
||||
written = graph._run_signature("stock", book) # what begin_checkpoint keys on
|
||||
cleared = []
|
||||
monkeypatch.setattr(tg, "clear_checkpoint", lambda d, t, dt, signature: cleared.append(signature))
|
||||
|
||||
graph._run_graph("AAPL", "2026-08-14", "stock", checkpoint_thread_id=None, portfolio=book)
|
||||
|
||||
assert cleared == [written]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_research_layer_sizes_against_a_standard_allocation():
|
||||
"""The research team is blind to the book, so its plan cannot promise
|
||||
position-relative sizing: it sizes against a standard allocation instead."""
|
||||
from tradingagents.agents.schemas import ResearchPlan
|
||||
|
||||
description = ResearchPlan.model_fields["strategic_actions"].description
|
||||
assert "standard allocation" in description
|
||||
assert "does not see the caller's holdings" in description
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_partial_portfolio_states_only_what_it_was_given():
|
||||
"""Cash omitted is not cash zero; the line is absent rather than invented."""
|
||||
text = PortfolioContext.model_validate({"positions": [{"ticker": "AAPL", "quantity": 5}]}).render("AAPL")
|
||||
assert "Cash" not in text
|
||||
assert "5" in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cli_rejects_an_unusable_portfolio_file_before_running(tmp_path, monkeypatch):
|
||||
from typer.testing import CliRunner
|
||||
|
||||
import cli.main as m
|
||||
|
||||
bad = tmp_path / "bad.json"
|
||||
bad.write_text('{"positions": [{"quantity": 5}]}')
|
||||
ran = []
|
||||
monkeypatch.setattr(m, "run_analysis", lambda **k: ran.append(k))
|
||||
|
||||
result = CliRunner().invoke(m.app, ["--portfolio", str(bad)])
|
||||
|
||||
assert result.exit_code == 1 and ran == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cli_passes_a_valid_portfolio_into_the_run(tmp_path, monkeypatch):
|
||||
from typer.testing import CliRunner
|
||||
|
||||
import cli.main as m
|
||||
|
||||
good = tmp_path / "good.json"
|
||||
good.write_text(json.dumps(HOLDING))
|
||||
ran = []
|
||||
monkeypatch.setattr(m, "run_analysis", lambda **k: ran.append(k))
|
||||
|
||||
result = CliRunner().invoke(m.app, ["--portfolio", str(good)])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert ran[0]["portfolio"].position_in("AAPL").quantity == 120
|
||||
@@ -0,0 +1,288 @@
|
||||
"""Jev post screening, against TypeSafe's documented request and response shapes."""
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from tradingagents.agents import post_screen as typesafe
|
||||
|
||||
QUESTIONS = {"is_urgent": {"type": "noul", "instructions": "Does this convey urgency?"}}
|
||||
ANSWERS = {"is_urgent": {"type": "noul", "noul": 0.95}}
|
||||
|
||||
|
||||
class _Response:
|
||||
def __init__(self, status, payload=None, headers=None):
|
||||
self.status_code = status
|
||||
self._payload = payload
|
||||
self.headers = headers or {}
|
||||
|
||||
def json(self):
|
||||
if self._payload is None:
|
||||
raise ValueError("no JSON")
|
||||
return self._payload
|
||||
|
||||
|
||||
class _Calls(list):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.queue = []
|
||||
self.sleeps = []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def post(monkeypatch):
|
||||
"""Queue responses on ``.queue``; the list records each call to requests.post."""
|
||||
calls = _Calls()
|
||||
queue = calls.queue
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls.append((url, kwargs))
|
||||
item = queue.pop(0)
|
||||
if isinstance(item, Exception):
|
||||
raise item
|
||||
return item
|
||||
|
||||
monkeypatch.setattr(typesafe.requests, "post", fake_post)
|
||||
monkeypatch.setattr(typesafe.time, "sleep", calls.sleeps.append)
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "ts-test")
|
||||
monkeypatch.delenv("TYPESAFE_DEFAULT_MODEL", raising=False)
|
||||
return calls
|
||||
|
||||
|
||||
def _ok():
|
||||
return _Response(200, {"model": "jev-1.13.0", "answers": ANSWERS,
|
||||
"usage": {"input_tokens": 296, "output_tokens": 20}})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_sends_the_documented_request_and_returns_the_answers(post):
|
||||
post.queue.append(_ok())
|
||||
|
||||
assert typesafe.system_one("Help! My payouts have been failing.", QUESTIONS) == ANSWERS
|
||||
|
||||
url, kwargs = post[0]
|
||||
assert url == "https://api.typesafe.ai/v1/systemone"
|
||||
assert kwargs["headers"]["Authorization"] == "Bearer ts-test"
|
||||
assert kwargs["json"] == {"state": "Help! My payouts have been failing.",
|
||||
"model": "jev-latest", "questions": QUESTIONS}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_model_follows_the_sdk_environment(post, monkeypatch):
|
||||
monkeypatch.setenv("TYPESAFE_DEFAULT_MODEL", "jev-1.13.0")
|
||||
post.queue.append(_ok())
|
||||
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
|
||||
assert post[0][1]["json"]["model"] == "jev-1.13.0"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("transient", [
|
||||
_Response(429), _Response(529), requests.ConnectionError(), requests.Timeout(),
|
||||
requests.exceptions.ChunkedEncodingError(),
|
||||
])
|
||||
def test_rate_limits_overload_and_dropped_connections_are_retried(post, transient):
|
||||
post.queue.extend([transient, _ok()])
|
||||
|
||||
assert typesafe.system_one("s", QUESTIONS) == ANSWERS
|
||||
assert len(post) == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_retry_after_header_sets_the_wait(post):
|
||||
post.queue.extend([_Response(429, headers={"retry-after": "7"}), _ok()])
|
||||
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
|
||||
assert post.sleeps == [7.0]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_long_retry_after_is_capped(post):
|
||||
post.queue.extend([_Response(529, headers={"retry-after": "600"}), _ok()])
|
||||
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
|
||||
assert post.sleeps == [typesafe._MAX_WAIT]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_retries_are_bounded(post):
|
||||
post.queue.extend([_Response(529)] * 3)
|
||||
|
||||
with pytest.raises(typesafe.TypeSafeError, match="HTTP 529"):
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
assert len(post) == 3
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("error", [requests.exceptions.InvalidHeader(), requests.exceptions.TooManyRedirects()])
|
||||
def test_other_request_errors_are_screening_failures_without_retry(post, error):
|
||||
post.queue.append(error)
|
||||
|
||||
with pytest.raises(typesafe.TypeSafeError, match=type(error).__name__):
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
assert len(post) == 1
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("status", [401, 422, 500])
|
||||
def test_other_failures_raise_without_retry(post, status):
|
||||
post.queue.append(_Response(status))
|
||||
|
||||
with pytest.raises(typesafe.TypeSafeError, match=f"HTTP {status}"):
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
assert len(post) == 1
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("payload", [
|
||||
None, {"model": "jev"}, {"answers": {"other": {}}},
|
||||
{"answers": {"is_urgent": {"type": "choice", "choice": "yes"}}},
|
||||
])
|
||||
def test_a_response_without_every_answer_is_an_error(post, payload):
|
||||
post.queue.append(_Response(200, payload))
|
||||
|
||||
with pytest.raises(typesafe.TypeSafeError, match="malformed"):
|
||||
typesafe.system_one("s", QUESTIONS)
|
||||
|
||||
|
||||
|
||||
def _post_answers(about: float, stance: str = "bullish", confidence: float = 0.9):
|
||||
return _Response(200, {"model": "jev-1.13.0", "answers": {
|
||||
"about": {"type": "noul", "noul": about},
|
||||
"stance": {"type": "choice", "choice": stance, "confidence": confidence,
|
||||
"probabilities": {stance: 1.0}},
|
||||
}})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def jev(post, monkeypatch):
|
||||
"""Answer each post by its text: ``post`` maps text -> response."""
|
||||
answers = {}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
post.append((url, kwargs))
|
||||
return answers[kwargs["json"]["state"]["post"]]
|
||||
|
||||
monkeypatch.setattr(typesafe.requests, "post", fake_post)
|
||||
monkeypatch.setattr(typesafe, "resolve_instrument_identity",
|
||||
lambda t: {"company_name": "NVIDIA Corporation"})
|
||||
return answers
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_no_key_no_screen(monkeypatch):
|
||||
monkeypatch.delenv("TYPESAFE_API_KEY", raising=False)
|
||||
assert typesafe.jev_screen("NVDA") is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_each_post_is_its_own_state_under_the_fixed_questions(jev, post):
|
||||
jev["NVDA to 200"] = _post_answers(0.9)
|
||||
|
||||
typesafe.jev_screen("NVDA")(["NVDA to 200"])
|
||||
|
||||
body = post[0][1]["json"]
|
||||
assert body["state"] == {"instrument": "NVIDIA Corporation (NVDA)", "post": "NVDA to 200"}
|
||||
assert body["questions"] == typesafe.QUESTIONS
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_screen_drops_clear_off_topic_posts_and_counts_confident_stances(jev):
|
||||
jev.update({
|
||||
"long NVDA": _post_answers(0.95, "bullish"),
|
||||
"NVDA puts": _post_answers(0.9, "bearish"),
|
||||
"maybe NVDA": _post_answers(0.4, "neutral"), # uncertain relevance: kept
|
||||
"NVDA?": _post_answers(0.8, "bullish", 0.3), # uncertain stance: unclear
|
||||
"NVDA!": _post_answers(0.8, "sideways"), # not an option: unclear
|
||||
"$AAPL $MSFT $NVDA pump": _post_answers(0.1, "bullish"),
|
||||
})
|
||||
|
||||
keep, note = typesafe.jev_screen("NVDA")(list(jev))
|
||||
|
||||
assert keep == [True, True, True, True, True, False]
|
||||
assert note == ("Screened by Jev: 5 of the 6 posts fetched are about NVIDIA Corporation (NVDA); "
|
||||
"their stance on its stock: 1 bullish, 1 bearish, 1 neutral, 2 unclear.")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_one_failed_request_leaves_every_post_unscreened(jev):
|
||||
jev.update({"a": _post_answers(0.1), "b": _Response(401)})
|
||||
|
||||
keep, note = typesafe.jev_screen("NVDA")(["a", "b"])
|
||||
|
||||
assert keep == [True, True]
|
||||
assert note == "<Jev screening unavailable (HTTP 401); posts are unscreened>"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_answer_missing_its_fields_reads_as_malformed(jev):
|
||||
jev["a"] = _Response(200, {"answers": {"about": {"type": "noul"},
|
||||
"stance": {"type": "choice"}}})
|
||||
|
||||
keep, note = typesafe.jev_screen("NVDA")(["a"])
|
||||
|
||||
assert keep == [True]
|
||||
assert note == "<Jev screening unavailable (malformed response); posts are unscreened>"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_first_failure_cancels_the_requests_not_yet_sent(jev, post, monkeypatch):
|
||||
monkeypatch.setattr(typesafe, "_WORKERS", 1)
|
||||
jev.update({"a": _Response(401), **{f"p{i}": _post_answers(0.9) for i in range(20)}})
|
||||
|
||||
typesafe.jev_screen("NVDA")(list(jev))
|
||||
|
||||
assert len(post) < 21
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_sentiment_analyst_hands_the_screen_to_both_social_fetchers(monkeypatch):
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
from tradingagents.agents.analysts import sentiment_analyst
|
||||
|
||||
screen = object()
|
||||
seen = []
|
||||
monkeypatch.setattr(sentiment_analyst, "jev_screen", lambda ticker: screen)
|
||||
monkeypatch.setattr(sentiment_analyst.get_news, "func", lambda *a: "news")
|
||||
for name in ("fetch_stocktwits_messages", "fetch_reddit_posts"):
|
||||
monkeypatch.setattr(sentiment_analyst, name, lambda *a, screen=None, **k: seen.append(screen) or "")
|
||||
|
||||
class _LLM:
|
||||
def with_structured_output(self, *a, **k):
|
||||
raise NotImplementedError
|
||||
|
||||
def invoke(self, messages):
|
||||
return AIMessage(content="report")
|
||||
|
||||
node = sentiment_analyst.create_sentiment_analyst(_LLM())
|
||||
node({"company_of_interest": "NVDA", "trade_date": "2026-01-09", "messages": []})
|
||||
|
||||
assert seen == [screen, screen]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_failure_does_not_wait_for_requests_still_in_flight(jev, monkeypatch):
|
||||
import threading
|
||||
import time
|
||||
|
||||
release = threading.Event()
|
||||
|
||||
class _Slow:
|
||||
status_code = 200
|
||||
headers = {}
|
||||
|
||||
def json(self):
|
||||
release.wait(5)
|
||||
return _post_answers(0.9).json()
|
||||
|
||||
jev.update({"slow": _Slow(), "bad": _Response(401)})
|
||||
started = time.monotonic()
|
||||
keep, note = typesafe.jev_screen("NVDA")(["slow", "bad"])
|
||||
elapsed = time.monotonic() - started
|
||||
release.set()
|
||||
|
||||
assert keep == [True, True] and "unavailable" in note
|
||||
assert elapsed < 2
|
||||
@@ -0,0 +1,87 @@
|
||||
"""What the agents are actually told.
|
||||
|
||||
Three problems the audit found: one analyst's brief reached the model as a Python
|
||||
tuple, every analyst was asked for a trade call that nothing reads, and a report
|
||||
that was never produced was presented as an empty labelled section, which invites
|
||||
the next agent to fill it in from nothing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
|
||||
ANALYSTS = ["market_analyst", "sentiment_analyst", "news_analyst", "fundamentals_analyst"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("name", ANALYSTS)
|
||||
def test_an_analyst_brief_is_text_not_a_python_object(name):
|
||||
"""A trailing comma made one brief a tuple, so the model was handed its repr
|
||||
(quotes, parens and all) instead of the instruction."""
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
mod = importlib.import_module(f"tradingagents.agents.analysts.{name}")
|
||||
tree = ast.parse(inspect.getsource(mod))
|
||||
briefs = [node.value for node in ast.walk(tree)
|
||||
if isinstance(node, ast.Assign)
|
||||
and getattr(node.targets[0], "id", "") == "system_message"]
|
||||
assert briefs, f"{name} has no system_message"
|
||||
for brief in briefs:
|
||||
assert not isinstance(brief, ast.Tuple), "the brief is a tuple, not text"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("name", ANALYSTS)
|
||||
def test_an_analyst_is_not_asked_for_a_trade_call_nothing_reads(name):
|
||||
"""The stop signal is never consumed, and asking for it makes an analyst
|
||||
open with a direction that then travels as evidence."""
|
||||
import inspect
|
||||
|
||||
mod = importlib.import_module(f"tradingagents.agents.analysts.{name}")
|
||||
assert "FINAL TRANSACTION PROPOSAL" not in inspect.getsource(mod)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("module, factory", [
|
||||
("tradingagents.agents.researchers.bull_researcher", "create_bull_researcher"),
|
||||
("tradingagents.agents.researchers.bear_researcher", "create_bear_researcher"),
|
||||
("tradingagents.agents.risk_mgmt.aggressive_debator", "create_aggressive_debator"),
|
||||
("tradingagents.agents.risk_mgmt.conservative_debator", "create_conservative_debator"),
|
||||
("tradingagents.agents.risk_mgmt.neutral_debator", "create_neutral_debator"),
|
||||
])
|
||||
def test_a_report_that_was_never_produced_says_so(module, factory):
|
||||
"""`--analysts market` leaves three reports empty; presenting them as blank
|
||||
sections invites the model to invent the contents."""
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
mod = importlib.import_module(module)
|
||||
seen = []
|
||||
|
||||
class _LLM:
|
||||
def invoke(self, prompt, *a, **k):
|
||||
seen.append(prompt if isinstance(prompt, str) else str(prompt))
|
||||
return AIMessage("argument")
|
||||
|
||||
def with_structured_output(self, *a, **k):
|
||||
raise NotImplementedError
|
||||
|
||||
state = {
|
||||
"company_of_interest": "NVDA", "trade_date": "2026-08-14", "asset_type": "stock",
|
||||
"instrument_context": "", "portfolio_context": "", "past_context": "",
|
||||
"market_report": "RSI 61, price 178.", "sentiment_report": "", "news_report": "",
|
||||
"fundamentals_report": "", "investment_plan": "P", "trader_investment_plan": "T",
|
||||
"investment_debate_state": {"bull_history": "", "bear_history": "", "history": "",
|
||||
"current_response": "", "judge_decision": "", "count": 0},
|
||||
"risk_debate_state": {"history": "", "latest_speaker": "", "count": 0,
|
||||
"aggressive_history": "", "conservative_history": "", "neutral_history": "",
|
||||
"current_aggressive_response": "", "current_conservative_response": "",
|
||||
"current_neutral_response": "", "judge_decision": ""},
|
||||
}
|
||||
getattr(mod, factory)(_LLM())(state)
|
||||
|
||||
prompt = " ".join(seen)
|
||||
assert "not part of this run" in prompt or "not available" in prompt, prompt[:400]
|
||||
assert "RSI 61" in prompt # the report that does exist is still passed through
|
||||
@@ -0,0 +1,60 @@
|
||||
"""The OpenAI-compatible provider registry is the single source of truth for the
|
||||
family; this guards each provider's resolved config (base URL, subclass, auth,
|
||||
Responses API) so a future edit can't silently break one.
|
||||
"""
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.openai_client import (
|
||||
OPENAI_COMPATIBLE_PROVIDERS,
|
||||
DeepSeekChatOpenAI,
|
||||
LocalCompatibleChatOpenAI,
|
||||
MinimaxChatOpenAI,
|
||||
NormalizedChatOpenAI,
|
||||
is_openai_compatible,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_registry_membership():
|
||||
assert is_openai_compatible("openai")
|
||||
assert is_openai_compatible("openai_compatible") # the generic endpoint
|
||||
# native (different API) clients are intentionally NOT in the registry
|
||||
assert not is_openai_compatible("anthropic")
|
||||
assert not is_openai_compatible("google")
|
||||
assert not is_openai_compatible("azure")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("provider,base_url,chat_class,responses", [
|
||||
("openai", None, NormalizedChatOpenAI, True),
|
||||
("xai", "https://api.x.ai/v1", NormalizedChatOpenAI, False),
|
||||
("deepseek", "https://api.deepseek.com", DeepSeekChatOpenAI, False),
|
||||
("qwen", "https://dashscope-intl.aliyuncs.com/compatible-mode/v1", NormalizedChatOpenAI, False),
|
||||
("qwen-cn", "https://dashscope.aliyuncs.com/compatible-mode/v1", NormalizedChatOpenAI, False),
|
||||
("glm", "https://api.z.ai/api/paas/v4/", NormalizedChatOpenAI, False),
|
||||
("glm-cn", "https://open.bigmodel.cn/api/paas/v4/", NormalizedChatOpenAI, False),
|
||||
("minimax", "https://api.minimax.io/v1", MinimaxChatOpenAI, False),
|
||||
("minimax-cn", "https://api.minimaxi.com/v1", MinimaxChatOpenAI, False),
|
||||
("openrouter", "https://openrouter.ai/api/v1", NormalizedChatOpenAI, False),
|
||||
("mistral", "https://api.mistral.ai/v1", NormalizedChatOpenAI, False),
|
||||
("kimi", "https://api.moonshot.ai/v1", NormalizedChatOpenAI, False),
|
||||
("groq", "https://api.groq.com/openai/v1", NormalizedChatOpenAI, False),
|
||||
("nvidia", "https://integrate.api.nvidia.com/v1", NormalizedChatOpenAI, False),
|
||||
("ollama", "http://localhost:11434/v1", LocalCompatibleChatOpenAI, False),
|
||||
])
|
||||
def test_registry_spec(provider, base_url, chat_class, responses):
|
||||
spec = OPENAI_COMPATIBLE_PROVIDERS[provider]
|
||||
assert spec.base_url == base_url
|
||||
assert spec.chat_class is chat_class
|
||||
assert spec.use_responses_api is responses
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_key_optionality():
|
||||
# Local/generic endpoints are key-optional; hosted APIs require a key.
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["ollama"].key_optional is True
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["openai_compatible"].key_optional is True
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["openai_compatible"].require_base_url is True
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["xai"].key_optional is False
|
||||
# OLLAMA_BASE_URL is the only base-URL env override.
|
||||
assert OPENAI_COMPATIBLE_PROVIDERS["ollama"].base_url_env == "OLLAMA_BASE_URL"
|
||||
@@ -0,0 +1,212 @@
|
||||
"""A decision is recorded as the call that was made, or as needing review.
|
||||
|
||||
Two readers used to disagree about the same text: the signal said REVIEW while
|
||||
the memory log wrote a fabricated Hold. Worse, prose that argued against a Buy
|
||||
before concluding Underweight was read as Buy, because the parser took the first
|
||||
rating word anywhere in the document. A wrong direction is worse than no
|
||||
direction, so an unclear decision is REVIEW everywhere.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
import cli.run as cli_run
|
||||
from tradingagents.agents.rating import RATING_REVIEW, extract_rating, parse_rating
|
||||
|
||||
INVERTED = ("The aggressive analyst pushed hard for a Buy on the AI backlog, but the "
|
||||
"conservative case on margin compression carried the debate. "
|
||||
"Final rating — Underweight. Trim to half weight over the next two weeks.")
|
||||
REFUSAL = "I'm sorry, I can't provide a rating for this security."
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("separator", [":", "-", "—", "–", ":", ": **"])
|
||||
def test_the_labelled_rating_wins_whatever_separates_it(separator):
|
||||
text = f"Buy arguments were raised and rejected.\n\nRating{separator}Underweight\n\nTrim."
|
||||
assert extract_rating(text) == "Underweight"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_rating_argued_against_is_not_read_as_the_decision():
|
||||
assert extract_rating(INVERTED) == "Underweight"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_prose_naming_several_ratings_without_a_label_needs_review():
|
||||
"""Nothing in the text says which one is the call, so guessing risks
|
||||
reporting the opposite of the decision."""
|
||||
text = "The bull wants Buy, the bear wants Sell, and the committee was split."
|
||||
assert extract_rating(text) is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_prose_naming_one_rating_is_taken_as_the_call():
|
||||
assert extract_rating("On balance we stay Underweight until margins recover.") == "Underweight"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_refusal_has_no_rating_and_is_not_defaulted():
|
||||
assert extract_rating(REFUSAL) is None
|
||||
assert parse_rating(REFUSAL) == RATING_REVIEW
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_scale_quoted_in_a_prompt_does_not_become_the_rating():
|
||||
"""A free-text answer that echoes the rating scale was read as the first
|
||||
tier listed in it."""
|
||||
text = ("**Rating Scale**: Buy, Overweight, Hold, Underweight, Sell.\n\n"
|
||||
"**Rating**: Sell\n\nExit the position.")
|
||||
assert extract_rating(text) == "Sell"
|
||||
|
||||
|
||||
# --- the readers agree ------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_memory_log_records_review_rather_than_a_tradeable_hold(tmp_path):
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
|
||||
log = TradingMemoryLog({"memory_log_path": str(tmp_path / "m.md")})
|
||||
log.store_decision("NVDA", "2026-01-05", REFUSAL)
|
||||
|
||||
entry = log.load_entries()[0]
|
||||
assert entry["rating"] == RATING_REVIEW
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_signal_and_the_log_agree_on_the_same_decision(tmp_path):
|
||||
from tradingagents.agents.rating import parse_rating
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
|
||||
log = TradingMemoryLog({"memory_log_path": str(tmp_path / "m.md")})
|
||||
for text in (INVERTED, REFUSAL, "**Rating**: Buy\n\nAccumulate."):
|
||||
log.store_decision("NVDA", f"2026-01-0{len(log.load_entries()) + 1}", text)
|
||||
|
||||
signals = [parse_rating(text)
|
||||
for text in (INVERTED, REFUSAL, "**Rating**: Buy\n\nAccumulate.")]
|
||||
assert [e["rating"] for e in log.load_entries()] == signals
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_unscored_decision_is_left_out_of_the_backtest_figures(tmp_path):
|
||||
"""REVIEW has no direction, so it cannot count for or against the system."""
|
||||
from tradingagents.backtest import summarize
|
||||
from tradingagents.decision_log import TradingMemoryLog
|
||||
|
||||
log = TradingMemoryLog({"memory_log_path": str(tmp_path / "m.md")})
|
||||
log.store_decision("NVDA", "2026-01-05", "**Rating**: Buy\n\nx")
|
||||
log.update_with_outcome("NVDA", "2026-01-05", 0.1, 0.04, 5, "note", "2026-02-01")
|
||||
log.store_decision("AAPL", "2026-01-05", REFUSAL)
|
||||
log.update_with_outcome("AAPL", "2026-01-05", 0.1, 0.04, 5, "note", "2026-02-01")
|
||||
|
||||
summary = summarize(tmp_path / "m.md")
|
||||
assert set(summary.by_rating) == {"Buy"}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_cli_says_when_a_run_produced_no_usable_rating(monkeypatch, tmp_path, capsys):
|
||||
"""The CLI is the primary entry point; an unreadable decision must be
|
||||
visible there, not only in the log."""
|
||||
import cli.main as m
|
||||
from cli.models import AnalystType
|
||||
|
||||
printed = []
|
||||
|
||||
class _Graph:
|
||||
graph = propagator = None
|
||||
|
||||
def create_run_state(self, *a, **k):
|
||||
return {"messages": []}
|
||||
|
||||
def record_decision(self, *a, **k):
|
||||
pass
|
||||
|
||||
def process_signal(self, text):
|
||||
from tradingagents.agents.rating import parse_rating
|
||||
return parse_rating(text)
|
||||
|
||||
def get_graph_args(self, callbacks=None):
|
||||
return {}
|
||||
|
||||
def begin_checkpoint(self, *a, **k):
|
||||
return None
|
||||
|
||||
def checkpoint_input(self, state):
|
||||
return state
|
||||
|
||||
def clear_checkpoint_on_success(self, *a, **k):
|
||||
pass
|
||||
|
||||
def end_checkpoint(self):
|
||||
pass
|
||||
|
||||
def stream(self, *a, **k):
|
||||
yield {"messages": [], "final_trade_decision": REFUSAL}
|
||||
|
||||
fake = _Graph()
|
||||
fake.graph = fake
|
||||
fake.propagator = fake
|
||||
monkeypatch.setattr(cli_run, "TradingAgentsGraph", lambda *a, **k: fake)
|
||||
monkeypatch.setattr(cli_run, "create_layout", lambda: None)
|
||||
monkeypatch.setattr(cli_run, "update_display", lambda *a, **k: None)
|
||||
monkeypatch.setattr(cli_run, "Live", type("L", (), {"__init__": lambda s, *a, **k: None,
|
||||
"__enter__": lambda s: s,
|
||||
"__exit__": lambda s, *a: False}))
|
||||
monkeypatch.setattr(m.console, "print", lambda *a, **k: printed.append(" ".join(str(x) for x in a)))
|
||||
monkeypatch.setattr(cli_run, "display_complete_report", lambda *a, **k: None)
|
||||
monkeypatch.setattr(m.typer, "prompt", lambda *a, **k: "N")
|
||||
monkeypatch.setattr(cli_run, "get_user_selections", lambda: {
|
||||
"ticker": "NVDA", "analysis_date": "2026-01-10",
|
||||
"analysts": [AnalystType.MARKET], "asset_type": "stock",
|
||||
})
|
||||
monkeypatch.setattr(cli_run, "_build_run_config", lambda s, c: {
|
||||
"data_cache_dir": str(tmp_path / "c"), "results_dir": str(tmp_path / "r")})
|
||||
|
||||
cli_run.run_analysis()
|
||||
|
||||
assert any("review" in line.lower() for line in printed), printed[-5:]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("module, factory, must_name", [
|
||||
("tradingagents.agents.managers.portfolio_manager", "create_portfolio_manager", "Rating"),
|
||||
("tradingagents.agents.managers.research_manager", "create_research_manager", "Recommendation"),
|
||||
("tradingagents.agents.trader.trader", "create_trader", "Action"),
|
||||
])
|
||||
def test_a_decision_prompt_states_the_shape_of_its_answer(module, factory, must_name):
|
||||
"""The field descriptions live in the schema, which a provider without
|
||||
structured output never sees. Without the format in the prompt body, the
|
||||
fallback answer is prose nobody can read a rating from."""
|
||||
import importlib
|
||||
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
mod = importlib.import_module(module)
|
||||
seen = []
|
||||
|
||||
class _LLM:
|
||||
def invoke(self, prompt, *a, **k):
|
||||
seen.append(prompt if isinstance(prompt, str) else str(prompt))
|
||||
return AIMessage("**Rating**: Hold\n\nnothing to do")
|
||||
|
||||
def with_structured_output(self, *a, **k):
|
||||
raise NotImplementedError # force the free-text path
|
||||
|
||||
state = {
|
||||
"company_of_interest": "NVDA", "trade_date": "2026-08-14", "asset_type": "stock",
|
||||
"instrument_context": "", "market_report": "M", "sentiment_report": "S",
|
||||
"news_report": "N", "fundamentals_report": "F", "investment_plan": "P",
|
||||
"trader_investment_plan": "T", "past_context": "", "portfolio_context": "",
|
||||
"investment_debate_state": {"bull_history": "b", "bear_history": "r", "history": "h",
|
||||
"current_response": "", "judge_decision": "", "count": 2},
|
||||
"risk_debate_state": {"history": "h", "latest_speaker": "", "count": 3,
|
||||
"aggressive_history": "", "conservative_history": "", "neutral_history": "",
|
||||
"current_aggressive_response": "", "current_conservative_response": "",
|
||||
"current_neutral_response": "", "judge_decision": ""},
|
||||
}
|
||||
getattr(mod, factory)(_LLM())(state)
|
||||
|
||||
prompt = " ".join(seen)
|
||||
assert "## Output" in prompt, "no output-format section in the prompt"
|
||||
section = prompt.split("## Output", 1)[1]
|
||||
assert f"**{must_name}**" in section, section[:300]
|
||||
@@ -0,0 +1,313 @@
|
||||
"""Tests for the Reddit RSS fetcher: one combined request, its 429 backoff, and
|
||||
chunked-transfer error handling (#1024)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import http.client
|
||||
from unittest.mock import patch
|
||||
from urllib.error import HTTPError
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.vendors import reddit
|
||||
|
||||
_SAMPLE_ATOM = """<?xml version="1.0" encoding="UTF-8"?>
|
||||
<feed xmlns="http://www.w3.org/2005/Atom">
|
||||
<entry>
|
||||
<title>NVDA earnings beat, stock pops</title>
|
||||
<published>2026-05-20T14:30:00+00:00</published>
|
||||
<content type="html"><!-- SC_OFF --><div class="md"><p>Great <b>quarter</b> for NVDA&#39;s datacenter unit.</p></div><!-- SC_ON --></content>
|
||||
</entry>
|
||||
<entry>
|
||||
<title>Is NVDA overvalued?</title>
|
||||
<published>2026-05-19T09:00:00Z</published>
|
||||
<content type="html"><p>Forward P/E discussion</p></content>
|
||||
</entry>
|
||||
</feed>
|
||||
"""
|
||||
|
||||
|
||||
def _resp(read_fn):
|
||||
"""A minimal context-manager response whose read() runs ``read_fn``."""
|
||||
class _Resp:
|
||||
def __enter__(self_inner):
|
||||
return self_inner
|
||||
|
||||
def __exit__(self_inner, *a):
|
||||
return False
|
||||
|
||||
def read(self_inner, size=-1):
|
||||
data = read_fn()
|
||||
return data if size is None or size < 0 else data[:size]
|
||||
return _Resp()
|
||||
|
||||
|
||||
def _atom_resp():
|
||||
return _resp(lambda: _SAMPLE_ATOM.encode("utf-8"))
|
||||
|
||||
|
||||
def _raise(exc):
|
||||
def _r():
|
||||
raise exc
|
||||
return _resp(_r)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestIsoToTimestamp:
|
||||
def test_parses_offset_and_z(self):
|
||||
assert reddit._iso_to_timestamp("2026-05-20T14:30:00+00:00") > 0
|
||||
assert reddit._iso_to_timestamp("2026-05-19T09:00:00Z") > 0
|
||||
|
||||
def test_none_and_garbage_return_none(self):
|
||||
assert reddit._iso_to_timestamp(None) is None
|
||||
assert reddit._iso_to_timestamp("not-a-date") is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStripHtml:
|
||||
def test_extracts_between_sc_markers_and_unescapes(self):
|
||||
raw = "<!-- SC_OFF --><div class=\"md\"><p>Great <b>quarter</b> & more</p></div><!-- SC_ON -->"
|
||||
assert reddit._strip_html(raw) == "Great quarter & more"
|
||||
|
||||
def test_empty(self):
|
||||
assert reddit._strip_html("") == ""
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRssParsing:
|
||||
def test_parses_atom_entries(self):
|
||||
with patch.object(reddit, "urlopen", return_value=_atom_resp()):
|
||||
posts = reddit._fetch_subreddit_rss("NVDA", "stocks", limit=5, timeout=5.0)
|
||||
assert len(posts) == 2
|
||||
assert posts[0]["title"] == "NVDA earnings beat, stock pops"
|
||||
assert posts[0]["created_utc"] > 0
|
||||
assert "datacenter unit" in posts[0]["selftext"]
|
||||
assert posts[0]["subreddit"] == "stocks"
|
||||
|
||||
def test_malformed_xml_reports_unavailable(self):
|
||||
with patch.object(reddit, "urlopen", return_value=_resp(lambda: b"<<not xml>>")):
|
||||
assert reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0) is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRss429Backoff:
|
||||
def test_429_then_success_retries_once(self):
|
||||
err = HTTPError("url", 429, "Too Many Requests", {}, None)
|
||||
with patch.object(reddit, "urlopen", side_effect=[err, _atom_resp()]) as op, \
|
||||
patch.object(reddit.time, "sleep") as slept:
|
||||
posts = reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0)
|
||||
assert op.call_count == 2 # original + exactly one retry
|
||||
slept.assert_called_once() # backed off before retrying
|
||||
assert len(posts) == 2
|
||||
|
||||
def test_429_twice_gives_up_after_one_retry(self):
|
||||
err = HTTPError("url", 429, "Too Many Requests", {}, None)
|
||||
with patch.object(reddit, "urlopen", side_effect=[err, err]) as op, \
|
||||
patch.object(reddit.time, "sleep"):
|
||||
posts = reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0)
|
||||
assert op.call_count == 2 # one retry, then gives up cleanly
|
||||
assert posts is None
|
||||
|
||||
def test_retry_after_header_is_honoured(self):
|
||||
err = HTTPError("url", 429, "Too Many Requests", {"Retry-After": "12"}, None)
|
||||
with patch.object(reddit, "urlopen", side_effect=[err, _atom_resp()]), \
|
||||
patch.object(reddit.time, "sleep") as slept:
|
||||
reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0)
|
||||
slept.assert_called_once_with(12.0)
|
||||
|
||||
def test_retry_after_zero_is_honoured_not_treated_as_absent(self):
|
||||
# A valid "Retry-After: 0" means retry at once; it must not fall through
|
||||
# to the fallback wait (the earlier `or 5.0` bug turned 0 into 5s).
|
||||
err = HTTPError("url", 429, "Too Many Requests", {"Retry-After": "0"}, None)
|
||||
with patch.object(reddit, "urlopen", side_effect=[err, _atom_resp()]), \
|
||||
patch.object(reddit.time, "sleep") as slept:
|
||||
reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0)
|
||||
slept.assert_called_once_with(0.0)
|
||||
|
||||
def test_headerless_429_fallback_is_jittered(self):
|
||||
# No Retry-After -> our own ~5s fallback, jittered so concurrent runs
|
||||
# don't retry in lockstep (kept within a tight band).
|
||||
err = HTTPError("url", 429, "Too Many Requests", {}, None)
|
||||
with patch.object(reddit, "urlopen", side_effect=[err, _atom_resp()]), \
|
||||
patch.object(reddit.time, "sleep") as slept:
|
||||
reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0)
|
||||
slept.assert_called_once()
|
||||
(wait,), _ = slept.call_args
|
||||
assert 48.0 <= wait <= 72.0 # 60s +/-20% jitter
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestChunkedTransferErrorsHandled:
|
||||
"""IncompleteRead/RemoteDisconnected come from http.client and are NOT
|
||||
OSErrors, so they were previously uncaught and crashed the pipeline (#1024)."""
|
||||
|
||||
def test_rss_incomplete_read_reports_unavailable(self):
|
||||
with patch.object(reddit, "urlopen", return_value=_raise(http.client.IncompleteRead(b""))):
|
||||
assert reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0) is None
|
||||
|
||||
def test_oversized_rss_feed_is_refused_not_parsed(self):
|
||||
# A hostile/misbehaving endpoint streaming an unbounded body must not be
|
||||
# read into memory before parsing; overflow degrades to an empty feed.
|
||||
big = _resp(lambda: b"x" * 100)
|
||||
with patch.object(reddit, "_MAX_FEED_BYTES", 10), \
|
||||
patch.object(reddit, "urlopen", return_value=big):
|
||||
assert reddit._fetch_subreddit_rss("NVDA", "stocks", 5, 5.0) is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFormatterHandlesRssPosts:
|
||||
def test_rss_posts_omit_fake_counts_and_note_source(self):
|
||||
rss_posts = [{
|
||||
"title": "NVDA pops", "score": None, "num_comments": None,
|
||||
"created_utc": reddit._iso_to_timestamp("2026-05-20T14:30:00Z"),
|
||||
"selftext": "great quarter", "source": "rss",
|
||||
}]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=rss_posts):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("stocks",))
|
||||
assert "↑" not in out # RSS has no scores; none are invented
|
||||
assert "NVDA pops" in out
|
||||
assert "great quarter" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCryptoSearchTerm:
|
||||
"""A crypto pair (BTC-USD) barely matches Reddit text; search the base (#1113)."""
|
||||
|
||||
def _captured_ticker(self, ticker):
|
||||
seen = {}
|
||||
|
||||
def fake_fetch(t, subs, limit, timeout, **kwargs):
|
||||
seen["ticker"] = t
|
||||
return []
|
||||
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", side_effect=fake_fetch):
|
||||
reddit.fetch_reddit_posts(ticker, subreddits=("stocks",))
|
||||
return seen["ticker"]
|
||||
|
||||
def test_crypto_pair_searches_base(self):
|
||||
assert self._captured_ticker("BTC-USD") == "BTC"
|
||||
|
||||
def test_equity_passes_through(self):
|
||||
assert self._captured_ticker("NVDA") == "NVDA"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOneRequestForAllSubreddits:
|
||||
"""Reddit's anonymous RSS allows about one request per minute per IP, so a
|
||||
request per subreddit spent a back-off on nearly every run. One combined
|
||||
feed (``r/a+b+c``) carries each entry's subreddit, so nothing is lost."""
|
||||
|
||||
def _post(self, sub, title="NVDA pops"):
|
||||
return {"title": title, "score": None, "num_comments": None,
|
||||
"created_utc": reddit._iso_to_timestamp("2026-05-20T14:30:00Z"),
|
||||
"selftext": "", "source": "rss", "subreddit": sub}
|
||||
|
||||
def test_all_subreddits_share_one_request(self):
|
||||
calls = []
|
||||
|
||||
def record(t, subs, limit, timeout):
|
||||
calls.append((subs, limit))
|
||||
return []
|
||||
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", side_effect=record):
|
||||
reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b", "c"), limit_per_sub=5)
|
||||
# One full page, so a busy subreddit cannot crowd the others out.
|
||||
assert calls == [("a+b+c", reddit._FEED_PAGE)]
|
||||
|
||||
def test_posts_are_grouped_back_by_subreddit(self):
|
||||
posts = [self._post("b", "FROM B"), self._post("a", "FROM A")]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=posts):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert out.index("r/a") < out.index("FROM A") < out.index("r/b") < out.index("FROM B")
|
||||
|
||||
def test_failed_request_is_unavailable_not_silence(self):
|
||||
# #1295: a throttled fetch must not read as "no posts found".
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=None):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert "Reddit unavailable" in out
|
||||
assert "no Reddit posts found" not in out
|
||||
|
||||
def test_genuine_empty_still_reports_no_posts(self):
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=[]):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert "no Reddit posts found" in out
|
||||
assert "unavailable" not in out
|
||||
|
||||
def test_subreddit_with_no_posts_is_listed_when_others_have_some(self):
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=[self._post("a")]):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert "r/b: <no posts found" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_posts_from_an_unrequested_or_unnamed_subreddit_are_not_dropped():
|
||||
posts = [
|
||||
{"title": "ELSEWHERE", "created_utc": None, "selftext": "", "subreddit": "options"},
|
||||
{"title": "NO LABEL", "created_utc": None, "selftext": "", "subreddit": ""},
|
||||
]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=posts):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert "ELSEWHERE" in out and "r/options" in out
|
||||
assert "NO LABEL" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_each_subreddit_keeps_its_own_quota():
|
||||
busy = [{"title": f"A{i}", "created_utc": None, "selftext": "", "subreddit": "a"} for i in range(9)]
|
||||
quiet = [{"title": "B0", "created_utc": None, "selftext": "", "subreddit": "b"}]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=busy + quiet):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"), limit_per_sub=3)
|
||||
assert "A0" in out and "A2" in out and "A3" not in out # capped per subreddit
|
||||
assert "B0" in out # not crowded out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_empty_subreddit_on_a_full_page_is_not_called_empty():
|
||||
# A full page may have cut a quieter subreddit's posts off, so its absence
|
||||
# from the page is not evidence of no posts.
|
||||
full = [{"title": f"A{i}", "created_utc": None, "selftext": "", "subreddit": "a"}
|
||||
for i in range(reddit._FEED_PAGE)]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=full):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert "r/b: <no posts found" not in out
|
||||
assert f"newest {reddit._FEED_PAGE}" in out
|
||||
|
||||
|
||||
def _screen_out(*dropped):
|
||||
"""A screen that drops posts whose text starts with one of ``dropped``."""
|
||||
def screen(texts):
|
||||
return [not t.startswith(dropped) for t in texts], "Screened: note"
|
||||
return screen
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_screened_out_posts_free_their_subreddit_slots():
|
||||
posts = [{"title": t, "created_utc": None, "selftext": "", "subreddit": "a"}
|
||||
for t in ("SPAM1", "SPAM2", "A1", "A2")]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=posts):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a",), limit_per_sub=2,
|
||||
screen=_screen_out("SPAM"))
|
||||
assert out.startswith("Screened: note")
|
||||
assert "A1" in out and "A2" in out and "SPAM" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_subreddit_emptied_by_screening_is_not_called_empty():
|
||||
posts = [{"title": "SPAM", "created_utc": None, "selftext": "", "subreddit": "b"},
|
||||
{"title": "A1", "created_utc": None, "selftext": "", "subreddit": "a"}]
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=posts):
|
||||
out = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"), screen=_screen_out("SPAM"))
|
||||
assert "r/b: <no posts about NVDA after screening>" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_unavailable_screen_keeps_every_post_and_says_so():
|
||||
posts = [{"title": "A1", "created_utc": None, "selftext": "", "subreddit": "a"}]
|
||||
|
||||
def unavailable(texts):
|
||||
return [True] * len(texts), "<Jev screening unavailable (HTTP 529); posts are unscreened>"
|
||||
|
||||
with patch.object(reddit, "_fetch_subreddit_rss", return_value=posts):
|
||||
screened = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"), screen=unavailable)
|
||||
plain = reddit.fetch_reddit_posts("NVDA", subreddits=("a", "b"))
|
||||
assert screened == "<Jev screening unavailable (HTTP 529); posts are unscreened>\n\n" + plain
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Report parity: the shared writer produces the report tree for the CLI and the
|
||||
programmatic API alike (#1037)."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
from tradingagents.reporting import write_report_tree
|
||||
|
||||
|
||||
def _state():
|
||||
return {
|
||||
"market_report": "MKT",
|
||||
"news_report": "NEWS",
|
||||
"investment_debate_state": {"judge_decision": "RM PLAN"},
|
||||
"trader_investment_plan": "TRADE",
|
||||
"risk_debate_state": {"judge_decision": "PM DECISION"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_write_report_tree_creates_files(tmp_path):
|
||||
out = write_report_tree(_state(), "AAPL", tmp_path)
|
||||
assert out.name == "complete_report.md"
|
||||
assert (tmp_path / "1_analysts" / "market.md").read_text() == "MKT"
|
||||
assert (tmp_path / "1_analysts" / "news.md").read_text() == "NEWS"
|
||||
assert (tmp_path / "2_research" / "manager.md").read_text() == "RM PLAN"
|
||||
assert (tmp_path / "3_trading" / "trader.md").read_text() == "TRADE"
|
||||
assert (tmp_path / "5_portfolio" / "decision.md").read_text() == "PM DECISION"
|
||||
complete = out.read_text()
|
||||
assert "Trading Analysis Report: AAPL" in complete
|
||||
assert "MKT" in complete and "PM DECISION" in complete
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_save_reports_explicit_path(tmp_path):
|
||||
# Unbound: with an explicit save_path, the method doesn't touch self/config.
|
||||
out = TradingAgentsGraph.save_reports(None, _state(), "AAPL", save_path=tmp_path)
|
||||
assert (tmp_path / "complete_report.md").exists()
|
||||
assert out == tmp_path / "complete_report.md"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_save_reports_defaults_under_results_dir(tmp_path):
|
||||
mock_self = SimpleNamespace(config={"results_dir": str(tmp_path)})
|
||||
out = TradingAgentsGraph.save_reports(mock_self, _state(), "AAPL")
|
||||
assert out.exists()
|
||||
assert out.parent.parent.name == "reports" # results_dir/reports/AAPL_<stamp>/...
|
||||
assert out.parent.name.startswith("AAPL_")
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Shared-router / path_map completeness (#1088).
|
||||
|
||||
Both `should_continue_risk_analysis` (three risk edges) and
|
||||
`should_continue_debate` (two research-debate edges) are single routers whose
|
||||
return set is larger than any one edge previously mapped. Each edge now shares a
|
||||
complete path map (`RISK_ANALYSIS_PATH_MAP` / `DEBATE_PATH_MAP`), so a
|
||||
fall-through return can never hit a missing entry -- which would crash LangGraph
|
||||
mid-run on prompt/i18n/refactor drift in the speaker labels.
|
||||
"""
|
||||
import pytest
|
||||
|
||||
from tradingagents.graph.conditional_logic import ConditionalLogic
|
||||
from tradingagents.graph.setup import DEBATE_PATH_MAP, RISK_ANALYSIS_PATH_MAP
|
||||
|
||||
|
||||
def _state(latest_speaker, count=0):
|
||||
return {"risk_debate_state": {"latest_speaker": latest_speaker, "count": count}}
|
||||
|
||||
|
||||
def _debate_state(current_response, count=0):
|
||||
return {"investment_debate_state": {"current_response": current_response, "count": count}}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("latest_speaker", [
|
||||
"Aggressive", "Aggressive Analyst",
|
||||
"Conservative", "Conservative Analyst",
|
||||
"Neutral", "Neutral Analyst",
|
||||
"", # drift: empty label
|
||||
"Aggressive Risk Analyst", # drift: node renamed
|
||||
"Agresivo", # drift: i18n / translated label
|
||||
])
|
||||
def test_router_return_always_routable(latest_speaker):
|
||||
logic = ConditionalLogic(max_risk_discuss_rounds=1)
|
||||
target = logic.should_continue_risk_analysis(_state(latest_speaker))
|
||||
assert target in RISK_ANALYSIS_PATH_MAP
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_router_terminates_at_round_limit():
|
||||
logic = ConditionalLogic(max_risk_discuss_rounds=1)
|
||||
# count >= 3 * rounds routes to the Portfolio Manager (debate ends)
|
||||
assert logic.should_continue_risk_analysis(_state("Neutral", count=3)) == "Portfolio Manager"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_path_map_covers_full_router_range():
|
||||
logic = ConditionalLogic(max_risk_discuss_rounds=1)
|
||||
returns = {
|
||||
logic.should_continue_risk_analysis(_state(s, c))
|
||||
for s in ("Aggressive", "Conservative", "Neutral", "drift")
|
||||
for c in (0, 99)
|
||||
}
|
||||
# Every value the router can emit is a key in the shared map...
|
||||
assert returns <= set(RISK_ANALYSIS_PATH_MAP)
|
||||
# ...and the terminal target is reachable.
|
||||
assert "Portfolio Manager" in returns
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("current_response", [
|
||||
"Bull", "Bull Researcher", "Bear", "Bear Researcher",
|
||||
"", # drift: empty label
|
||||
"Optimista", # drift: i18n / translated label
|
||||
])
|
||||
def test_debate_router_return_always_routable(current_response):
|
||||
logic = ConditionalLogic(max_debate_rounds=1)
|
||||
target = logic.should_continue_debate(_debate_state(current_response))
|
||||
assert target in DEBATE_PATH_MAP
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_debate_path_map_covers_full_router_range():
|
||||
logic = ConditionalLogic(max_debate_rounds=1)
|
||||
returns = {
|
||||
logic.should_continue_debate(_debate_state(s, c))
|
||||
for s in ("Bull", "Bear", "drift")
|
||||
for c in (0, 99)
|
||||
}
|
||||
assert returns <= set(DEBATE_PATH_MAP)
|
||||
assert "Research Manager" in returns # terminal reachable
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Tests for the ticker path-component validator that blocks directory traversal."""
|
||||
|
||||
import os
|
||||
import unittest
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.symbols import safe_ticker_component
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestSafeTickerComponent(unittest.TestCase):
|
||||
def test_accepts_common_ticker_formats(self):
|
||||
for ticker in ("AAPL", "BRK-B", "BRK.A", "0700.HK", "7203.T", "BHP.AX", "^GSPC"):
|
||||
self.assertEqual(safe_ticker_component(ticker), ticker)
|
||||
|
||||
def test_accepts_futures_and_forex_formats(self):
|
||||
# Futures use '=' (GC=F gold, CL=F crude), forex/CFD symbols use '+'.
|
||||
for ticker in ("GC=F", "CL=F", "ES=F", "XAUUSD+", "EURUSD+"):
|
||||
self.assertEqual(safe_ticker_component(ticker), ticker)
|
||||
|
||||
def test_rejects_path_separators(self):
|
||||
for bad in (".", "..", "../etc", "a/b", "a\\b", "/abs", "..\\..\\x"):
|
||||
with self.assertRaises(ValueError):
|
||||
safe_ticker_component(bad)
|
||||
|
||||
def test_rejects_null_byte_and_whitespace(self):
|
||||
for bad in ("AAP L", "AAPL\x00", "AAPL\n", "\tAAPL"):
|
||||
with self.assertRaises(ValueError):
|
||||
safe_ticker_component(bad)
|
||||
|
||||
def test_rejects_empty_or_non_string(self):
|
||||
for bad in ("", None, 123, b"AAPL"):
|
||||
with self.assertRaises(ValueError):
|
||||
safe_ticker_component(bad)
|
||||
|
||||
def test_rejects_overlong_input(self):
|
||||
with self.assertRaises(ValueError):
|
||||
safe_ticker_component("A" * 33)
|
||||
|
||||
def test_rejects_dot_only_values(self):
|
||||
# '.' and '..' pass the regex but traverse when used as a path
|
||||
# component (e.g. ``Path(results_dir) / ticker / "logs"``).
|
||||
for bad in (".", "..", "...", "...."):
|
||||
with self.assertRaises(ValueError):
|
||||
safe_ticker_component(bad)
|
||||
|
||||
def test_traversal_string_does_not_escape_join(self):
|
||||
"""Sanity: sanitized values stay within base when joined."""
|
||||
base = os.path.realpath("/tmp/cache")
|
||||
ticker = safe_ticker_component("AAPL")
|
||||
joined = os.path.realpath(os.path.join(base, f"{ticker}.csv"))
|
||||
self.assertTrue(joined.startswith(base + os.sep))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,295 @@
|
||||
"""SEC EDGAR fundamentals: statements as they were filed, not as they read today.
|
||||
|
||||
Every other fundamentals vendor serves the current value of a past period and
|
||||
cuts on the fiscal period end, so a run sees figures the company had not yet
|
||||
filed, and later restatements replace what was actually published. EDGAR carries
|
||||
the filing date of every fact, so a run can be limited to what was on file by its
|
||||
own date.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.errors import NoMarketDataError
|
||||
from tradingagents.dataflows.vendors import sec_edgar
|
||||
|
||||
_REAL_FETCH = sec_edgar._fetch_json
|
||||
|
||||
TICKER_MAP = {"0": {"cik_str": 320193, "ticker": "AAPL", "title": "Apple Inc."}}
|
||||
|
||||
|
||||
def _fact(end, val, filed, form="10-K", fp="FY", start=None):
|
||||
fact = {"end": end, "val": val, "filed": filed, "form": form, "fy": int(end[:4]), "fp": fp}
|
||||
if start:
|
||||
fact["start"] = start
|
||||
return fact
|
||||
|
||||
|
||||
FACTS = {
|
||||
"cik": 320193,
|
||||
"entityName": "Apple Inc.",
|
||||
"facts": {"us-gaap": {
|
||||
"Assets": {"units": {"USD": [
|
||||
_fact("2008-09-27", 39_572_000_000, "2008-11-05"),
|
||||
_fact("2008-09-27", 36_171_000_000, "2010-01-25", form="10-K/A"),
|
||||
_fact("2022-03-26", 350_662_000_000, "2022-04-29", form="10-Q", fp="Q2"),
|
||||
_fact("2024-09-28", 364_980_000_000, "2024-11-01"),
|
||||
]}},
|
||||
"Liabilities": {"units": {"USD": [_fact("2024-09-28", 308_030_000_000, "2024-11-01")]}},
|
||||
"EarningsPerShareDiluted": {"units": {"USD/shares": [
|
||||
_fact("2024-09-28", 6.08, "2024-11-01", start="2023-09-30"),
|
||||
]}},
|
||||
"RevenueFromContractWithCustomerExcludingAssessedTax": {"units": {"USD": [
|
||||
# One filing reports the quarter and the year to date under one end date.
|
||||
_fact("2025-12-31", 81_300_000_000, "2026-01-29", form="10-Q", fp="Q2", start="2025-10-01"),
|
||||
_fact("2025-12-31", 158_900_000_000, "2026-01-29", form="10-Q", fp="Q2", start="2025-07-01"),
|
||||
_fact("2024-09-28", 391_035_000_000, "2024-11-01", start="2023-09-30"),
|
||||
]}},
|
||||
}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_network_or_cache(tmp_path, monkeypatch):
|
||||
monkeypatch.setattr(sec_edgar, "get_config", lambda: {"data_cache_dir": str(tmp_path)})
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json", lambda url: TICKER_MAP if "company_tickers" in url else FACTS)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_us_filer_resolves_to_its_cik():
|
||||
assert sec_edgar.cik_for("AAPL") == "0000320193"
|
||||
assert sec_edgar.cik_for("aapl") == "0000320193"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_non_filer_is_reported_as_such_not_as_missing_data():
|
||||
with pytest.raises(NoMarketDataError, match="not a US SEC filer"):
|
||||
sec_edgar.get_balance_sheet("0700.HK", "annual", "2026-01-01")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_restated_figure_reads_as_it_did_at_the_time():
|
||||
"""The value published then, not the correction filed later."""
|
||||
as_filed = sec_edgar.get_balance_sheet("AAPL", "annual", "2009-06-30")
|
||||
restated = sec_edgar.get_balance_sheet("AAPL", "annual", "2011-01-01")
|
||||
assert "39572" in as_filed and "36171" not in as_filed
|
||||
assert "36171" in restated
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_period_that_ended_but_was_not_filed_yet_is_not_served():
|
||||
"""The fiscal year ended 2024-09-28; it reached the public on 2024-11-01."""
|
||||
before = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-10-15")
|
||||
after = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15")
|
||||
assert "2024-09-28" not in before
|
||||
assert "2024-09-28" in after and "364980" in after
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_quarter_is_not_confused_with_the_year_to_date():
|
||||
"""One filing carries both spans under the same end date (#MSFT-shaped)."""
|
||||
out = sec_edgar.get_income_statement("AAPL", "quarterly", "2026-06-01")
|
||||
assert "81300" in out
|
||||
assert "158900" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_line_the_filer_does_not_tag_is_named_unavailable():
|
||||
out = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15")
|
||||
assert "Stockholders Equity" in out and "unavailable" in out
|
||||
assert "364980" in out # the rest of the statement still returns
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_report_states_the_vintage_rule():
|
||||
out = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15")
|
||||
assert "filed on or before 2024-11-15" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_filer_with_no_usable_facts_reads_differently_from_a_non_filer(monkeypatch):
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json",
|
||||
lambda url: TICKER_MAP if "company_tickers" in url else {"facts": {}})
|
||||
with pytest.raises(NoMarketDataError, match="no us-gaap facts"):
|
||||
sec_edgar.get_balance_sheet("AAPL", "annual", "2026-01-01")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_company_facts_are_fetched_once_per_company_not_once_per_date(tmp_path, monkeypatch):
|
||||
"""A sweep asks for many dates; the filing history is the same file."""
|
||||
calls = []
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json",
|
||||
lambda url: calls.append(url) or (TICKER_MAP if "company_tickers" in url else FACTS))
|
||||
for date in ("2024-11-15", "2025-01-15", "2025-06-15"):
|
||||
sec_edgar.get_balance_sheet("AAPL", "annual", date)
|
||||
assert len([u for u in calls if "companyfacts" in u]) == 1
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_it_works_unconfigured_and_takes_the_caller_s_own_contact(monkeypatch):
|
||||
"""SEC returns 403 for a User-Agent with no contact address, so the default
|
||||
carries one; a caller who sets their own replaces it."""
|
||||
monkeypatch.delenv("SEC_EDGAR_USER_AGENT", raising=False)
|
||||
assert "@" in sec_edgar._user_agent()
|
||||
|
||||
monkeypatch.setenv("SEC_EDGAR_USER_AGENT", "MyDesk research@example.com")
|
||||
assert sec_edgar._user_agent() == "MyDesk research@example.com"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_throttle_lets_the_next_vendor_try(monkeypatch):
|
||||
"""SEC throttles by refusing the request; the router then tries yfinance."""
|
||||
import requests
|
||||
|
||||
from tradingagents.dataflows.errors import VendorRateLimitError
|
||||
|
||||
def _throttled(*a, **k):
|
||||
raise requests.HTTPError(response=mock.Mock(status_code=429))
|
||||
|
||||
monkeypatch.setattr(sec_edgar.requests, "get", _throttled)
|
||||
with pytest.raises(VendorRateLimitError):
|
||||
_REAL_FETCH("https://data.sec.gov/api/xbrl/companyfacts/CIK0000320193.json")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_values_do_not_break_the_columns():
|
||||
"""Figures run to the billions; a thousands separator would split the field."""
|
||||
out = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15")
|
||||
body = [row for row in out.splitlines() if row.startswith("Total Assets")][0]
|
||||
assert body.count(",") == out.splitlines()[3].count(",")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_per_share_figure_keeps_its_own_unit():
|
||||
"""Statements are reported in millions, but EPS is dollars per share: scaling
|
||||
it the same way prints a real figure as zero."""
|
||||
out = sec_edgar.get_income_statement("AAPL", "annual", "2024-11-15")
|
||||
row = [r for r in out.splitlines() if r.startswith("Diluted EPS")][0]
|
||||
assert "6.08" in row
|
||||
assert "USD/shares" in row or "per share" in row
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_every_row_has_one_cell_per_period():
|
||||
"""An untagged line still has to line up with the columns, or the table is
|
||||
misread by position."""
|
||||
out = sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15")
|
||||
table = [r for r in out.splitlines() if r and not r.startswith("#")]
|
||||
widths = {row.count(",") for row in table}
|
||||
assert len(widths) == 1, table
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_server_error_lets_the_next_vendor_try(monkeypatch):
|
||||
import requests
|
||||
|
||||
from tradingagents.dataflows.errors import VendorError
|
||||
|
||||
def _server_error(*a, **k):
|
||||
raise requests.HTTPError(response=mock.Mock(status_code=503))
|
||||
|
||||
monkeypatch.setattr(sec_edgar.requests, "get", _server_error)
|
||||
with pytest.raises(VendorError): # not a bare HTTPError
|
||||
_REAL_FETCH("https://data.sec.gov/api/xbrl/companyfacts/CIK0000320193.json")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_older_periods_fall_back_to_the_tag_the_filer_used_then():
|
||||
"""Filers renamed lines when the revenue standard changed, so one tag covers
|
||||
only recent years. Each period takes one tag, never a sum of two."""
|
||||
facts = {"Revenues": {"units": {"USD": [_fact("2015-09-26", 233_715_000_000, "2015-10-28",
|
||||
start="2014-09-28")]}},
|
||||
"RevenueFromContractWithCustomerExcludingAssessedTax": {"units": {"USD": [
|
||||
_fact("2024-09-28", 391_035_000_000, "2024-11-01", start="2023-09-30")]}}}
|
||||
values, unit = sec_edgar._as_of(facts, ("RevenueFromContractWithCustomerExcludingAssessedTax",
|
||||
"Revenues"), "2026-01-01", (300, 400))
|
||||
assert values == {"2015-09-26": 233_715_000_000, "2024-09-28": 391_035_000_000}
|
||||
assert unit == "USD"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_period_reported_under_two_tags_takes_the_preferred_one_not_both():
|
||||
facts = {"Revenues": {"units": {"USD": [_fact("2024-09-28", 111, "2024-11-01", start="2023-09-30")]}},
|
||||
"RevenueFromContractWithCustomerExcludingAssessedTax": {"units": {"USD": [
|
||||
_fact("2024-09-28", 999, "2024-11-01", start="2023-09-30")]}}}
|
||||
values, _ = sec_edgar._as_of(facts, ("RevenueFromContractWithCustomerExcludingAssessedTax",
|
||||
"Revenues"), "2026-01-01", (300, 400))
|
||||
assert values == {"2024-09-28": 999}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_default_identification_tracks_the_installed_version(monkeypatch):
|
||||
"""A release should identify itself, not a version frozen in the source."""
|
||||
monkeypatch.delenv("SEC_EDGAR_USER_AGENT", raising=False)
|
||||
monkeypatch.setattr(sec_edgar.metadata, "version", lambda name: "9.9.9")
|
||||
assert sec_edgar._user_agent() == "TradingAgents/9.9.9 (contact@example.com)"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_uninstalled_checkout_still_identifies_itself(monkeypatch):
|
||||
monkeypatch.delenv("SEC_EDGAR_USER_AGENT", raising=False)
|
||||
|
||||
def _missing(name):
|
||||
raise sec_edgar.metadata.PackageNotFoundError(name)
|
||||
|
||||
monkeypatch.setattr(sec_edgar.metadata, "version", _missing)
|
||||
assert "@" in sec_edgar._user_agent()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_capital_expenditure_is_found_under_either_tag_filers_use(monkeypatch):
|
||||
"""NVIDIA and Amazon report purchases of productive assets, not of property and equipment."""
|
||||
facts = {"facts": {"us-gaap": {"PaymentsToAcquireProductiveAssets": {"units": {"USD": [
|
||||
_fact("2024-12-31", 70_000_000, "2025-02-10", start="2024-01-01")]}}}}}
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json",
|
||||
lambda url: TICKER_MAP if "company_tickers" in url else facts)
|
||||
out = sec_edgar.get_cashflow("AAPL", "annual", "2025-03-01")
|
||||
assert [r for r in out.splitlines() if r.startswith("Capital Expenditure")] == ["Capital Expenditure,70"]
|
||||
|
||||
|
||||
def _columns(out):
|
||||
return [line for line in out.splitlines() if line.startswith(",")][0].split(",")[1:]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_an_annual_balance_sheet_has_no_quarter_end_columns():
|
||||
"""A balance has no span, so a 10-Q's quarter-end balance passed as annual."""
|
||||
annual = _columns(sec_edgar.get_balance_sheet("AAPL", "annual", "2024-11-15"))
|
||||
quarterly = _columns(sec_edgar.get_balance_sheet("AAPL", "quarterly", "2024-11-15"))
|
||||
assert "2022-03-26" not in annual and "2024-09-28" in annual
|
||||
assert "2022-03-26" in quarterly
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_twelve_month_total_from_a_quarterly_report_is_not_a_fiscal_year(monkeypatch):
|
||||
"""Amazon's 10-Qs report trailing twelve months, which passed the annual span
|
||||
check and read as fiscal years overlapping the real ones."""
|
||||
facts = {"facts": {"us-gaap": {"NetCashProvidedByUsedInOperatingActivities": {"units": {"USD": [
|
||||
_fact("2024-12-31", 115_000_000_000, "2025-02-07", start="2024-01-01"),
|
||||
_fact("2025-03-31", 113_000_000_000, "2025-05-02", form="10-Q", fp="Q1", start="2024-04-01"),
|
||||
]}}}}}
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json",
|
||||
lambda url: TICKER_MAP if "company_tickers" in url else facts)
|
||||
assert _columns(sec_edgar.get_cashflow("AAPL", "annual", "2025-06-01")) == ["2024-12-31"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_recast_outside_the_annual_report_still_counts_from_its_filing(monkeypatch):
|
||||
"""Filers recast past years in an 8-K after a split or spin-off. The annual
|
||||
report decides the columns; the value is the latest filing of any form."""
|
||||
facts = {"facts": {"us-gaap": {"EarningsPerShareDiluted": {"units": {"USD/shares": [
|
||||
_fact("2017-03-31", 16.97, "2017-06-15", form="20-F", start="2016-04-01"),
|
||||
_fact("2017-03-31", 2.12, "2019-09-30", form="6-K", start="2016-04-01"),
|
||||
]}}}}}
|
||||
monkeypatch.setattr(sec_edgar, "_fetch_json",
|
||||
lambda url: TICKER_MAP if "company_tickers" in url else facts)
|
||||
|
||||
def eps(date):
|
||||
out = sec_edgar.get_income_statement("AAPL", "annual", date)
|
||||
return [line for line in out.splitlines() if line.startswith("Diluted EPS")][0].split(",")[1]
|
||||
|
||||
assert eps("2019-01-01") == "16.97"
|
||||
assert eps("2020-01-01") == "2.12"
|
||||
@@ -0,0 +1,94 @@
|
||||
"""The rating heuristic that reads the decision's 5-tier rating.
|
||||
|
||||
The Portfolio Manager's rendered decision always carries a ``**Rating**: X``
|
||||
header, so the rating is read deterministically; no second model call is made.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.agents.rating import RATING_REVIEW, RATINGS_5_TIER, extract_rating, parse_rating
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Heuristic parser
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestParseRating:
|
||||
def test_explicit_label_buy(self):
|
||||
assert parse_rating("Rating: Buy\nReasoning here.") == "Buy"
|
||||
|
||||
def test_explicit_label_overweight(self):
|
||||
assert parse_rating("Rating: Overweight\nDetails.") == "Overweight"
|
||||
|
||||
def test_explicit_label_with_markdown_bold_value(self):
|
||||
# Regression: Rating: **Sell** — markdown around the value.
|
||||
assert parse_rating("Rating: **Sell**\nExit immediately.") == "Sell"
|
||||
|
||||
def test_explicit_label_with_markdown_bold_label(self):
|
||||
assert parse_rating("**Rating**: Underweight\nTrim exposure.") == "Underweight"
|
||||
|
||||
def test_rendered_pm_markdown_shape(self):
|
||||
# The exact shape produced by render_pm_decision must always parse.
|
||||
text = (
|
||||
"**Rating**: Buy\n\n"
|
||||
"**Executive Summary**: Enter at $189-192, 6% portfolio cap.\n\n"
|
||||
"**Investment Thesis**: AI capex cycle intact; institutional flows constructive."
|
||||
)
|
||||
assert parse_rating(text) == "Buy"
|
||||
|
||||
def test_explicit_label_wins_over_prose_with_markdown(self):
|
||||
text = (
|
||||
"The buy thesis is weakened by guidance.\n"
|
||||
"Rating: **Sell**\n"
|
||||
"Exit before earnings."
|
||||
)
|
||||
assert parse_rating(text) == "Sell"
|
||||
|
||||
def test_no_rating_is_flagged_for_review_not_defaulted(self):
|
||||
# A decision nobody can read is not a Hold; recording one invents a call.
|
||||
assert parse_rating("No clear directional signal at this time.") == RATING_REVIEW
|
||||
|
||||
def test_no_rating_custom_default(self):
|
||||
assert parse_rating("Plain prose.", default="Underweight") == "Underweight"
|
||||
|
||||
def test_all_five_tiers_recognised(self):
|
||||
for r in RATINGS_5_TIER:
|
||||
assert parse_rating(f"Rating: {r}") == r
|
||||
|
||||
def test_fullwidth_colon_is_parsed_not_reviewed(self):
|
||||
# `Rating:Overweight` (fullwidth colon) is read, not sent to review (#1170).
|
||||
assert parse_rating("Rating:Overweight\n理由はこちら。") == "Overweight"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestExtractRating:
|
||||
def test_returns_none_when_absent(self):
|
||||
assert extract_rating("No directional call here.") is None
|
||||
assert extract_rating("") is None
|
||||
|
||||
def test_whole_word_only(self):
|
||||
# substrings inside larger words must not match
|
||||
assert extract_rating("The buyer was holding shares.") is None
|
||||
|
||||
def test_parse_rating_defaults_to_review(self):
|
||||
# The memory log tags an unreadable decision REVIEW, never a tradeable rating.
|
||||
assert parse_rating("No rating here.") == RATING_REVIEW
|
||||
assert parse_rating("No rating here.", default="Underweight") == "Underweight"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGraphSignalContract:
|
||||
"""The graph-facing signal (TradingAgentsGraph.process_signal) honors the
|
||||
documented "5-tier or REVIEW" contract, not just the parser in isolation."""
|
||||
|
||||
def _bare_graph(self):
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
g = object.__new__(TradingAgentsGraph)
|
||||
return g
|
||||
|
||||
def test_graph_surfaces_review(self):
|
||||
assert self._bare_graph().process_signal("no rating in here") == RATING_REVIEW
|
||||
|
||||
def test_graph_returns_rating(self):
|
||||
assert self._bare_graph().process_signal("**Rating**: Sell") == "Sell"
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Historical social sentiment must not leak current data into a backtest (#1220).
|
||||
|
||||
StockTwits and Reddit fetchers pull only recent items, so for a historical run
|
||||
they must be trimmed to the analysis window (and yield a clear placeholder when
|
||||
nothing qualifies) rather than showing today's chatter as if it were from the
|
||||
as-of date. All three sources share dataflows.date_window.in_window.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.date_window import in_window
|
||||
from tradingagents.dataflows.vendors import reddit, stocktwits
|
||||
|
||||
|
||||
class _JsonResp:
|
||||
"""Minimal urlopen() context-manager stub returning a JSON body."""
|
||||
|
||||
def __init__(self, payload):
|
||||
self._body = json.dumps(payload).encode()
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return self._body
|
||||
|
||||
|
||||
# --- shared window helper ---------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_in_window_bounds_and_exclusive_upper():
|
||||
start = datetime(2026, 5, 1)
|
||||
end = datetime(2026, 5, 9)
|
||||
assert in_window(datetime(2026, 5, 5, tzinfo=timezone.utc), start, end) is True
|
||||
assert in_window(datetime(2026, 5, 9, 23, 59, tzinfo=timezone.utc), start, end) is True
|
||||
# exactly midnight after end -> excluded (no leak)
|
||||
assert in_window(datetime(2026, 5, 10, 0, 0, tzinfo=timezone.utc), start, end) is False
|
||||
# offset-aware converted, not truncated: 05-10T01:00+05:00 == 05-09T20:00Z
|
||||
assert in_window(datetime.fromisoformat("2026-05-10T01:00:00+05:00"), start, end) is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_in_window_undated_excluded_in_backtest_kept_live():
|
||||
old = datetime(2026, 5, 9)
|
||||
assert in_window(None, datetime(2026, 5, 1), old) is False # historical
|
||||
now = datetime.now(timezone.utc)
|
||||
assert in_window(None, now, now) is True # live
|
||||
|
||||
|
||||
# --- StockTwits -------------------------------------------------------------
|
||||
|
||||
def _msg(created_iso, sentiment=None):
|
||||
return {
|
||||
"created_at": created_iso,
|
||||
"user": {"username": "u"},
|
||||
"entities": {"sentiment": {"basic": sentiment}},
|
||||
"body": "text",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stocktwits_historical_window_excludes_recent(monkeypatch):
|
||||
# All messages are "today"; a run as-of a past week must show none of them.
|
||||
recent = [_msg("2026-08-30T12:00:00Z", "Bullish"), _msg("2026-08-29T09:00:00Z")]
|
||||
monkeypatch.setattr(stocktwits, "urlopen", lambda *a, **k: _JsonResp({"messages": recent}))
|
||||
out = stocktwits.fetch_stocktwits_messages("AAPL", start_date="2026-05-01", end_date="2026-05-08")
|
||||
assert "2026-05-01..2026-05-08" in out
|
||||
assert "Bullish: 1" not in out # the recent bullish message did not leak
|
||||
# Coverage starts after the window: unavailable, never a claim of silence.
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stocktwits_live_window_keeps_in_range(monkeypatch):
|
||||
msgs = [_msg("2026-05-05T12:00:00Z", "Bullish"), _msg("2026-05-07T09:00:00Z", "Bearish")]
|
||||
monkeypatch.setattr(stocktwits, "urlopen", lambda *a, **k: _JsonResp({"messages": msgs}))
|
||||
out = stocktwits.fetch_stocktwits_messages("AAPL", start_date="2026-05-01", end_date="2026-05-08")
|
||||
assert "Total: 2" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stocktwits_no_window_is_unfiltered(monkeypatch):
|
||||
msgs = [_msg("2026-08-30T12:00:00Z", "Bullish")]
|
||||
monkeypatch.setattr(stocktwits, "urlopen", lambda *a, **k: _JsonResp({"messages": msgs}))
|
||||
out = stocktwits.fetch_stocktwits_messages("AAPL") # live caller, no dates
|
||||
assert "Total: 1" in out
|
||||
|
||||
|
||||
# --- Reddit -----------------------------------------------------------------
|
||||
|
||||
def _epoch(date_str):
|
||||
return int(datetime.strptime(date_str, "%Y-%m-%d").replace(tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_historical_window_excludes_recent(monkeypatch):
|
||||
posts = [{"title": "NOW", "created_utc": _epoch("2026-08-30"), "source": "rss"}]
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: posts)
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date="2026-05-01", end_date="2026-05-08",
|
||||
)
|
||||
assert "NOW" not in out
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_live_window_keeps_in_range(monkeypatch):
|
||||
posts = [{"title": "INRANGE", "created_utc": _epoch("2026-05-05"), "source": "rss"}]
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: posts)
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date="2026-05-01", end_date="2026-05-08",
|
||||
)
|
||||
assert "INRANGE" in out
|
||||
|
||||
|
||||
# --- coverage vs absence --------------------------------------------------------
|
||||
# The public feeds only serve recent items. When everything fetched postdates the
|
||||
# window the source cannot answer for that date; reporting "no posts" there is a
|
||||
# claim about the market that was never observed.
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stocktwits_covered_but_empty_window_is_a_real_absence(monkeypatch):
|
||||
# The stream reaches back before the window (an older message exists) yet
|
||||
# nothing falls inside it: that is genuine silence.
|
||||
msgs = [_msg("2026-08-30T12:00:00Z"), _msg("2026-04-20T12:00:00Z")]
|
||||
monkeypatch.setattr(stocktwits, "urlopen", lambda *a, **k: _JsonResp({"messages": msgs}))
|
||||
out = stocktwits.fetch_stocktwits_messages("AAPL", start_date="2026-05-01", end_date="2026-05-08")
|
||||
assert "no StockTwits messages" in out
|
||||
assert "unavailable" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_covered_but_empty_window_is_a_real_absence(monkeypatch):
|
||||
posts = [{"title": "NOW", "created_utc": _epoch("2026-08-30"), "source": "rss"},
|
||||
{"title": "OLD", "created_utc": _epoch("2026-04-20"), "source": "rss"}]
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: posts)
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date="2026-05-01", end_date="2026-05-08",
|
||||
)
|
||||
assert "no reddit posts" in out.lower()
|
||||
assert "unavailable" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_empty_feed_for_an_old_window_is_unavailable(monkeypatch):
|
||||
# Search is limited to the last week, so an empty response says nothing
|
||||
# about a window from months ago: there are no timestamps to go on, and the
|
||||
# lookback bound alone must decide.
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: [])
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date="2024-05-01", end_date="2024-05-08",
|
||||
)
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_live_empty_feed_is_a_real_absence(monkeypatch):
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: [])
|
||||
out = reddit.fetch_reddit_posts("AAPL", subreddits=("stocks",))
|
||||
assert "no reddit posts" in out.lower() and "past 7 days" in out
|
||||
assert "unavailable" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_stocktwits_empty_stream_for_a_past_window_is_unavailable(monkeypatch):
|
||||
monkeypatch.setattr(stocktwits, "urlopen", lambda *a, **k: _JsonResp({"messages": []}))
|
||||
out = stocktwits.fetch_stocktwits_messages("AAPL", start_date="2026-05-01", end_date="2026-05-08")
|
||||
assert "unavailable" in out and "not an absence" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_window_straddling_the_lookback_is_unavailable(monkeypatch):
|
||||
# Ten days ago through five days ago: the week-long search never reaches the
|
||||
# first three days, so an empty result cannot stand for the whole window.
|
||||
from datetime import timedelta
|
||||
today = datetime.now(timezone.utc).date()
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: [])
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date=str(today - timedelta(days=10)), end_date=str(today - timedelta(days=5)),
|
||||
)
|
||||
assert "unavailable" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_standard_week_window_empty_is_a_real_absence(monkeypatch):
|
||||
# The graph's window is [trade_date - 7, trade_date]; the week-long search
|
||||
# covers it, so an empty result is genuine silence.
|
||||
from datetime import timedelta
|
||||
today = datetime.now(timezone.utc).date()
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: [])
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date=str(today - timedelta(days=7)), end_date=str(today),
|
||||
)
|
||||
assert "no reddit posts" in out.lower() and "unavailable" not in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_reddit_full_page_does_not_vouch_for_older_days(monkeypatch):
|
||||
# 100 posts from today say nothing about five days ago: the page may have
|
||||
# cut older matches off, so the window stays unavailable.
|
||||
from datetime import timedelta
|
||||
today = datetime.now(timezone.utc).date()
|
||||
ts = _epoch(str(today))
|
||||
page = [{"title": f"T{i}", "created_utc": ts, "subreddit": "stocks"} for i in range(reddit._FEED_PAGE)]
|
||||
monkeypatch.setattr(reddit, "_fetch_subreddit_rss", lambda *a, **k: page)
|
||||
out = reddit.fetch_reddit_posts(
|
||||
"AAPL", subreddits=("stocks",),
|
||||
start_date=str(today - timedelta(days=6)), end_date=str(today - timedelta(days=5)),
|
||||
)
|
||||
assert "unavailable" in out and "no reddit posts" not in out.lower()
|
||||
@@ -0,0 +1,123 @@
|
||||
"""StockTwits fetch: transport-error resilience (#1024) and crypto symbol
|
||||
mapping (#1113).
|
||||
|
||||
StockTwits lists crypto under ``<BASE>.X`` (Yahoo's ``BTC-USD`` 404s), and any
|
||||
transport error must degrade to a placeholder rather than raise.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import http.client
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
from urllib.error import HTTPError
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.vendors import stocktwits
|
||||
|
||||
|
||||
def _raise(exc):
|
||||
class _Resp:
|
||||
def __enter__(self_inner):
|
||||
return self_inner
|
||||
|
||||
def __exit__(self_inner, *a):
|
||||
return False
|
||||
|
||||
def read(self_inner):
|
||||
raise exc
|
||||
return _Resp()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStockTwitsResilience:
|
||||
@pytest.mark.parametrize(
|
||||
"exc",
|
||||
[
|
||||
http.client.IncompleteRead(b""),
|
||||
HTTPError("url", 503, "down", {}, None),
|
||||
TimeoutError("slow"),
|
||||
],
|
||||
)
|
||||
def test_transport_errors_return_placeholder(self, exc):
|
||||
with patch.object(stocktwits, "urlopen", return_value=_raise(exc)):
|
||||
out = stocktwits.fetch_stocktwits_messages("NVDA")
|
||||
assert "unavailable" in out.lower()
|
||||
assert out.startswith("<stocktwits unavailable")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStockTwitsCryptoSymbols:
|
||||
@pytest.mark.parametrize(
|
||||
("ticker", "expected"),
|
||||
[
|
||||
("BTC-USD", "BTC.X"),
|
||||
("eth-usd", "ETH.X"),
|
||||
("SOL-USD", "SOL.X"),
|
||||
("BTCUSD", "BTC.X"), # undashed broker form
|
||||
("BTC-USDT", "BTC.X"), # stablecoin quote
|
||||
("AMD", "AMD"),
|
||||
("BRK-B", "BRK-B"), # dashed class share: untouched
|
||||
("GOLD", "GOLD"), # real equity (aliases elsewhere): untouched here
|
||||
("XYZ-USD", "XYZ-USD"), # unknown base: not treated as crypto
|
||||
],
|
||||
)
|
||||
def test_symbol_mapping(self, ticker, expected):
|
||||
assert stocktwits._stocktwits_symbol(ticker) == expected
|
||||
|
||||
def test_crypto_pair_requests_dot_x_endpoint(self):
|
||||
seen = {}
|
||||
|
||||
def fake_urlopen(req, timeout=None):
|
||||
seen["url"] = req.full_url
|
||||
raise TimeoutError("stop after capturing the URL")
|
||||
|
||||
with patch.object(stocktwits, "urlopen", side_effect=fake_urlopen):
|
||||
stocktwits.fetch_stocktwits_messages("BTC-USD")
|
||||
assert "/symbol/BTC.X.json" in seen["url"]
|
||||
|
||||
|
||||
def _stream(*bodies):
|
||||
payload = {"messages": [
|
||||
{"body": b, "created_at": "2026-01-09T15:00:00Z", "user": {"username": "u"},
|
||||
"entities": {"sentiment": {"basic": "Bullish"}}}
|
||||
for b in bodies
|
||||
]}
|
||||
|
||||
class _Resp:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return json.dumps(payload).encode()
|
||||
return _Resp()
|
||||
|
||||
|
||||
def _drop_spam(texts):
|
||||
return [not t.startswith("SPAM") for t in texts], "Screened: note"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStockTwitsScreening:
|
||||
def test_screened_out_messages_leave_the_block_and_its_counts(self):
|
||||
with patch.object(stocktwits, "urlopen", return_value=_stream("SPAM", "long NVDA")):
|
||||
out = stocktwits.fetch_stocktwits_messages("NVDA", screen=_drop_spam)
|
||||
assert out.startswith("Screened: note")
|
||||
assert "long NVDA" in out and "SPAM" not in out
|
||||
assert "Total: 1 most-recent" in out
|
||||
|
||||
def test_all_screened_out_is_not_called_empty(self):
|
||||
with patch.object(stocktwits, "urlopen", return_value=_stream("SPAM")):
|
||||
out = stocktwits.fetch_stocktwits_messages("NVDA", screen=_drop_spam)
|
||||
assert "none of the 1 StockTwits messages is about $NVDA" in out
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_html_entities_in_message_bodies_are_decoded():
|
||||
with patch.object(stocktwits, "urlopen", return_value=_stream("S&P wasn't up")):
|
||||
out = stocktwits.fetch_stocktwits_messages("NVDA")
|
||||
assert "S&P wasn't up" in out
|
||||
@@ -0,0 +1,148 @@
|
||||
"""Agents on the schema-only structured-output path must not invite tool calls (#1130).
|
||||
|
||||
`with_structured_output` binds exactly one tool (the schema). A prompt that
|
||||
primes tool use makes models emit an unknown `web_search` call, which discards
|
||||
the structured attempt and forces a free-text retry — an extra LLM round trip
|
||||
and the loss of typed output.
|
||||
|
||||
These assert the constraint reaches the *rendered* prompt each agent actually
|
||||
sends, not merely that the constant is referenced in the module.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import tradingagents.agents.analysts.sentiment_analyst as sentiment
|
||||
from tradingagents.agents.managers.portfolio_manager import create_portfolio_manager
|
||||
from tradingagents.agents.managers.research_manager import create_research_manager
|
||||
from tradingagents.agents.structured import NO_EXTERNAL_TOOLS
|
||||
from tradingagents.agents.trader.trader import create_trader
|
||||
|
||||
|
||||
def _capturing_llm(captured: dict, result):
|
||||
"""LLM whose structured binding records the prompt it was handed."""
|
||||
structured = MagicMock()
|
||||
structured.invoke.side_effect = lambda prompt: (
|
||||
captured.__setitem__("prompt", prompt) or result
|
||||
)
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.return_value = structured
|
||||
return llm
|
||||
|
||||
|
||||
def _prompt_text(prompt) -> str:
|
||||
"""Flatten a captured prompt (str, message list, or objects) to text."""
|
||||
if isinstance(prompt, str):
|
||||
return prompt
|
||||
parts = []
|
||||
for m in prompt:
|
||||
parts.append(m.get("content", "") if isinstance(m, dict) else getattr(m, "content", ""))
|
||||
return "\n".join(str(p) for p in parts)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_trader_prompt_states_constraint():
|
||||
from tradingagents.agents.schemas import TraderAction, TraderProposal
|
||||
|
||||
captured = {}
|
||||
llm = _capturing_llm(captured, TraderProposal(action=TraderAction.BUY, reasoning="x"))
|
||||
create_trader(llm)({
|
||||
"company_of_interest": "NVDA",
|
||||
"investment_plan": "**Recommendation**: Buy",
|
||||
"market_report": "Current price $189.5; ATR 4.2.",
|
||||
})
|
||||
assert NO_EXTERNAL_TOOLS in _prompt_text(captured["prompt"])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_research_manager_prompt_states_constraint():
|
||||
from tradingagents.agents.schemas import PortfolioRating, ResearchPlan
|
||||
|
||||
captured = {}
|
||||
llm = _capturing_llm(
|
||||
captured,
|
||||
ResearchPlan(
|
||||
recommendation=PortfolioRating.BUY, rationale="x", strategic_actions="y"
|
||||
),
|
||||
)
|
||||
create_research_manager(llm)({
|
||||
"company_of_interest": "NVDA",
|
||||
"investment_debate_state": {
|
||||
"history": "h", "bull_history": "b", "bear_history": "r",
|
||||
"current_response": "", "judge_decision": "", "count": 1,
|
||||
},
|
||||
})
|
||||
assert NO_EXTERNAL_TOOLS in _prompt_text(captured["prompt"])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_portfolio_manager_prompt_states_constraint():
|
||||
from tradingagents.agents.schemas import PortfolioDecision, PortfolioRating
|
||||
|
||||
captured = {}
|
||||
llm = _capturing_llm(
|
||||
captured,
|
||||
PortfolioDecision(
|
||||
rating=PortfolioRating.HOLD,
|
||||
executive_summary="x",
|
||||
investment_thesis="y",
|
||||
),
|
||||
)
|
||||
risk = {
|
||||
"history": "h", "aggressive_history": "a", "conservative_history": "c",
|
||||
"neutral_history": "n", "current_aggressive_response": "",
|
||||
"current_conservative_response": "", "current_neutral_response": "",
|
||||
"latest_speaker": "Neutral", "count": 1,
|
||||
}
|
||||
create_portfolio_manager(llm)({
|
||||
"company_of_interest": "NVDA",
|
||||
"risk_debate_state": risk,
|
||||
"investment_plan": "plan",
|
||||
"trader_investment_plan": "trader plan",
|
||||
})
|
||||
assert NO_EXTERNAL_TOOLS in _prompt_text(captured["prompt"])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_sentiment_prompt_states_constraint(monkeypatch):
|
||||
from tradingagents.agents.schemas import SentimentBand, SentimentReport
|
||||
|
||||
# Pre-fetched sources are stubbed so the prompt builds without network I/O.
|
||||
monkeypatch.setattr(sentiment, "fetch_stocktwits_messages", lambda *a, **k: "st")
|
||||
monkeypatch.setattr(sentiment, "fetch_reddit_posts", lambda *a, **k: "rd")
|
||||
monkeypatch.setattr(sentiment.get_news, "func", lambda *a, **k: "news", raising=False)
|
||||
|
||||
captured = {}
|
||||
llm = _capturing_llm(captured, SentimentReport(
|
||||
overall_band=SentimentBand.BULLISH, overall_score=7.5,
|
||||
confidence="high", narrative="n",
|
||||
))
|
||||
sentiment.create_sentiment_analyst(llm)({
|
||||
"company_of_interest": "NVDA", "trade_date": "2026-01-15",
|
||||
"asset_type": "stock", "messages": [],
|
||||
})
|
||||
text = _prompt_text(captured["prompt"])
|
||||
assert NO_EXTERNAL_TOOLS in text
|
||||
# This agent binds no tools, so tool-range wording must not reappear.
|
||||
assert "tool-call date ranges" not in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_tool_using_analysts_keep_their_date_guidance():
|
||||
# The analysts that really do call tools keep the wording that anchors their
|
||||
# tool date ranges (#836) — this fix is scoped to no-tool agents.
|
||||
import tradingagents.agents.analysts.market_analyst as market
|
||||
import tradingagents.agents.analysts.news_analyst as news
|
||||
for module in (market, news):
|
||||
assert "tool-call date ranges" in inspect.getsource(module)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_constraint_text_is_unambiguous():
|
||||
assert "do not call external tools" in NO_EXTERNAL_TOOLS.lower()
|
||||
# No template braces: it is embedded in ChatPromptTemplate strings, where
|
||||
# braces would be parsed as input variables.
|
||||
assert "{" not in NO_EXTERNAL_TOOLS and "}" not in NO_EXTERNAL_TOOLS
|
||||
@@ -0,0 +1,548 @@
|
||||
"""Tests for structured-output agents (Trader, Research Manager, Sentiment Analyst).
|
||||
|
||||
The Portfolio Manager has its own coverage in tests/test_memory_log.py
|
||||
(which exercises the full memory-log → PM injection cycle). This file
|
||||
covers the parallel schemas, render functions, and graceful-fallback
|
||||
behavior we added for the Trader, Research Manager, and Sentiment Analyst
|
||||
so they share the same deterministic output shape.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from tradingagents.agents.analysts.sentiment_analyst import create_sentiment_analyst
|
||||
from tradingagents.agents.managers.portfolio_manager import create_portfolio_manager
|
||||
from tradingagents.agents.managers.research_manager import create_research_manager
|
||||
from tradingagents.agents.schemas import (
|
||||
PortfolioDecision,
|
||||
PortfolioRating,
|
||||
ResearchPlan,
|
||||
SentimentBand,
|
||||
SentimentReport,
|
||||
TraderAction,
|
||||
TraderProposal,
|
||||
render_research_plan,
|
||||
render_sentiment_report,
|
||||
render_trader_proposal,
|
||||
)
|
||||
from tradingagents.agents.trader.trader import create_trader
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Render functions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRenderTraderProposal:
|
||||
def test_minimal_required_fields(self):
|
||||
p = TraderProposal(action=TraderAction.HOLD, reasoning="Balanced setup; no edge.")
|
||||
md = render_trader_proposal(p)
|
||||
assert "**Action**: Hold" in md
|
||||
assert "**Reasoning**: Balanced setup; no edge." in md
|
||||
# The trailing FINAL TRANSACTION PROPOSAL line is preserved for the
|
||||
# analyst stop-signal text and any external code that greps for it.
|
||||
assert "FINAL TRANSACTION PROPOSAL: **HOLD**" in md
|
||||
|
||||
def test_optional_fields_included_when_present(self):
|
||||
p = TraderProposal(
|
||||
action=TraderAction.BUY,
|
||||
reasoning="Strong technicals + fundamentals.",
|
||||
entry_price=189.5,
|
||||
stop_loss=178.0,
|
||||
position_sizing="6% of portfolio",
|
||||
)
|
||||
md = render_trader_proposal(p)
|
||||
assert "**Action**: Buy" in md
|
||||
assert "**Entry Price**: 189.5" in md
|
||||
assert "**Stop Loss**: 178.0" in md
|
||||
assert "**Position Sizing**: 6% of portfolio" in md
|
||||
assert "FINAL TRANSACTION PROPOSAL: **BUY**" in md
|
||||
|
||||
def test_optional_fields_are_named_as_not_provided(self):
|
||||
"""An omitted line reads as a field nobody asked for; the reader cannot
|
||||
tell it from a level the trader declined to set."""
|
||||
p = TraderProposal(action=TraderAction.SELL, reasoning="Guidance cut.")
|
||||
md = render_trader_proposal(p)
|
||||
for field in ("Entry Price", "Stop Loss", "Position Sizing"):
|
||||
assert f"**{field}**: not provided" in md
|
||||
assert "FINAL TRANSACTION PROPOSAL: **SELL**" in md
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestNullishFloatCoercion:
|
||||
"""A weak LLM may write "None"/"N/A" into an optional float field (#1058);
|
||||
coerce those to None so the structured call validates instead of erroring."""
|
||||
|
||||
def test_trader_nullish_strings_coerce_to_none(self):
|
||||
for sentinel in ("None", "N/A", "null", "-", "", "TBD"):
|
||||
p = TraderProposal(
|
||||
action=TraderAction.HOLD,
|
||||
reasoning="x",
|
||||
entry_price=sentinel,
|
||||
stop_loss=sentinel,
|
||||
)
|
||||
assert p.entry_price is None
|
||||
assert p.stop_loss is None
|
||||
|
||||
def test_trader_real_numeric_string_still_parses(self):
|
||||
p = TraderProposal(action=TraderAction.BUY, reasoning="x", entry_price="189.5")
|
||||
assert p.entry_price == 189.5
|
||||
|
||||
def test_pm_nullish_price_target_coerces_to_none(self):
|
||||
d = PortfolioDecision(
|
||||
rating=PortfolioRating.OVERWEIGHT,
|
||||
executive_summary="s",
|
||||
investment_thesis="t",
|
||||
price_target="N/A",
|
||||
)
|
||||
assert d.price_target is None
|
||||
|
||||
def test_percentage_answer_to_a_price_field_becomes_none(self):
|
||||
# The Trader is asked for concrete levels and may answer a price field
|
||||
# with a distance ("15%"), which failed the whole proposal (#1288).
|
||||
# A percentage cannot be salvaged: 15% must not become a $15 stop.
|
||||
for pct in ("15%", " 7.5% ", "-10%"):
|
||||
p = TraderProposal(
|
||||
action=TraderAction.BUY,
|
||||
reasoning="x",
|
||||
entry_price=pct,
|
||||
stop_loss=pct,
|
||||
)
|
||||
assert p.entry_price is None
|
||||
assert p.stop_loss is None
|
||||
|
||||
def test_human_formatted_price_is_reduced_to_its_number(self):
|
||||
p = TraderProposal(
|
||||
action=TraderAction.BUY,
|
||||
reasoning="x",
|
||||
entry_price="$1,234.50",
|
||||
stop_loss="1,180",
|
||||
)
|
||||
assert p.entry_price == 1234.50
|
||||
assert p.stop_loss == 1180.0
|
||||
|
||||
def test_one_bad_field_no_longer_fails_the_whole_proposal(self):
|
||||
# Previously a single '15%' raised, forcing a free-text retry that lost
|
||||
# the action and reasoning; now the rest of the proposal survives.
|
||||
p = TraderProposal(
|
||||
action=TraderAction.SELL,
|
||||
reasoning="downgrade on margin compression",
|
||||
entry_price="612.40",
|
||||
stop_loss="15%",
|
||||
)
|
||||
assert p.action is TraderAction.SELL
|
||||
assert p.entry_price == 612.40
|
||||
assert p.stop_loss is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRenderResearchPlan:
|
||||
def test_required_fields(self):
|
||||
p = ResearchPlan(
|
||||
recommendation=PortfolioRating.OVERWEIGHT,
|
||||
rationale="Bull case carried; tailwinds intact.",
|
||||
strategic_actions="Build position over two weeks; cap at 5%.",
|
||||
)
|
||||
md = render_research_plan(p)
|
||||
assert "**Recommendation**: Overweight" in md
|
||||
assert "**Rationale**: Bull case carried" in md
|
||||
assert "**Strategic Actions**: Build position" in md
|
||||
|
||||
def test_all_5_tier_ratings_render(self):
|
||||
for rating in PortfolioRating:
|
||||
p = ResearchPlan(
|
||||
recommendation=rating,
|
||||
rationale="r",
|
||||
strategic_actions="s",
|
||||
)
|
||||
md = render_research_plan(p)
|
||||
assert f"**Recommendation**: {rating.value}" in md
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Trader agent: structured happy path + fallback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_trader_state():
|
||||
return {
|
||||
"company_of_interest": "NVDA",
|
||||
"investment_plan": "**Recommendation**: Buy\n**Rationale**: ...\n**Strategic Actions**: ...",
|
||||
"market_report": "Current price $189.5; 14-day ATR 4.2; support $178, resistance $196.",
|
||||
}
|
||||
|
||||
|
||||
def _structured_trader_llm(captured: dict, proposal: TraderProposal | None = None):
|
||||
"""Build a MagicMock LLM whose with_structured_output binding captures the
|
||||
prompt and returns a real TraderProposal so render_trader_proposal works.
|
||||
"""
|
||||
if proposal is None:
|
||||
proposal = TraderProposal(
|
||||
action=TraderAction.BUY,
|
||||
reasoning="Strong setup.",
|
||||
)
|
||||
structured = MagicMock()
|
||||
structured.invoke.side_effect = lambda prompt: (
|
||||
captured.__setitem__("prompt", prompt) or proposal
|
||||
)
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.return_value = structured
|
||||
return llm
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_invoke_structured_falls_back_when_result_is_none():
|
||||
# A thinking model can answer in plain text, leaving the parser with None.
|
||||
# That must fall back to free text, not crash on render(None) (#1051).
|
||||
from tradingagents.agents.structured import invoke_structured_or_freetext
|
||||
|
||||
structured = MagicMock()
|
||||
structured.invoke.return_value = None
|
||||
plain = MagicMock()
|
||||
plain.invoke.return_value = MagicMock(content="FREETEXT")
|
||||
|
||||
out = invoke_structured_or_freetext(
|
||||
structured, plain, "prompt", render=lambda r: r.rating, agent_name="t"
|
||||
)
|
||||
assert out == "FREETEXT"
|
||||
plain.invoke.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestTraderAgent:
|
||||
def test_structured_path_produces_rendered_markdown(self):
|
||||
captured = {}
|
||||
proposal = TraderProposal(
|
||||
action=TraderAction.BUY,
|
||||
reasoning="AI capex cycle intact; institutional flows constructive.",
|
||||
entry_price=189.5,
|
||||
stop_loss=178.0,
|
||||
position_sizing="6% of portfolio",
|
||||
)
|
||||
llm = _structured_trader_llm(captured, proposal)
|
||||
trader = create_trader(llm)
|
||||
result = trader(_make_trader_state())
|
||||
plan = result["trader_investment_plan"]
|
||||
assert "**Action**: Buy" in plan
|
||||
assert "**Entry Price**: 189.5" in plan
|
||||
assert "FINAL TRANSACTION PROPOSAL: **BUY**" in plan
|
||||
# The same rendered markdown is also added to messages for downstream agents.
|
||||
assert plan in result["messages"][0].content
|
||||
|
||||
def test_prompt_includes_investment_plan(self):
|
||||
captured = {}
|
||||
llm = _structured_trader_llm(captured)
|
||||
trader = create_trader(llm)
|
||||
trader(_make_trader_state())
|
||||
# The investment plan is in the user message of the captured prompt.
|
||||
prompt = captured["prompt"]
|
||||
assert any("Proposed Investment Plan" in m["content"] for m in prompt)
|
||||
|
||||
def test_prompt_includes_market_report_for_price_levels(self):
|
||||
# #1167: the Trader must see the technical market report so entry/stop
|
||||
# levels are grounded in real price structure, not just the digested plan.
|
||||
captured = {}
|
||||
trader = create_trader(_structured_trader_llm(captured))
|
||||
trader(_make_trader_state())
|
||||
user = " ".join(m["content"] for m in captured["prompt"] if m["role"] == "user")
|
||||
system = " ".join(m["content"] for m in captured["prompt"] if m["role"] == "system")
|
||||
assert "Technical Market Report:" in user
|
||||
assert "14-day ATR 4.2" in user # the actual report content reached the Trader
|
||||
assert "support $178, resistance $196" in user
|
||||
assert "Ground concrete price levels" in system
|
||||
|
||||
def test_empty_market_report_omits_the_section_and_grounding(self):
|
||||
# #1167: when the market analyst wasn't selected the report is empty, so
|
||||
# don't tell the Trader to ground levels in a report it doesn't have.
|
||||
captured = {}
|
||||
state = _make_trader_state()
|
||||
state["market_report"] = ""
|
||||
create_trader(_structured_trader_llm(captured))(state)
|
||||
text = " ".join(m["content"] for m in captured["prompt"])
|
||||
assert "Technical Market Report:" not in text
|
||||
assert "Ground concrete price levels" not in text
|
||||
assert "Proposed Investment Plan" in text # still present
|
||||
|
||||
def test_falls_back_to_freetext_when_structured_unavailable(self):
|
||||
plain_response = (
|
||||
"**Action**: Sell\n\nGuidance cut hits margins.\n\n"
|
||||
"FINAL TRANSACTION PROPOSAL: **SELL**"
|
||||
)
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.side_effect = NotImplementedError("provider unsupported")
|
||||
llm.invoke.return_value = MagicMock(content=plain_response)
|
||||
trader = create_trader(llm)
|
||||
result = trader(_make_trader_state())
|
||||
assert result["trader_investment_plan"] == plain_response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Research Manager agent: structured happy path + fallback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_rm_state():
|
||||
return {
|
||||
"company_of_interest": "NVDA",
|
||||
"investment_debate_state": {
|
||||
"history": "Bull and bear arguments here.",
|
||||
"bull_history": "Bull says...",
|
||||
"bear_history": "Bear says...",
|
||||
"current_response": "",
|
||||
"judge_decision": "",
|
||||
"count": 1,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _structured_rm_llm(captured: dict, plan: ResearchPlan | None = None):
|
||||
if plan is None:
|
||||
plan = ResearchPlan(
|
||||
recommendation=PortfolioRating.HOLD,
|
||||
rationale="Balanced view across both sides.",
|
||||
strategic_actions="Hold current position; reassess after earnings.",
|
||||
)
|
||||
structured = MagicMock()
|
||||
structured.invoke.side_effect = lambda prompt: (
|
||||
captured.__setitem__("prompt", prompt) or plan
|
||||
)
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.return_value = structured
|
||||
return llm
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestResearchManagerAgent:
|
||||
def test_structured_path_produces_rendered_markdown(self):
|
||||
captured = {}
|
||||
plan = ResearchPlan(
|
||||
recommendation=PortfolioRating.OVERWEIGHT,
|
||||
rationale="Bull case is stronger; AI tailwind intact.",
|
||||
strategic_actions="Build position gradually over two weeks.",
|
||||
)
|
||||
llm = _structured_rm_llm(captured, plan)
|
||||
rm = create_research_manager(llm)
|
||||
result = rm(_make_rm_state())
|
||||
ip = result["investment_plan"]
|
||||
assert "**Recommendation**: Overweight" in ip
|
||||
assert "**Rationale**: Bull case" in ip
|
||||
assert "**Strategic Actions**: Build position" in ip
|
||||
|
||||
def test_prompt_uses_5_tier_rating_scale(self):
|
||||
"""The RM prompt must list all five tiers so the schema enum matches user expectations."""
|
||||
captured = {}
|
||||
llm = _structured_rm_llm(captured)
|
||||
rm = create_research_manager(llm)
|
||||
rm(_make_rm_state())
|
||||
prompt = captured["prompt"]
|
||||
for tier in ("Buy", "Overweight", "Hold", "Underweight", "Sell"):
|
||||
assert f"**{tier}**" in prompt, f"missing {tier} in prompt"
|
||||
|
||||
def test_falls_back_to_freetext_when_structured_unavailable(self):
|
||||
plain_response = "**Recommendation**: Sell\n\n**Rationale**: ...\n\n**Strategic Actions**: ..."
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.side_effect = NotImplementedError("provider unsupported")
|
||||
llm.invoke.return_value = MagicMock(content=plain_response)
|
||||
rm = create_research_manager(llm)
|
||||
result = rm(_make_rm_state())
|
||||
assert result["investment_plan"] == plain_response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sentiment Analyst: schema, render, structured happy path + fallback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRenderSentimentReport:
|
||||
def test_header_contains_band_and_score(self):
|
||||
report = SentimentReport(
|
||||
overall_band=SentimentBand.BULLISH,
|
||||
overall_score=7.2,
|
||||
confidence="high",
|
||||
narrative="Source breakdown here.",
|
||||
)
|
||||
md = render_sentiment_report(report)
|
||||
assert "**Overall Sentiment:** **Bullish**" in md
|
||||
assert "(Score: 7.2/10)" in md
|
||||
|
||||
def test_header_contains_confidence(self):
|
||||
report = SentimentReport(
|
||||
overall_band=SentimentBand.NEUTRAL,
|
||||
overall_score=5.0,
|
||||
confidence="low",
|
||||
narrative="Limited data.",
|
||||
)
|
||||
assert "**Confidence:** Low" in render_sentiment_report(report)
|
||||
|
||||
def test_narrative_preserved_in_output(self):
|
||||
narrative = "## Breakdown\n\nStockTwits: 70% bullish.\n\n| Signal | Direction |\n|---|---|\n| News | Neutral |"
|
||||
report = SentimentReport(
|
||||
overall_band=SentimentBand.MILDLY_BULLISH,
|
||||
overall_score=6.0,
|
||||
confidence="medium",
|
||||
narrative=narrative,
|
||||
)
|
||||
assert narrative in render_sentiment_report(report)
|
||||
|
||||
def test_all_six_bands_render(self):
|
||||
for band in SentimentBand:
|
||||
report = SentimentReport(
|
||||
overall_band=band, overall_score=5.0,
|
||||
confidence="medium", narrative="n",
|
||||
)
|
||||
assert band.value in render_sentiment_report(report)
|
||||
|
||||
def test_score_out_of_range_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
SentimentReport(
|
||||
overall_band=SentimentBand.BULLISH, overall_score=11.0,
|
||||
confidence="high", narrative="n",
|
||||
)
|
||||
|
||||
|
||||
def _make_sentiment_state():
|
||||
return {
|
||||
"company_of_interest": "NVDA",
|
||||
"trade_date": "2026-01-15",
|
||||
"asset_type": "stock",
|
||||
"messages": [],
|
||||
}
|
||||
|
||||
|
||||
def _structured_sentiment_llm(captured: dict, report: SentimentReport | None = None):
|
||||
"""MagicMock LLM whose structured binding captures the prompt and returns
|
||||
a real SentimentReport so render_sentiment_report works."""
|
||||
if report is None:
|
||||
report = SentimentReport(
|
||||
overall_band=SentimentBand.BULLISH, overall_score=7.5,
|
||||
confidence="high",
|
||||
narrative="StockTwits 75% bullish. News constructive. Reddit upbeat.",
|
||||
)
|
||||
structured = MagicMock()
|
||||
structured.invoke.side_effect = lambda prompt: (
|
||||
captured.__setitem__("prompt", prompt) or report
|
||||
)
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.return_value = structured
|
||||
return llm
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestSentimentAnalystAgent:
|
||||
@pytest.fixture(autouse=True)
|
||||
def _stub_prefetched_sources(self, monkeypatch):
|
||||
"""Stub the sources the analyst pre-fetches before prompting.
|
||||
|
||||
create_sentiment_analyst fetches news, StockTwits and Reddit itself, so
|
||||
without this these tests hit the live network. A real Reddit 429 then
|
||||
backs the fetcher off for a minute per subreddit, which is what turned
|
||||
this file into a multi-minute hang.
|
||||
"""
|
||||
from tradingagents.agents.analysts import sentiment_analyst as sentiment
|
||||
|
||||
monkeypatch.setattr(sentiment, "fetch_stocktwits_messages", lambda *a, **k: "st")
|
||||
monkeypatch.setattr(sentiment, "fetch_reddit_posts", lambda *a, **k: "rd")
|
||||
monkeypatch.setattr(sentiment.get_news, "func", lambda *a, **k: "news", raising=False)
|
||||
|
||||
def test_structured_path_produces_rendered_markdown(self):
|
||||
captured = {}
|
||||
report = SentimentReport(
|
||||
overall_band=SentimentBand.MILDLY_BEARISH, overall_score=4.0,
|
||||
confidence="medium", narrative="Mixed signals across sources.",
|
||||
)
|
||||
analyst = create_sentiment_analyst(_structured_sentiment_llm(captured, report))
|
||||
sr = analyst(_make_sentiment_state())["sentiment_report"]
|
||||
assert "**Overall Sentiment:** **Mildly Bearish**" in sr
|
||||
assert "(Score: 4.0/10)" in sr
|
||||
assert "Mixed signals across sources." in sr
|
||||
|
||||
def test_sentiment_report_also_in_messages(self):
|
||||
captured = {}
|
||||
analyst = create_sentiment_analyst(_structured_sentiment_llm(captured))
|
||||
result = analyst(_make_sentiment_state())
|
||||
assert len(result["messages"]) == 1
|
||||
assert result["sentiment_report"] == result["messages"][0].content
|
||||
|
||||
def test_prompt_contains_ticker(self):
|
||||
captured = {}
|
||||
create_sentiment_analyst(_structured_sentiment_llm(captured))(_make_sentiment_state())
|
||||
assert any("NVDA" in str(m) for m in captured["prompt"])
|
||||
|
||||
def test_falls_back_to_freetext_when_structured_unavailable(self):
|
||||
plain = "**Overall Sentiment:** **Bearish** (Score: 3.0/10)\n**Confidence:** Low\n\nLimited data."
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.side_effect = NotImplementedError("provider unsupported")
|
||||
llm.invoke.return_value = MagicMock(content=plain)
|
||||
assert create_sentiment_analyst(llm)(_make_sentiment_state())["sentiment_report"] == plain
|
||||
|
||||
def test_falls_back_to_freetext_when_structured_call_fails(self):
|
||||
plain = "Fallback free-text sentiment."
|
||||
structured = MagicMock()
|
||||
structured.invoke.side_effect = ValueError("bad JSON from model")
|
||||
llm = MagicMock()
|
||||
llm.with_structured_output.return_value = structured
|
||||
llm.invoke.return_value = MagicMock(content=plain)
|
||||
assert create_sentiment_analyst(llm)(_make_sentiment_state())["sentiment_report"] == plain
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("source", [
|
||||
pytest.param(lambda: ResearchPlan.model_fields["recommendation"].description, id="ResearchPlan.recommendation"),
|
||||
pytest.param(lambda: PortfolioDecision.model_fields["rating"].description, id="PortfolioDecision.rating"),
|
||||
pytest.param(lambda: inspect.getsource(create_research_manager), id="research_manager prompt"),
|
||||
pytest.param(lambda: inspect.getsource(create_portfolio_manager), id="portfolio_manager prompt"),
|
||||
])
|
||||
def test_conflict_alone_is_not_a_hold_trigger(source):
|
||||
# The debate always contains conflicting arguments, so a Hold condition that
|
||||
# conflict satisfies fires on every run and swallows directional calls
|
||||
# (#1321). All four decision sites must state the same rule.
|
||||
text = " ".join(source().split())
|
||||
assert "conflict alone is not a reason to Hold" in text or \
|
||||
"Conflicting arguments alone are not a reason to Hold" in text
|
||||
assert "materially conflicting" not in text
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("written", ["150-160", "150 to 160", "around 150", "150/160", "~150"])
|
||||
def test_a_price_written_as_a_range_drops_only_that_field(written):
|
||||
"""Anything that is not a single number becomes None. Letting it through
|
||||
fails the whole decision's validation, and the run falls back to free text,
|
||||
losing every other field the model got right."""
|
||||
from tradingagents.agents.schemas import PortfolioDecision, PortfolioRating
|
||||
|
||||
decision = PortfolioDecision(rating=PortfolioRating.BUY, executive_summary="s",
|
||||
investment_thesis="t", price_target=written)
|
||||
assert decision.price_target is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_price_that_is_a_number_survives():
|
||||
from tradingagents.agents.schemas import PortfolioDecision, PortfolioRating
|
||||
|
||||
decision = PortfolioDecision(rating=PortfolioRating.BUY, executive_summary="s",
|
||||
investment_thesis="t", price_target="$1,150.25")
|
||||
assert decision.price_target == 1150.25
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_a_field_the_model_did_not_give_says_so():
|
||||
"""An omitted line and a line never asked for read the same to an analyst."""
|
||||
from tradingagents.agents.schemas import PortfolioDecision, PortfolioRating, render_pm_decision
|
||||
|
||||
rendered = render_pm_decision(PortfolioDecision(
|
||||
rating=PortfolioRating.HOLD, executive_summary="s", investment_thesis="t"))
|
||||
assert "Price Target" in rendered and "not provided" in rendered.lower()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_the_trader_names_the_levels_it_did_not_give():
|
||||
from tradingagents.agents.schemas import TraderAction, TraderProposal, render_trader_proposal
|
||||
|
||||
rendered = render_trader_proposal(TraderProposal(action=TraderAction.HOLD, reasoning="r"))
|
||||
for field in ("Entry Price", "Stop Loss", "Position Sizing"):
|
||||
assert field in rendered
|
||||
assert rendered.lower().count("not provided") == 3
|
||||
@@ -0,0 +1,17 @@
|
||||
"""The suite runs the same on any machine: no test reaches the network."""
|
||||
|
||||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@pytest.mark.parametrize("connect", [
|
||||
lambda: socket.create_connection(("192.0.2.1", 80), timeout=1),
|
||||
lambda: socket.socket().connect_ex(("192.0.2.1", 80)),
|
||||
], ids=["connect", "connect_ex"])
|
||||
def test_a_test_cannot_reach_the_network(connect):
|
||||
"""A test that silently depends on a live vendor passes or fails with the
|
||||
machine it runs on; conftest refuses the connection instead."""
|
||||
with pytest.raises(OSError, match="reach the network"):
|
||||
connect()
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Symbol normalization must apply on every yfinance path, not just price fetch.
|
||||
|
||||
Regression tests for #983 (instrument identity), #984 (reflection returns), and
|
||||
the news path: a broker symbol like XAUUSD must resolve to the same Yahoo symbol
|
||||
(GC=F) that the price path uses, so identity, realized-return, and news lookups
|
||||
hit the right instrument instead of failing/mismatching.
|
||||
"""
|
||||
import pandas as pd
|
||||
|
||||
import tradingagents.agents.context as au
|
||||
import tradingagents.dataflows.vendors.yahoo.market as yahoo_market
|
||||
import tradingagents.dataflows.vendors.yahoo.news as ynews
|
||||
from tradingagents.graph import settlement
|
||||
|
||||
|
||||
def test_identity_lookup_normalizes_symbol(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
class FakeTicker:
|
||||
def __init__(self, symbol):
|
||||
seen["symbol"] = symbol
|
||||
|
||||
@property
|
||||
def info(self):
|
||||
return {"longName": "Gold Futures", "quoteType": "FUTURE"}
|
||||
|
||||
monkeypatch.setattr(yahoo_market.yf, "Ticker", FakeTicker)
|
||||
au.resolve_instrument_identity.cache_clear()
|
||||
|
||||
identity = au.resolve_instrument_identity("XAUUSD")
|
||||
|
||||
assert seen["symbol"] == "GC=F" # normalized, not the raw broker symbol
|
||||
assert identity.get("company_name") == "Gold Futures"
|
||||
|
||||
|
||||
def test_fetch_returns_normalizes_symbol(monkeypatch):
|
||||
queried = []
|
||||
|
||||
class FakeTicker:
|
||||
def __init__(self, symbol):
|
||||
queried.append(symbol)
|
||||
|
||||
def history(self, *args, **kwargs):
|
||||
prices = [100.0, 101.0, 102.0, 103.0, 104.0, 105.0, 106.0]
|
||||
idx = pd.date_range(start="2025-01-02", periods=len(prices), freq="D")
|
||||
return pd.DataFrame({"Close": prices}, index=idx)
|
||||
|
||||
monkeypatch.setattr(yahoo_market.yf, "Ticker", FakeTicker)
|
||||
|
||||
raw, alpha, days, resolved = settlement.fetch_returns(
|
||||
"XAUUSD", "2025-01-02", holding_days=5, benchmark="SPY"
|
||||
)
|
||||
|
||||
assert queried[0] == "GC=F" # stock symbol normalized (#984)
|
||||
assert queried[1] == "SPY" # benchmark left as the canonical symbol
|
||||
assert raw is not None and days is not None
|
||||
assert resolved == "2025-01-07" # resolution date recorded (#1251)
|
||||
|
||||
|
||||
def test_news_lookup_normalizes_symbol(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
class FakeTicker:
|
||||
def __init__(self, symbol):
|
||||
seen["symbol"] = symbol
|
||||
|
||||
def get_news(self, count):
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(ynews.yf, "Ticker", FakeTicker)
|
||||
monkeypatch.setattr(ynews, "yf_retry", lambda fn: fn())
|
||||
|
||||
out = ynews.get_news_yfinance("XAUUSD", "2025-01-01", "2025-01-10")
|
||||
|
||||
assert seen["symbol"] == "GC=F" # news queried with the canonical symbol
|
||||
assert "XAUUSD" in out # the user's ticker stays in the report
|
||||
assert "GC=F" in out # provenance noted
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Tests for symbol normalization and the no-data routing sentinel."""
|
||||
|
||||
import unittest
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.dataflows.errors import NoMarketDataError
|
||||
from tradingagents.dataflows.symbols import crypto_base, normalize_symbol
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestNormalizeSymbol(unittest.TestCase):
|
||||
def test_plain_equities_unchanged(self):
|
||||
for sym in ("AAPL", "MSFT", "TSM", "BRK.B", "0700.HK", "^GSPC", "GC=F"):
|
||||
self.assertEqual(normalize_symbol(sym), sym)
|
||||
|
||||
def test_lowercases_are_upper(self):
|
||||
self.assertEqual(normalize_symbol("aapl"), "AAPL")
|
||||
self.assertEqual(normalize_symbol(" msft "), "MSFT")
|
||||
|
||||
def test_metal_aliases_map_to_futures(self):
|
||||
self.assertEqual(normalize_symbol("XAUUSD"), "GC=F")
|
||||
self.assertEqual(normalize_symbol("XAUUSD+"), "GC=F") # broker CFD suffix
|
||||
self.assertEqual(normalize_symbol("xauusd+"), "GC=F")
|
||||
self.assertEqual(normalize_symbol("GOLD"), "GC=F")
|
||||
self.assertEqual(normalize_symbol("XAGUSD"), "SI=F")
|
||||
|
||||
def test_energy_and_index_aliases(self):
|
||||
self.assertEqual(normalize_symbol("USOIL"), "CL=F")
|
||||
self.assertEqual(normalize_symbol("SPX500"), "^GSPC")
|
||||
self.assertEqual(normalize_symbol("NAS100"), "^NDX")
|
||||
self.assertEqual(normalize_symbol("US30"), "^DJI")
|
||||
|
||||
def test_forex_pairs_get_x_suffix(self):
|
||||
self.assertEqual(normalize_symbol("EURUSD"), "EURUSD=X")
|
||||
self.assertEqual(normalize_symbol("GBPJPY"), "GBPJPY=X")
|
||||
self.assertEqual(normalize_symbol("eurusd"), "EURUSD=X")
|
||||
|
||||
def test_crypto_pairs_get_dash_usd(self):
|
||||
self.assertEqual(normalize_symbol("BTCUSD"), "BTC-USD")
|
||||
self.assertEqual(normalize_symbol("ETHUSD"), "ETH-USD")
|
||||
|
||||
def test_six_letter_non_currency_left_alone(self):
|
||||
# GOOGLE-style 6-letter tickers that aren't two currency codes
|
||||
# must not be mangled into a fake forex pair.
|
||||
self.assertEqual(normalize_symbol("ABCDEF"), "ABCDEF")
|
||||
|
||||
def test_empty_input_passthrough(self):
|
||||
self.assertEqual(normalize_symbol(""), "")
|
||||
|
||||
def test_hk_five_digit_code_repadded_to_four(self):
|
||||
# HKEX lists up to 5-digit codes; Yahoo only accepts 4 (#957).
|
||||
self.assertEqual(normalize_symbol("09992.HK"), "9992.HK")
|
||||
self.assertEqual(normalize_symbol("00700.HK"), "0700.HK")
|
||||
self.assertEqual(normalize_symbol("00001.HK"), "0001.HK")
|
||||
|
||||
def test_hk_four_digit_code_unchanged(self):
|
||||
self.assertEqual(normalize_symbol("0700.HK"), "0700.HK")
|
||||
self.assertEqual(normalize_symbol("9992.HK"), "9992.HK")
|
||||
self.assertEqual(normalize_symbol("80737.HK"), "80737.HK")
|
||||
|
||||
def test_hk_short_code_padded_to_four(self):
|
||||
self.assertEqual(normalize_symbol("700.HK"), "0700.HK")
|
||||
|
||||
def test_hk_code_case_insensitive_suffix(self):
|
||||
self.assertEqual(normalize_symbol("09992.hk"), "9992.HK")
|
||||
|
||||
def test_shanghai_sh_suffix_maps_to_yahoo_ss(self):
|
||||
self.assertEqual(normalize_symbol("600519.sh"), "600519.SS")
|
||||
self.assertEqual(normalize_symbol("600519.SS"), "600519.SS")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestNoMarketDataError(unittest.TestCase):
|
||||
def test_message_includes_resolution(self):
|
||||
err = NoMarketDataError("XAUUSD+", "GC=F", "no rows")
|
||||
self.assertIn("XAUUSD+", str(err))
|
||||
self.assertIn("GC=F", str(err))
|
||||
self.assertEqual(err.symbol, "XAUUSD+")
|
||||
self.assertEqual(err.canonical, "GC=F")
|
||||
|
||||
def test_canonical_defaults_to_symbol(self):
|
||||
err = NoMarketDataError("FOOBAR")
|
||||
self.assertEqual(err.canonical, "FOOBAR")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestCryptoBase(unittest.TestCase):
|
||||
def test_resolves_known_crypto_forms(self):
|
||||
for raw in ("BTC-USD", "BTCUSD", "btc-usdt", "BTC-USDC", "BTCUSD+"):
|
||||
self.assertEqual(crypto_base(raw), "BTC")
|
||||
self.assertEqual(crypto_base("ETH-USD"), "ETH")
|
||||
self.assertEqual(crypto_base("sol-usd"), "SOL")
|
||||
|
||||
def test_non_crypto_returns_none(self):
|
||||
# Plain equities, class shares, and real tickers that alias elsewhere
|
||||
# (GOLD -> gold future on the Yahoo path) must NOT read as crypto.
|
||||
for raw in ("AAPL", "BRK-B", "GOLD", "XYZ-USD", "EURUSD", "", None):
|
||||
self.assertIsNone(crypto_base(raw))
|
||||
|
||||
def test_agrees_with_normalize_symbol(self):
|
||||
# crypto_base is the shared primitive behind the -USD normalization.
|
||||
self.assertEqual(normalize_symbol("BTCUSD"), "BTC-USD")
|
||||
self.assertEqual(crypto_base("BTCUSD"), "BTC")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Tests for the configurable sampling temperature (#178/#168).
|
||||
|
||||
Temperature is a cross-provider knob: when set it must reach the underlying
|
||||
chat client; when unset the provider keeps its own default.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.llm_clients.factory import create_llm_client
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestTemperatureForwarding:
|
||||
@pytest.mark.parametrize(
|
||||
"provider,model",
|
||||
[
|
||||
# gpt-4.1 is intentionally a non-reasoning model: the GPT-5 family
|
||||
# are reasoning models and correctly drop temperature (see
|
||||
# test_openai_reasoning_effort), so forwarding is tested on gpt-4.1.
|
||||
("openai", "gpt-4.1"),
|
||||
("anthropic", "claude-sonnet-5"),
|
||||
("google", "gemini-3.5-flash"),
|
||||
("deepseek", "deepseek-chat"),
|
||||
],
|
||||
)
|
||||
def test_temperature_reaches_client_when_set(self, provider, model):
|
||||
llm = create_llm_client(
|
||||
provider=provider, model=model, temperature=0.0, api_key="placeholder"
|
||||
).get_llm()
|
||||
assert llm.temperature == 0.0
|
||||
|
||||
def test_temperature_omitted_leaves_provider_default(self):
|
||||
# Not passing temperature must not force it to a value.
|
||||
llm = create_llm_client(
|
||||
provider="openai", model="gpt-4.1", api_key="placeholder"
|
||||
).get_llm()
|
||||
# langchain's default is unset/None, not 0.0
|
||||
assert llm.temperature is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestTemperatureEnvOverlay:
|
||||
def test_env_sets_temperature(self, monkeypatch):
|
||||
import tradingagents.default_config as dc
|
||||
monkeypatch.setenv("TRADINGAGENTS_TEMPERATURE", "0.2")
|
||||
importlib.reload(dc)
|
||||
# Stored on config (string from env is fine; consumed via float()).
|
||||
assert dc.DEFAULT_CONFIG["temperature"] in ("0.2", 0.2)
|
||||
assert float(dc.DEFAULT_CONFIG["temperature"]) == 0.2
|
||||
monkeypatch.delenv("TRADINGAGENTS_TEMPERATURE", raising=False)
|
||||
importlib.reload(dc)
|
||||
|
||||
def test_default_temperature_is_none(self, monkeypatch):
|
||||
import tradingagents.default_config as dc
|
||||
monkeypatch.delenv("TRADINGAGENTS_TEMPERATURE", raising=False)
|
||||
importlib.reload(dc)
|
||||
assert dc.DEFAULT_CONFIG["temperature"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestProviderKwargsTemperature:
|
||||
"""build_llm_kwargs float-coerces and forwards temperature, or omits it."""
|
||||
|
||||
def _kwargs_for(self, temperature):
|
||||
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
|
||||
|
||||
def test_float_passthrough(self):
|
||||
assert self._kwargs_for(0.0)["temperature"] == 0.0
|
||||
|
||||
def test_none_omitted(self):
|
||||
assert "temperature" not in self._kwargs_for(None)
|
||||
|
||||
def test_empty_string_omitted(self):
|
||||
assert "temperature" not in self._kwargs_for("")
|
||||
@@ -0,0 +1,28 @@
|
||||
import unittest
|
||||
|
||||
import pytest
|
||||
|
||||
from cli.prompts import normalize_ticker_symbol
|
||||
from tradingagents.agents.context import build_instrument_context
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TickerSymbolHandlingTests(unittest.TestCase):
|
||||
def test_normalize_ticker_symbol_preserves_exchange_suffix(self):
|
||||
self.assertEqual(normalize_ticker_symbol(" cnc.to "), "CNC.TO")
|
||||
|
||||
def test_build_instrument_context_mentions_exact_symbol(self):
|
||||
context = build_instrument_context("7203.T")
|
||||
self.assertIn("7203.T", context)
|
||||
self.assertIn("exchange suffix", context)
|
||||
|
||||
def test_single_get_ticker_no_shadow(self):
|
||||
# A second get_ticker with an empty prompt (a bare "?") once shadowed
|
||||
# the descriptive one; the selection flow must use the one in prompts.
|
||||
import cli.prompts
|
||||
import cli.selections
|
||||
self.assertIs(cli.selections.get_ticker, cli.prompts.get_ticker)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user