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.

192 lines
6.4 KiB
Python

"""
Feature Materializer — computes derived attention features from wiki + gdelt raw data.
Materializes into attention_features_daily for the given ticker + event_date.
"""
import logging
import statistics
from datetime import date, timedelta
from typing import Optional
from sqlalchemy import select, func
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.attention import (
AttentionFeaturesDaily,
CompanyEntityMap,
GdeltArticleRaw,
WikiPageviewsDaily,
)
logger = logging.getLogger(__name__)
async def _get_wiki_views_series(
db: AsyncSession,
wiki_title: str,
target_date: date,
lookback: int,
) -> list[int]:
"""Return list of view counts for [target_date - lookback, target_date - 1]."""
start = target_date - timedelta(days=lookback)
end = target_date - timedelta(days=1)
result = await db.execute(
select(WikiPageviewsDaily.views)
.where(
WikiPageviewsDaily.wiki_title == wiki_title,
WikiPageviewsDaily.date >= start,
WikiPageviewsDaily.date <= end,
)
.order_by(WikiPageviewsDaily.date)
)
return list(result.scalars().all())
async def materialize_features(
db: AsyncSession,
ticker: str,
event_date: date,
) -> AttentionFeaturesDaily:
"""Compute and upsert daily attention features for ticker on event_date.
Returns the upserted AttentionFeaturesDaily row.
"""
ticker = ticker.upper()
entity_result = await db.execute(
select(CompanyEntityMap).where(CompanyEntityMap.ticker == ticker)
)
entity = entity_result.scalars().first()
wiki_title = entity.wiki_title if entity else None
# --- Wiki features ---
wiki_views: Optional[int] = None
wiki_spike_10d: Optional[float] = None
wiki_zscore_20d: Optional[float] = None
if wiki_title:
# Fetch today's views
today_result = await db.execute(
select(WikiPageviewsDaily.views).where(
WikiPageviewsDaily.wiki_title == wiki_title,
WikiPageviewsDaily.date == event_date,
)
)
wiki_views = today_result.scalars().first()
if wiki_views is not None:
# Spike: views / median of previous 10 days
prev_10 = await _get_wiki_views_series(db, wiki_title, event_date, 10)
if prev_10:
median_10 = statistics.median(prev_10)
wiki_spike_10d = wiki_views / median_10 if median_10 > 0 else None
# Z-score over previous 20 days
prev_20 = await _get_wiki_views_series(db, wiki_title, event_date, 20)
if len(prev_20) >= 3:
mean_20 = statistics.mean(prev_20)
stdev_20 = statistics.stdev(prev_20)
if stdev_20 > 0:
wiki_zscore_20d = (wiki_views - mean_20) / stdev_20
# --- GDELT features ---
# 1-day: articles published on event_date only
day_start = event_date
day_end = event_date
count_1d_result = await db.execute(
select(func.count()).where(
GdeltArticleRaw.matched_ticker == ticker,
func.date(GdeltArticleRaw.published_at) >= day_start,
func.date(GdeltArticleRaw.published_at) <= day_end,
)
)
gdelt_article_count_1d = count_1d_result.scalar() or 0
# 3-day window: event_date ± 1
window_start = event_date - timedelta(days=1)
window_end = event_date + timedelta(days=1)
count_3d_result = await db.execute(
select(func.count()).where(
GdeltArticleRaw.matched_ticker == ticker,
func.date(GdeltArticleRaw.published_at) >= window_start,
func.date(GdeltArticleRaw.published_at) <= window_end,
)
)
gdelt_article_count_3d = count_3d_result.scalar() or 0
# Unique domains in 3-day window
domains_result = await db.execute(
select(func.count(func.distinct(GdeltArticleRaw.domain))).where(
GdeltArticleRaw.matched_ticker == ticker,
GdeltArticleRaw.domain.isnot(None),
func.date(GdeltArticleRaw.published_at) >= window_start,
func.date(GdeltArticleRaw.published_at) <= window_end,
)
)
gdelt_unique_domains_3d = domains_result.scalar() or 0
# US articles in 3-day window
us_result = await db.execute(
select(func.count()).where(
GdeltArticleRaw.matched_ticker == ticker,
GdeltArticleRaw.sourcecountry == "US",
func.date(GdeltArticleRaw.published_at) >= window_start,
func.date(GdeltArticleRaw.published_at) <= window_end,
)
)
gdelt_us_article_count_3d = us_result.scalar() or 0
# Upsert
stmt = (
pg_insert(AttentionFeaturesDaily)
.values(
ticker=ticker,
date=event_date,
wiki_views=wiki_views,
wiki_spike_10d=wiki_spike_10d,
wiki_zscore_20d=wiki_zscore_20d,
gdelt_article_count_1d=gdelt_article_count_1d,
gdelt_article_count_3d=gdelt_article_count_3d,
gdelt_unique_domains_3d=gdelt_unique_domains_3d,
gdelt_us_article_count_3d=gdelt_us_article_count_3d,
)
.on_conflict_do_update(
constraint="uq_attention_features_daily",
set_=dict(
wiki_views=wiki_views,
wiki_spike_10d=wiki_spike_10d,
wiki_zscore_20d=wiki_zscore_20d,
gdelt_article_count_1d=gdelt_article_count_1d,
gdelt_article_count_3d=gdelt_article_count_3d,
gdelt_unique_domains_3d=gdelt_unique_domains_3d,
gdelt_us_article_count_3d=gdelt_us_article_count_3d,
),
)
.returning(AttentionFeaturesDaily)
)
await db.execute(stmt)
await db.commit()
# Always re-fetch with populate_existing=True to avoid SQLAlchemy identity-map
# returning stale pre-upsert values (same pattern as resolve_entity).
row_result = await db.execute(
select(AttentionFeaturesDaily)
.where(
AttentionFeaturesDaily.ticker == ticker,
AttentionFeaturesDaily.date == event_date,
)
.execution_options(populate_existing=True)
)
row = row_result.scalars().first()
logger.info(
"Materialized features for %s on %s: wiki=%s gdelt_1d=%d gdelt_3d=%d",
ticker, event_date, wiki_views, gdelt_article_count_1d, gdelt_article_count_3d,
)
return row