mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-27 06:56:39 +03:00
feat(sentiment): screen social posts with TypeSafe's Jev (#1376)
- with TYPESAFE_API_KEY set, each StockTwits and Reddit post is asked whether it is about the company and its stance on the stock - posts clearly about something else are dropped before the per-source cut, and each block opens with a stance count - any failed request leaves the source's posts unscreened and says so; without a key nothing changes
This commit is contained in:
@@ -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
|
||||
@@ -271,3 +271,43 @@ def test_empty_subreddit_on_a_full_page_is_not_called_empty():
|
||||
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
|
||||
|
||||
@@ -8,6 +8,7 @@ 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
|
||||
|
||||
@@ -75,3 +76,41 @@ class TestStockTwitsCryptoSymbols:
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user