You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
223 lines
7.3 KiB
Python
223 lines
7.3 KiB
Python
"""
|
|
Multi-source headline ingest — vendor payload → NewsHeadline rows.
|
|
|
|
For each article, one DB row per (source, source_id, ticker) is emitted so
|
|
ticker-indexed queries are O(log n). The full vendor symbol set is preserved
|
|
on `tickers_all`.
|
|
|
|
Dedup is enforced by `uq_news_headline_source_ticker` + ON CONFLICT DO NOTHING,
|
|
making re-ingest idempotent (daily backfill rerunning over an overlapping
|
|
window is safe).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
from datetime import datetime, timezone
|
|
from typing import Iterable
|
|
|
|
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.database import AsyncSessionLocal
|
|
from app.models.news_headline import NewsHeadline
|
|
from app.services.news.alpaca_news_client import AlpacaNewsArticle
|
|
from app.services.news.category_normalizer import (
|
|
normalize_alpaca,
|
|
normalize_finnhub,
|
|
normalize_stocktwits,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# NewsHeadline has ~17 insertable columns. asyncpg param ceiling 32,767.
|
|
# 32767 // 17 ≈ 1927; round down for safety + future column adds.
|
|
_CHUNK = 1500
|
|
|
|
|
|
async def insert_headline_rows(rows: list[dict], db: AsyncSession | None = None) -> int:
|
|
"""Insert pre-built rows with chunked ON CONFLICT DO NOTHING.
|
|
|
|
If `db` is provided, it's used directly and the caller commits.
|
|
Otherwise a short-lived session is opened and committed here.
|
|
"""
|
|
if not rows:
|
|
return 0
|
|
|
|
if db is not None:
|
|
return await _insert_chunked(db, rows)
|
|
|
|
async with AsyncSessionLocal() as sess:
|
|
inserted = await _insert_chunked(sess, rows)
|
|
await sess.commit()
|
|
return inserted
|
|
|
|
|
|
async def _insert_chunked(sess: AsyncSession, rows: list[dict]) -> int:
|
|
inserted = 0
|
|
for i in range(0, len(rows), _CHUNK):
|
|
stmt = pg_insert(NewsHeadline).values(rows[i : i + _CHUNK])
|
|
stmt = stmt.on_conflict_do_nothing(constraint="uq_news_headline_source_ticker")
|
|
result = await sess.execute(stmt)
|
|
inserted += result.rowcount or 0
|
|
await asyncio.sleep(0)
|
|
return inserted
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Vendor → row builders
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def alpaca_articles_to_rows(
|
|
articles: Iterable[AlpacaNewsArticle],
|
|
ingested_at: datetime | None = None,
|
|
) -> list[dict]:
|
|
"""Expand each Alpaca article into one row per ticker."""
|
|
now = ingested_at or datetime.now(timezone.utc)
|
|
out: list[dict] = []
|
|
for art in articles:
|
|
if not art.symbols:
|
|
continue
|
|
vendor_cats = [art.source] if art.source else []
|
|
unified = normalize_alpaca(vendor_cats, art.headline)
|
|
# Alpaca's first symbol is conventionally primary
|
|
primary = art.symbols[0].upper()
|
|
for sym in art.symbols:
|
|
sym_u = sym.strip().upper()
|
|
if not sym_u:
|
|
continue
|
|
out.append({
|
|
"source": "alpaca_benzinga",
|
|
"source_id": str(art.id),
|
|
"ticker": sym_u,
|
|
"tickers_all": [s.upper() for s in art.symbols if s],
|
|
"published_at": art.created_at,
|
|
"headline": art.headline,
|
|
"summary": art.summary,
|
|
"url": art.url,
|
|
"language": "en",
|
|
"vendor_categories": vendor_cats,
|
|
"categories": unified,
|
|
"raw_sentiment": None, # Alpaca free tier does not include sentiment
|
|
"is_primary": sym_u == primary,
|
|
"ingested_at": now,
|
|
})
|
|
return out
|
|
|
|
|
|
def finnhub_articles_to_rows(
|
|
ticker: str,
|
|
articles: Iterable[dict],
|
|
ingested_at: datetime | None = None,
|
|
) -> list[dict]:
|
|
"""Convert Finnhub /company-news payloads. Single-ticker per article."""
|
|
now = ingested_at or datetime.now(timezone.utc)
|
|
out: list[dict] = []
|
|
sym = ticker.strip().upper()
|
|
for art in articles:
|
|
if not art.get("id"):
|
|
continue
|
|
ts = art.get("datetime")
|
|
if ts is None:
|
|
continue
|
|
try:
|
|
published_at = datetime.fromtimestamp(int(ts), tz=timezone.utc)
|
|
except (TypeError, ValueError):
|
|
continue
|
|
headline = art.get("headline") or ""
|
|
category = art.get("category")
|
|
unified = normalize_finnhub(category, headline)
|
|
out.append({
|
|
"source": "finnhub",
|
|
"source_id": str(art["id"]),
|
|
"ticker": sym,
|
|
"tickers_all": [sym],
|
|
"published_at": published_at,
|
|
"headline": headline,
|
|
"summary": art.get("summary"),
|
|
"url": art.get("url"),
|
|
"language": "en",
|
|
"vendor_categories": [category] if category else [],
|
|
"categories": unified,
|
|
"raw_sentiment": None,
|
|
"is_primary": True,
|
|
"ingested_at": now,
|
|
})
|
|
return out
|
|
|
|
|
|
def stocktwits_messages_to_rows(
|
|
ticker: str,
|
|
messages: Iterable[dict],
|
|
ingested_at: datetime | None = None,
|
|
) -> list[dict]:
|
|
"""
|
|
Convert StockTwits stream messages (streams/symbol/{ticker}.json).
|
|
|
|
Sentiment tag is the message's entities.sentiment.basic in {"Bullish", "Bearish"}.
|
|
Mapped to raw_sentiment ∈ {+1, -1}; absent tag → NULL (neutral).
|
|
"""
|
|
now = ingested_at or datetime.now(timezone.utc)
|
|
out: list[dict] = []
|
|
sym = ticker.strip().upper()
|
|
for msg in messages:
|
|
if not msg.get("id"):
|
|
continue
|
|
ts = msg.get("created_at")
|
|
if not ts:
|
|
continue
|
|
try:
|
|
published_at = _parse_st_iso(ts)
|
|
except Exception:
|
|
continue
|
|
body = (msg.get("body") or "").strip()
|
|
if not body:
|
|
continue
|
|
sentiment = _extract_st_sentiment(msg)
|
|
unified = normalize_stocktwits(body)
|
|
# Collect all symbol tags for tickers_all
|
|
all_syms = [s.get("symbol", "").upper() for s in (msg.get("symbols") or []) if s.get("symbol")]
|
|
if sym not in all_syms:
|
|
all_syms.insert(0, sym)
|
|
out.append({
|
|
"source": "stocktwits",
|
|
"source_id": str(msg["id"]),
|
|
"ticker": sym,
|
|
"tickers_all": all_syms,
|
|
"published_at": published_at,
|
|
"headline": body[:300], # truncate for headline; full body in summary
|
|
"summary": body,
|
|
"url": f"https://stocktwits.com/message/{msg['id']}",
|
|
"language": "en",
|
|
"vendor_categories": [sentiment] if sentiment else [],
|
|
"categories": unified,
|
|
"raw_sentiment": _sentiment_to_score(sentiment),
|
|
"is_primary": all_syms[0] == sym if all_syms else True,
|
|
"ingested_at": now,
|
|
})
|
|
return out
|
|
|
|
|
|
def _parse_st_iso(s: str) -> datetime:
|
|
if s.endswith("Z"):
|
|
s = s[:-1] + "+00:00"
|
|
dt = datetime.fromisoformat(s)
|
|
if dt.tzinfo is None:
|
|
dt = dt.replace(tzinfo=timezone.utc)
|
|
return dt
|
|
|
|
|
|
def _extract_st_sentiment(msg: dict) -> str | None:
|
|
entities = msg.get("entities") or {}
|
|
sentiment = entities.get("sentiment") or {}
|
|
return sentiment.get("basic") # "Bullish" | "Bearish" | None
|
|
|
|
|
|
def _sentiment_to_score(tag: str | None) -> float | None:
|
|
if tag == "Bullish":
|
|
return 1.0
|
|
if tag == "Bearish":
|
|
return -1.0
|
|
return None
|