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.
393 lines
13 KiB
Python
393 lines
13 KiB
Python
"""
|
|
News v2 endpoints — multi-source headline + session aggregate API.
|
|
|
|
Mounted at `/api/v1/news/v2/*` to avoid conflict with the legacy
|
|
`GET /news/{ticker}` aggregator (which uses a wildcard path).
|
|
|
|
Endpoints:
|
|
GET /headlines raw headline rows
|
|
GET /session_aggregate single (ticker, session_date, window)
|
|
POST /session_aggregate/batch many tickers in one shot
|
|
GET /coverage per-source ingest depth probe
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
from datetime import date, datetime, timezone
|
|
from typing import Optional
|
|
|
|
from fastapi import APIRouter, Body, Depends, HTTPException, Query
|
|
from fastapi.responses import Response
|
|
from pydantic import BaseModel, Field, field_validator
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.database import get_db
|
|
from app.models.news_headline import NewsHeadline
|
|
from app.services.news.session_aggregator import (
|
|
aggregate_session,
|
|
aggregate_session_batch,
|
|
)
|
|
from app.services.news.session_window import session_window
|
|
from app.utils.cache import with_cache
|
|
|
|
logger = logging.getLogger("app.api.v1.news_v2")
|
|
|
|
router = APIRouter()
|
|
|
|
# Concurrency gates — prevent event-loop saturation under bulk scans
|
|
_HEADLINES_SEMAPHORE = asyncio.Semaphore(8)
|
|
_AGG_SEMAPHORE = asyncio.Semaphore(8)
|
|
_SEMAPHORE_WAIT_TIMEOUT = 10.0
|
|
|
|
_VALID_SOURCES = {"alpaca_benzinga", "stocktwits", "finnhub", "gdelt"}
|
|
_VALID_WINDOWS = {"premarket", "intraday", "post", "full_session"}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Schemas
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class HeadlineItem(BaseModel):
|
|
source: str
|
|
source_id: str
|
|
ticker: str
|
|
tickers_all: list[str] | None = None
|
|
published_at: str
|
|
headline: str
|
|
summary: str | None = None
|
|
url: str | None = None
|
|
language: str | None = None
|
|
vendor_categories: list[str] | None = None
|
|
categories: list[str] | None = None
|
|
raw_sentiment: float | None = None
|
|
is_primary: bool
|
|
ingested_at: str
|
|
|
|
|
|
class HeadlinesResponse(BaseModel):
|
|
items: list[HeadlineItem]
|
|
next_cursor: str | None = None
|
|
|
|
|
|
class SocialStatsItem(BaseModel):
|
|
message_count: int = 0
|
|
bull_count: int = 0
|
|
bear_count: int = 0
|
|
bull_bear_ratio: float | None = None
|
|
|
|
|
|
class SessionAggregateItem(BaseModel):
|
|
ticker: str
|
|
session_date: str
|
|
window: str
|
|
headline_count: int = 0
|
|
primary_count: int = 0
|
|
first_headline_at: str | None = None
|
|
last_headline_at: str | None = None
|
|
category_counts: dict[str, int] = Field(default_factory=dict)
|
|
sentiment_mean: float | None = None
|
|
sentiment_recency_weighted: float | None = None
|
|
social: SocialStatsItem = Field(default_factory=SocialStatsItem)
|
|
sources_present: list[str] = Field(default_factory=list)
|
|
|
|
|
|
class SessionAggregateBatchRequest(BaseModel):
|
|
session_date: date
|
|
window: str
|
|
symbols: list[str]
|
|
sources: list[str] | None = None
|
|
|
|
@field_validator("window")
|
|
@classmethod
|
|
def _validate_window(cls, v: str) -> str:
|
|
if v not in _VALID_WINDOWS:
|
|
raise ValueError(f"window must be one of {sorted(_VALID_WINDOWS)}")
|
|
return v
|
|
|
|
@field_validator("symbols")
|
|
@classmethod
|
|
def _validate_symbols(cls, v: list[str]) -> list[str]:
|
|
if not v:
|
|
raise ValueError("symbols must not be empty")
|
|
if len(v) > 200:
|
|
raise ValueError("symbols max 200 per request")
|
|
return v
|
|
|
|
|
|
class SessionAggregateBatchResponse(BaseModel):
|
|
items: dict[str, SessionAggregateItem]
|
|
|
|
|
|
class CoverageResponse(BaseModel):
|
|
source: str
|
|
symbol: str | None = None
|
|
earliest: str | None = None
|
|
latest: str | None = None
|
|
ingested_count: int
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# A. /headlines — raw rows
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@router.get(
|
|
"/headlines",
|
|
response_model=HeadlinesResponse,
|
|
summary="Raw multi-source news headlines",
|
|
description=(
|
|
"Multi-source raw headline rows. Filter by symbols, time window, and "
|
|
"source. Sources: `alpaca_benzinga`, `stocktwits`, `finnhub`, `gdelt`."
|
|
),
|
|
)
|
|
@with_cache(
|
|
namespace="news_v2:headlines",
|
|
ttl=300,
|
|
key_params=["symbols", "start", "end", "sources", "limit", "cursor"],
|
|
)
|
|
async def get_headlines(
|
|
response: Response,
|
|
symbols: Optional[str] = Query(None, description="CSV ticker list, max 50 (e.g. AAPL,MSFT)"),
|
|
start: Optional[datetime] = Query(None, description="Start time (UTC ISO)"),
|
|
end: Optional[datetime] = Query(None, description="End time (UTC ISO)"),
|
|
sources: Optional[str] = Query(None, description=f"CSV source filter, subset of {sorted(_VALID_SOURCES)}"),
|
|
limit: int = Query(100, ge=1, le=500),
|
|
cursor: Optional[str] = Query(None, description="published_at_lt cursor (ISO datetime)"),
|
|
force_refresh: bool = Query(False),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
sym_list = _csv_to_list(symbols, max_items=50, kind="symbols")
|
|
src_list = _csv_to_list(sources, max_items=10, kind="sources")
|
|
if src_list:
|
|
bad = [s for s in src_list if s not in _VALID_SOURCES]
|
|
if bad:
|
|
raise HTTPException(400, f"Unknown sources: {bad}")
|
|
|
|
cursor_dt = _parse_cursor(cursor)
|
|
|
|
try:
|
|
await asyncio.wait_for(_HEADLINES_SEMAPHORE.acquire(), timeout=_SEMAPHORE_WAIT_TIMEOUT)
|
|
except asyncio.TimeoutError:
|
|
raise HTTPException(429, "Server busy — try again later")
|
|
try:
|
|
stmt = select(NewsHeadline).order_by(NewsHeadline.published_at.desc()).limit(limit)
|
|
if sym_list:
|
|
stmt = stmt.where(NewsHeadline.ticker.in_([s.upper() for s in sym_list]))
|
|
if start is not None:
|
|
stmt = stmt.where(NewsHeadline.published_at >= _ensure_utc(start))
|
|
if end is not None:
|
|
stmt = stmt.where(NewsHeadline.published_at < _ensure_utc(end))
|
|
if src_list:
|
|
stmt = stmt.where(NewsHeadline.source.in_(src_list))
|
|
if cursor_dt is not None:
|
|
stmt = stmt.where(NewsHeadline.published_at < cursor_dt)
|
|
|
|
result = await db.execute(stmt)
|
|
rows = result.scalars().all()
|
|
finally:
|
|
_HEADLINES_SEMAPHORE.release()
|
|
|
|
items = [_row_to_item(r) for r in rows]
|
|
next_cursor = items[-1].published_at if len(items) == limit else None
|
|
return HeadlinesResponse(items=items, next_cursor=next_cursor)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# B. /session_aggregate — single ticker
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@router.get(
|
|
"/session_aggregate",
|
|
response_model=SessionAggregateItem,
|
|
summary="Session-aggregated news for one ticker",
|
|
)
|
|
@with_cache(
|
|
namespace="news_v2:session_agg",
|
|
ttl=600,
|
|
key_params=["symbol", "session_date", "window", "sources"],
|
|
)
|
|
async def get_session_aggregate(
|
|
response: Response,
|
|
symbol: str = Query(..., description="Ticker symbol"),
|
|
session_date: date = Query(..., description="ET session date (YYYY-MM-DD)"),
|
|
window: str = Query("premarket", description=f"One of {sorted(_VALID_WINDOWS)}"),
|
|
sources: Optional[str] = Query(None, description=f"CSV source filter, subset of {sorted(_VALID_SOURCES)}"),
|
|
force_refresh: bool = Query(False),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
if window not in _VALID_WINDOWS:
|
|
raise HTTPException(400, f"window must be one of {sorted(_VALID_WINDOWS)}")
|
|
|
|
src_list = _csv_to_list(sources, max_items=10, kind="sources")
|
|
if src_list:
|
|
bad = [s for s in src_list if s not in _VALID_SOURCES]
|
|
if bad:
|
|
raise HTTPException(400, f"Unknown sources: {bad}")
|
|
|
|
try:
|
|
await asyncio.wait_for(_AGG_SEMAPHORE.acquire(), timeout=_SEMAPHORE_WAIT_TIMEOUT)
|
|
except asyncio.TimeoutError:
|
|
raise HTTPException(429, "Server busy — try again later")
|
|
try:
|
|
try:
|
|
agg = await aggregate_session(
|
|
db=db,
|
|
ticker=symbol,
|
|
session_date=session_date,
|
|
window=window, # type: ignore[arg-type]
|
|
sources=src_list,
|
|
)
|
|
except ValueError as e:
|
|
raise HTTPException(400, str(e))
|
|
finally:
|
|
_AGG_SEMAPHORE.release()
|
|
|
|
return SessionAggregateItem(**agg.to_dict())
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# C. POST /session_aggregate/batch — many tickers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@router.post(
|
|
"/session_aggregate/batch",
|
|
response_model=SessionAggregateBatchResponse,
|
|
summary="Session-aggregated news for many tickers in one call",
|
|
description=(
|
|
"Batch variant. Caching is intentionally NOT applied at this layer — "
|
|
"fithia2 maintains a client-side disk cache as the primary defense; "
|
|
"Oracle absorbs only burst load. Use the GET single endpoint for "
|
|
"Redis-cached single-ticker reads."
|
|
),
|
|
)
|
|
async def post_session_aggregate_batch(
|
|
payload: SessionAggregateBatchRequest = Body(...),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
src_list = payload.sources or None
|
|
if src_list:
|
|
bad = [s for s in src_list if s not in _VALID_SOURCES]
|
|
if bad:
|
|
raise HTTPException(400, f"Unknown sources: {bad}")
|
|
|
|
try:
|
|
await asyncio.wait_for(_AGG_SEMAPHORE.acquire(), timeout=_SEMAPHORE_WAIT_TIMEOUT)
|
|
except asyncio.TimeoutError:
|
|
raise HTTPException(429, "Server busy — try again later")
|
|
try:
|
|
try:
|
|
results = await aggregate_session_batch(
|
|
db=db,
|
|
tickers=payload.symbols,
|
|
session_date=payload.session_date,
|
|
window=payload.window, # type: ignore[arg-type]
|
|
sources=src_list,
|
|
)
|
|
except ValueError as e:
|
|
raise HTTPException(400, str(e))
|
|
finally:
|
|
_AGG_SEMAPHORE.release()
|
|
|
|
return SessionAggregateBatchResponse(
|
|
items={t: SessionAggregateItem(**a.to_dict()) for t, a in results.items()}
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# D. /coverage — per-source ingest depth
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@router.get(
|
|
"/coverage",
|
|
response_model=CoverageResponse,
|
|
summary="Per-source ingest coverage probe",
|
|
)
|
|
@with_cache(
|
|
namespace="news_v2:coverage",
|
|
ttl=300,
|
|
key_params=["source", "symbol"],
|
|
)
|
|
async def get_coverage(
|
|
response: Response,
|
|
source: str = Query(..., description=f"One of {sorted(_VALID_SOURCES)}"),
|
|
symbol: Optional[str] = Query(None, description="Optional ticker filter"),
|
|
force_refresh: bool = Query(False),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
if source not in _VALID_SOURCES:
|
|
raise HTTPException(400, f"Unknown source: {source}")
|
|
|
|
stmt = select(
|
|
func.min(NewsHeadline.published_at),
|
|
func.max(NewsHeadline.published_at),
|
|
func.count(NewsHeadline.id),
|
|
).where(NewsHeadline.source == source)
|
|
if symbol:
|
|
stmt = stmt.where(NewsHeadline.ticker == symbol.strip().upper())
|
|
|
|
result = await db.execute(stmt)
|
|
row = result.one()
|
|
earliest, latest, count = row
|
|
|
|
return CoverageResponse(
|
|
source=source,
|
|
symbol=symbol.strip().upper() if symbol else None,
|
|
earliest=earliest.isoformat() if earliest else None,
|
|
latest=latest.isoformat() if latest else None,
|
|
ingested_count=int(count or 0),
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _csv_to_list(s: str | None, *, max_items: int, kind: str) -> list[str]:
|
|
if not s:
|
|
return []
|
|
items = [x.strip() for x in s.split(",") if x.strip()]
|
|
if len(items) > max_items:
|
|
raise HTTPException(400, f"{kind} max {max_items} per request")
|
|
return items
|
|
|
|
|
|
def _parse_cursor(cursor: str | None) -> datetime | None:
|
|
if not cursor:
|
|
return None
|
|
try:
|
|
if cursor.endswith("Z"):
|
|
cursor = cursor[:-1] + "+00:00"
|
|
dt = datetime.fromisoformat(cursor)
|
|
if dt.tzinfo is None:
|
|
dt = dt.replace(tzinfo=timezone.utc)
|
|
return dt
|
|
except Exception:
|
|
raise HTTPException(400, "cursor must be ISO datetime")
|
|
|
|
|
|
def _ensure_utc(dt: datetime) -> datetime:
|
|
if dt.tzinfo is None:
|
|
return dt.replace(tzinfo=timezone.utc)
|
|
return dt.astimezone(timezone.utc)
|
|
|
|
|
|
def _row_to_item(r: NewsHeadline) -> HeadlineItem:
|
|
return HeadlineItem(
|
|
source=r.source,
|
|
source_id=r.source_id,
|
|
ticker=r.ticker,
|
|
tickers_all=list(r.tickers_all) if r.tickers_all else None,
|
|
published_at=_ensure_utc(r.published_at).isoformat(),
|
|
headline=r.headline,
|
|
summary=r.summary,
|
|
url=r.url,
|
|
language=r.language,
|
|
vendor_categories=list(r.vendor_categories) if r.vendor_categories else None,
|
|
categories=list(r.categories) if r.categories else None,
|
|
raw_sentiment=r.raw_sentiment,
|
|
is_primary=bool(r.is_primary),
|
|
ingested_at=_ensure_utc(r.ingested_at).isoformat(),
|
|
)
|