From 7f28266bd635ad7bfb3e55ce35b70391b43e8575 Mon Sep 17 00:00:00 2001 From: I Luk Kim Date: Mon, 20 Apr 2026 16:03:40 -0700 Subject: [PATCH] =?UTF-8?q?feat:=20company=20metadata=20=EC=95=88=EC=A0=95?= =?UTF-8?q?=ED=99=94=20=E2=80=94=20=EC=8B=A4=EC=A0=9C=20sector/industry/ex?= =?UTF-8?q?change=20=EB=B3=B4=EA=B0=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CompanyMetadataService 신규: Redis → registry → yfinance 순서 조회, semaphore(5) 보호, 24h cache, companies + registry UPSERT - GET /company/{ticker} + POST /company/bulk (최대 100개, 세션 격리) - companies 테이블에 exchange/country/market_cap 컬럼 추가 (alembic) - CompanyInfo 스키마 exchange/country/market_cap 필드 추가 - _is_placeholder() 체크: 기존 "Technology/Software/XYZ Corporation" 행 재보강 - financial endpoint: financials/price 실패 시 company block은 유지 (500 방지) - bulk endpoint: 각 ticker마다 독립 AsyncSession으로 동시 세션 충돌 방지 - scripts/backfill_registry_sector.py: 9376개 NULL-sector 일괄 보강 스크립트 검증: AU→Basic Materials, USAS→Basic Materials, CPRX→Healthcare, HE→Utilities, ACHR→Industrials, ZZZZZZ→404, AAPL 기존 데이터 유지 Co-Authored-By: Claude Sonnet 4.6 --- ...6g7_companies_add_exchange_country_mcap.py | 26 ++ app/api/v1/api.py | 5 +- app/api/v1/endpoints/company.py | 96 +++++++ app/api/v1/endpoints/financial.py | 5 +- app/models/financial.py | 3 + app/schemas/company.py | 35 +++ app/schemas/financial.py | 7 +- app/services/company_metadata_service.py | 253 ++++++++++++++++++ app/services/financial_service.py | 148 +++++----- scripts/backfill_registry_sector.py | 93 +++++++ 10 files changed, 598 insertions(+), 73 deletions(-) create mode 100644 alembic/versions/k2c3d4e5f6g7_companies_add_exchange_country_mcap.py create mode 100644 app/api/v1/endpoints/company.py create mode 100644 app/schemas/company.py create mode 100644 app/services/company_metadata_service.py create mode 100644 scripts/backfill_registry_sector.py diff --git a/alembic/versions/k2c3d4e5f6g7_companies_add_exchange_country_mcap.py b/alembic/versions/k2c3d4e5f6g7_companies_add_exchange_country_mcap.py new file mode 100644 index 0000000..13b8c63 --- /dev/null +++ b/alembic/versions/k2c3d4e5f6g7_companies_add_exchange_country_mcap.py @@ -0,0 +1,26 @@ +"""companies: add exchange, country, market_cap columns + +Revision ID: k2c3d4e5f6g7 +Revises: j1b2c3d4e5f6 +Create Date: 2026-04-20 + +""" +from alembic import op +import sqlalchemy as sa + +revision = 'k2c3d4e5f6g7' +down_revision = 'j1b2c3d4e5f6' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column('companies', sa.Column('exchange', sa.String(20), nullable=True)) + op.add_column('companies', sa.Column('country', sa.String(100), nullable=True)) + op.add_column('companies', sa.Column('market_cap', sa.Float(), nullable=True)) + + +def downgrade() -> None: + op.drop_column('companies', 'market_cap') + op.drop_column('companies', 'country') + op.drop_column('companies', 'exchange') diff --git a/app/api/v1/api.py b/app/api/v1/api.py index ff1c090..1de9ef8 100644 --- a/app/api/v1/api.py +++ b/app/api/v1/api.py @@ -3,7 +3,7 @@ API v1 router """ from fastapi import APIRouter -from app.api.v1.endpoints import financial, price, catalog, health, migration, database, error_logs, request_logs, news, etf, stocks, fred, filings, alpaca, finra, overlay, screener, attention, insider, earnings, universe, dividends +from app.api.v1.endpoints import financial, price, catalog, health, migration, database, error_logs, request_logs, news, etf, stocks, fred, filings, alpaca, finra, overlay, screener, attention, insider, earnings, universe, dividends, company api_router = APIRouter() @@ -30,4 +30,5 @@ api_router.include_router(attention.router, prefix="/attention", tags=["attentio api_router.include_router(insider.router, prefix="/insider", tags=["insider"]) api_router.include_router(earnings.router, prefix="/earnings", tags=["earnings"]) api_router.include_router(universe.router, prefix="/universe", tags=["universe"]) -api_router.include_router(dividends.router, prefix="/dividends", tags=["dividends"]) \ No newline at end of file +api_router.include_router(dividends.router, prefix="/dividends", tags=["dividends"]) +api_router.include_router(company.router, prefix="/company", tags=["company"]) \ No newline at end of file diff --git a/app/api/v1/endpoints/company.py b/app/api/v1/endpoints/company.py new file mode 100644 index 0000000..bfe0271 --- /dev/null +++ b/app/api/v1/endpoints/company.py @@ -0,0 +1,96 @@ +""" +Company metadata endpoints — sector/industry/exchange/market_cap for any ticker. +""" + +import asyncio +import logging +from typing import List + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.database import AsyncSessionLocal, get_db +from app.schemas.company import ( + BulkCompanyItem, + BulkCompanyRequest, + BulkCompanyResponse, + CompanyMetadataResponse, +) +from app.services import company_metadata_service as cms + +logger = logging.getLogger(__name__) + +router = APIRouter() + +_BULK_MAX = 100 + + +@router.get( + "/{ticker}", + response_model=CompanyMetadataResponse, + summary="Get company metadata", + description=( + "Returns sector, industry, exchange, market_cap, country, and other metadata " + "for a ticker. Valid tickers without financial statements still return 200. " + "Unknown tickers return 404." + ), +) +async def get_company( + ticker: str, + db: AsyncSession = Depends(get_db), +): + ticker = ticker.upper() + try: + data = await cms.get_metadata(db, ticker) + except ValueError as e: + if "invalid ticker" in str(e).lower(): + raise HTTPException(status_code=404, detail=f"Unknown ticker: {ticker}") + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + logger.error("company metadata error for %s: %s", ticker, e) + raise HTTPException(status_code=500, detail="Failed to retrieve company metadata") + + return CompanyMetadataResponse(**data) + + +@router.post( + "/bulk", + response_model=BulkCompanyResponse, + summary="Bulk company metadata", + description=( + f"Fetch metadata for up to {_BULK_MAX} tickers in one request. " + "Partial failures are allowed — each item has either `data` or `error`." + ), +) +async def bulk_company( + request: BulkCompanyRequest, +): + tickers = [t.upper() for t in request.tickers] + if not tickers: + raise HTTPException(status_code=400, detail="tickers list is empty") + if len(tickers) > _BULK_MAX: + raise HTTPException( + status_code=400, + detail=f"Too many tickers: max {_BULK_MAX}, got {len(tickers)}", + ) + + async def _fetch_one(ticker: str) -> BulkCompanyItem: + async with AsyncSessionLocal() as session: + try: + data = await cms.get_metadata(session, ticker) + return BulkCompanyItem(ticker=ticker, data=CompanyMetadataResponse(**data)) + except ValueError as e: + return BulkCompanyItem(ticker=ticker, error=str(e)) + except Exception as e: + logger.warning("bulk company error for %s: %s", ticker, e) + return BulkCompanyItem(ticker=ticker, error="lookup failed") + + results: List[BulkCompanyItem] = await asyncio.gather(*[_fetch_one(t) for t in tickers]) + + success_count = sum(1 for r in results if r.data is not None) + return BulkCompanyResponse( + results=results, + total=len(results), + success_count=success_count, + error_count=len(results) - success_count, + ) diff --git a/app/api/v1/endpoints/financial.py b/app/api/v1/endpoints/financial.py index ac6d127..43bfa37 100644 --- a/app/api/v1/endpoints/financial.py +++ b/app/api/v1/endpoints/financial.py @@ -167,9 +167,12 @@ async def get_financial_data( ticker=company.ticker, name=company.name, cik=company.cik, + exchange=getattr(company, "exchange", None), sector=company.sector, industry=company.industry, - business_description=company.business_description + country=getattr(company, "country", None), + market_cap=getattr(company, "market_cap", None), + business_description=company.business_description, ) # Filter financial data by period type diff --git a/app/models/financial.py b/app/models/financial.py index a6ba573..5f1c767 100644 --- a/app/models/financial.py +++ b/app/models/financial.py @@ -15,8 +15,11 @@ class Company(Base): ticker = Column(String(10), unique=True, nullable=False, index=True) name = Column(String(255), nullable=False) cik = Column(String(20), unique=True, nullable=True) + exchange = Column(String(20), nullable=True) sector = Column(String(100), nullable=True) industry = Column(String(100), nullable=True) + country = Column(String(100), nullable=True) + market_cap = Column(Float, nullable=True) business_description = Column(String, nullable=True) created_at = Column(TIMESTAMP(timezone=True), default=lambda: datetime.now(timezone.utc)) updated_at = Column(TIMESTAMP(timezone=True), default=lambda: datetime.now(timezone.utc), onupdate=lambda: datetime.now(timezone.utc)) diff --git a/app/schemas/company.py b/app/schemas/company.py new file mode 100644 index 0000000..c454d19 --- /dev/null +++ b/app/schemas/company.py @@ -0,0 +1,35 @@ +""" +Company metadata schemas for the dedicated /company endpoints. +""" + +from typing import Any, Dict, List, Optional +from pydantic import BaseModel + + +class CompanyMetadataResponse(BaseModel): + ticker: str + name: Optional[str] = None + cik: Optional[str] = None + exchange: Optional[str] = None + sector: Optional[str] = None + industry: Optional[str] = None + country: Optional[str] = None + market_cap: Optional[float] = None + business_description: Optional[str] = None + + +class BulkCompanyRequest(BaseModel): + tickers: List[str] + + +class BulkCompanyItem(BaseModel): + ticker: str + data: Optional[CompanyMetadataResponse] = None + error: Optional[str] = None + + +class BulkCompanyResponse(BaseModel): + results: List[BulkCompanyItem] + total: int + success_count: int + error_count: int diff --git a/app/schemas/financial.py b/app/schemas/financial.py index 80085ba..c00644d 100644 --- a/app/schemas/financial.py +++ b/app/schemas/financial.py @@ -282,12 +282,15 @@ class BulkPriceDataRequest(BaseModel): # Response Schemas class CompanyInfo(BaseModel): model_config = ConfigDict(from_attributes=True) - + ticker: str - name: str + name: Optional[str] = None cik: Optional[str] = None + exchange: Optional[str] = None sector: Optional[str] = None industry: Optional[str] = None + country: Optional[str] = None + market_cap: Optional[float] = None business_description: Optional[str] = None class FinancialDataPoint(BaseModel): diff --git a/app/services/company_metadata_service.py b/app/services/company_metadata_service.py new file mode 100644 index 0000000..8c8ed50 --- /dev/null +++ b/app/services/company_metadata_service.py @@ -0,0 +1,253 @@ +""" +CompanyMetadataService — authoritative source for company metadata. + +Lookup order: + 1. Redis cache (24h TTL) + 2. universe_ticker_registry (DB) — if sector NOT NULL, return directly + 3. yfinance .info enrichment (semaphore=5) — upsert results back to registry + companies + 4. If yfinance fails, return whatever registry has (may have NULL sector) + 5. If ticker not in registry AND yfinance fails/invalid → raise ValueError("invalid ticker") +""" + +import asyncio +import logging +from datetime import datetime, timezone +from typing import Optional + +import yfinance as yf +from sqlalchemy import select +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.financial import Company +from app.models.universe_snapshot import UniverseSnapshot, UniverseTickerRegistry +from app.utils.cache import build_cache_key, get_cached_response, set_cached_response + +logger = logging.getLogger(__name__) + +_YFINANCE_SEMAPHORE = asyncio.Semaphore(5) + +_EXCHANGE_MAP = { + "NYQ": "NYSE", "NMS": "NASDAQ", "NGM": "NASDAQ", + "NCM": "NASDAQ", "ASE": "AMEX", "PCX": "NYSE_ARCA", +} +_SKIP_QUOTE_TYPES = {"ETF", "MUTUALFUND", "INDEX", "CURRENCY", "FUTURE", "OPTION"} +_CACHE_TTL = 60 * 60 * 24 # 24h + + +def _canonical_exchange(raw: Optional[str]) -> Optional[str]: + if not raw: + return None + return _EXCHANGE_MAP.get(raw.upper(), raw.upper()) or None + + +async def _fetch_yfinance_info(ticker: str) -> Optional[dict]: + """Fetch yfinance .info in a thread pool with semaphore and 20s timeout.""" + loop = asyncio.get_event_loop() + + def _sync_fetch(): + try: + t = yf.Ticker(ticker) + return t.info + except Exception as e: + logger.warning("yfinance .info failed for %s: %s", ticker, e) + return None + + async with _YFINANCE_SEMAPHORE: + try: + return await asyncio.wait_for( + loop.run_in_executor(None, _sync_fetch), + timeout=20, + ) + except asyncio.TimeoutError: + logger.warning("yfinance .info timed out for %s", ticker) + return None + except Exception as e: + logger.warning("yfinance enrichment error for %s: %s", ticker, e) + return None + + +async def _upsert_registry(db: AsyncSession, ticker: str, data: dict) -> None: + stmt = pg_insert(UniverseTickerRegistry).values( + ticker=ticker, + name=data.get("name"), + cik=data.get("cik"), + sector=data.get("sector"), + industry=data.get("industry"), + exchange=data.get("exchange"), + is_active=True, + updated_at=datetime.now(timezone.utc), + ).on_conflict_do_update( + constraint="uq_universe_ticker_registry", + set_={ + "name": data.get("name"), + "sector": data.get("sector"), + "industry": data.get("industry"), + "exchange": data.get("exchange"), + "updated_at": datetime.now(timezone.utc), + }, + ) + await db.execute(stmt) + + +async def _upsert_company(db: AsyncSession, ticker: str, data: dict) -> None: + result = await db.execute(select(Company).where(Company.ticker == ticker)) + company = result.scalar_one_or_none() + now = datetime.now(timezone.utc) + + if company: + if data.get("name") and (not company.name or company.name == f"{ticker} Corporation"): + company.name = data["name"] + if data.get("sector"): + company.sector = data["sector"] + if data.get("industry"): + company.industry = data["industry"] + if data.get("exchange"): + company.exchange = data["exchange"] + if data.get("country"): + company.country = data["country"] + if data.get("market_cap"): + company.market_cap = data["market_cap"] + if data.get("business_description"): + company.business_description = data["business_description"] + company.updated_at = now + else: + company = Company( + ticker=ticker, + name=data.get("name") or f"{ticker} Corporation", + cik=data.get("cik"), + exchange=data.get("exchange"), + sector=data.get("sector"), + industry=data.get("industry"), + country=data.get("country"), + market_cap=data.get("market_cap"), + business_description=data.get("business_description"), + created_at=now, + updated_at=now, + ) + db.add(company) + + await db.commit() + + +async def get_metadata(db: AsyncSession, ticker: str) -> dict: + """ + Return company metadata dict for a ticker. + Raises ValueError("invalid ticker: {ticker}") if ticker is unknown. + + Returned dict keys: ticker, name, cik, exchange, sector, industry, + country, market_cap, business_description + """ + ticker = ticker.upper() + cache_key = build_cache_key("company:meta", ticker) + + # --- 1. Redis cache --- + cached = await get_cached_response(cache_key) + if cached: + body, _ = cached + return body + + # --- 2. Registry DB --- + reg_result = await db.execute( + select(UniverseTickerRegistry).where(UniverseTickerRegistry.ticker == ticker) + ) + reg = reg_result.scalar_one_or_none() + + # --- 3. market_cap from latest snapshot (opportunistic) --- + snap_market_cap: Optional[float] = None + snap_result = await db.execute( + select(UniverseSnapshot) + .where(UniverseSnapshot.ticker == ticker) + .order_by(UniverseSnapshot.snapshot_date.desc()) + .limit(1) + ) + snap = snap_result.scalar_one_or_none() + if snap: + snap_market_cap = snap.market_cap + + if reg and reg.sector: + # Fast path: registry has sector — no yfinance needed + data = _build_from_registry(ticker, reg, snap_market_cap) + await set_cached_response(cache_key, data, ttl_seconds=_CACHE_TTL) + return data + + # --- 4. yfinance enrichment --- + info = await _fetch_yfinance_info(ticker) + + if info: + quote_type = info.get("quoteType") or "" + if quote_type in _SKIP_QUOTE_TYPES: + # Valid but not an equity — return minimal + data = _build_minimal(ticker, info, snap_market_cap) + await set_cached_response(cache_key, data, ttl_seconds=_CACHE_TTL) + return data + + if not quote_type and not info.get("longName") and not info.get("shortName"): + # Likely invalid ticker + if not reg: + raise ValueError(f"invalid ticker: {ticker}") + # Fall through to registry-only result + + enriched = { + "name": info.get("longName") or info.get("shortName"), + "cik": reg.cik if reg else None, + "exchange": _canonical_exchange(info.get("exchange")), + "sector": info.get("sector"), + "industry": info.get("industry"), + "country": info.get("country"), + "market_cap": info.get("marketCap") or snap_market_cap, + "business_description": info.get("longBusinessSummary"), + } + + # Persist enrichment + try: + await _upsert_registry(db, ticker, enriched) + await _upsert_company(db, ticker, enriched) + except Exception as e: + logger.warning("Failed to persist enrichment for %s: %s", ticker, e) + + data = { + "ticker": ticker, + **enriched, + } + await set_cached_response(cache_key, data, ttl_seconds=_CACHE_TTL) + return data + + # --- 5. yfinance failed — use registry if available --- + if reg: + data = _build_from_registry(ticker, reg, snap_market_cap) + # Short TTL so we retry enrichment soon + await set_cached_response(cache_key, data, ttl_seconds=60 * 15) + return data + + raise ValueError(f"invalid ticker: {ticker}") + + +def _build_from_registry( + ticker: str, reg: UniverseTickerRegistry, market_cap: Optional[float] +) -> dict: + return { + "ticker": ticker, + "name": reg.name, + "cik": reg.cik, + "exchange": reg.exchange, + "sector": reg.sector, + "industry": reg.industry, + "country": None, + "market_cap": market_cap, + "business_description": None, + } + + +def _build_minimal(ticker: str, info: dict, market_cap: Optional[float]) -> dict: + return { + "ticker": ticker, + "name": info.get("longName") or info.get("shortName"), + "cik": None, + "exchange": _canonical_exchange(info.get("exchange")), + "sector": info.get("sector"), + "industry": info.get("industry"), + "country": info.get("country"), + "market_cap": info.get("marketCap") or market_cap, + "business_description": None, + } diff --git a/app/services/financial_service.py b/app/services/financial_service.py index 2ab250d..79c0bc8 100644 --- a/app/services/financial_service.py +++ b/app/services/financial_service.py @@ -95,19 +95,26 @@ class FinancialService: start_date, end_date, quarters, period, ticker ) - # Get or create company + # Get or create company — always succeeds (may have NULL sector for truly unknown tickers) company = await self._get_or_create_company(db, ticker) - - # Get financial data from database or generate realistic data - financial_data = await self._get_or_generate_financial_data( - db, ticker, resolved_start, resolved_end, force_refresh - ) - - # Get price data for calculations - price_data = await self._get_price_data_for_period( - db, ticker, resolved_start, resolved_end - ) - + + # Fetch financials and price data independently; failures return empty lists. + try: + financial_data = await self._get_or_generate_financial_data( + db, ticker, resolved_start, resolved_end, force_refresh + ) + except Exception as e: + logger.warning("Financial data fetch failed for %s: %s", ticker, e) + financial_data = [] + + try: + price_data = await self._get_price_data_for_period( + db, ticker, resolved_start, resolved_end + ) + except Exception as e: + logger.warning("Price data fetch failed for %s: %s", ticker, e) + price_data = [] + # Calculate metrics using real price data calculated_metrics = await self._calculate_real_metrics( db, ticker, financial_data, price_data, force_refresh @@ -119,73 +126,78 @@ class FinancialService: "calculated_metrics": calculated_metrics } - async def _get_or_create_company(self, db: AsyncSession, ticker: str) -> Company: - """Get or create company record""" - result = await db.execute( - select(Company).where(Company.ticker == ticker) + @staticmethod + def _is_placeholder(company: Company) -> bool: + """True when the row was created by the old hardcoded-defaults path.""" + return bool( + company.sector == "Technology" + and company.industry == "Software" + and company.name + and company.name.endswith(" Corporation") ) + + async def _get_or_create_company(self, db: AsyncSession, ticker: str) -> Company: + """Get or create company record, enriching via CompanyMetadataService.""" + result = await db.execute(select(Company).where(Company.ticker == ticker)) + company = result.scalar_one_or_none() + # Fast path: row exists with real sector (not a hardcoded placeholder) + if company and company.sector and not self._is_placeholder(company): + return company + + # Enrich via registry + yfinance + try: + from app.services import company_metadata_service as cms + meta = await cms.get_metadata(db, ticker) + except ValueError: + # Invalid ticker — still create a minimal placeholder so the rest of the + # financial pipeline doesn't break. + meta = { + "name": f"{ticker} Corporation", + "cik": None, + "exchange": None, + "sector": None, + "industry": None, + "country": None, + "market_cap": None, + "business_description": None, + } + except Exception as e: + logger.warning("CompanyMetadataService failed for %s: %s — using placeholder", ticker, e) + meta = { + "name": f"{ticker} Corporation", + "cik": None, + "exchange": None, + "sector": None, + "industry": None, + "country": None, + "market_cap": None, + "business_description": None, + } + + # cms.get_metadata already upserts the Company row; re-fetch or create if missing. + result = await db.execute(select(Company).where(Company.ticker == ticker)) company = result.scalar_one_or_none() - if not company: - # Create company with basic info (in real implementation, this would fetch from SEC) - company_info = self._get_default_company_info(ticker) + now = datetime.now(timezone.utc) company = Company( ticker=ticker, - name=company_info["name"], - cik=company_info["cik"], - sector=company_info["sector"], - industry=company_info["industry"], - business_description=company_info["business_description"], - created_at=datetime.now(timezone.utc), - updated_at=datetime.now(timezone.utc) + name=meta.get("name") or f"{ticker} Corporation", + cik=meta.get("cik"), + exchange=meta.get("exchange"), + sector=meta.get("sector"), + industry=meta.get("industry"), + country=meta.get("country"), + market_cap=meta.get("market_cap"), + business_description=meta.get("business_description"), + created_at=now, + updated_at=now, ) db.add(company) await db.commit() await db.refresh(company) - + return company - def _get_default_company_info(self, ticker: str) -> Dict: - """Get default company info (placeholder for real SEC data)""" - company_defaults = { - 'AAPL': { - 'name': 'Apple Inc.', - 'cik': '0000320193', - 'sector': 'Technology', - 'industry': 'Consumer Electronics', - 'business_description': 'Technology company designing and manufacturing consumer electronics' - }, - 'MSFT': { - 'name': 'Microsoft Corporation', - 'cik': '0000789019', - 'sector': 'Technology', - 'industry': 'Software—Infrastructure', - 'business_description': 'Software and cloud services company' - }, - 'TSLA': { - 'name': 'Tesla Inc.', - 'cik': '0001318605', - 'sector': 'Consumer Cyclical', - 'industry': 'Auto Manufacturers', - 'business_description': 'Electric vehicle and clean energy company' - }, - 'NVDA': { - 'name': 'NVIDIA Corporation', - 'cik': '0001045810', - 'sector': 'Technology', - 'industry': 'Semiconductors', - 'business_description': 'Semiconductor company specializing in graphics processing units' - } - } - - return company_defaults.get(ticker, { - 'name': f'{ticker} Corporation', - 'cik': f'000{hash(ticker) % 1000000:06d}', - 'sector': 'Technology', - 'industry': 'Software', - 'business_description': f'{ticker} technology company' - }) - async def _get_or_generate_financial_data( self, db: AsyncSession, diff --git a/scripts/backfill_registry_sector.py b/scripts/backfill_registry_sector.py new file mode 100644 index 0000000..69de62e --- /dev/null +++ b/scripts/backfill_registry_sector.py @@ -0,0 +1,93 @@ +#!/usr/bin/env python3 +""" +Backfill sector/industry/exchange/country for all active tickers in universe_ticker_registry +that currently have sector=NULL. + +Usage: + python scripts/backfill_registry_sector.py [--batch 50] [--dry-run] + +Expected runtime: ~30 min for 9376 tickers at semaphore=5. +""" + +import asyncio +import argparse +import logging +import sys +import os + +sys.path.insert(0, os.path.dirname(os.path.dirname(__file__))) + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker + +from app.core.config import settings +from app.models.universe_snapshot import UniverseTickerRegistry +from app.services import company_metadata_service as cms + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +logger = logging.getLogger(__name__) + + +async def backfill(batch_size: int, dry_run: bool) -> None: + engine = create_async_engine(settings.DATABASE_URL, echo=False) + Session = async_sessionmaker(engine, expire_on_commit=False) + + async with Session() as db: + result = await db.execute( + select(UniverseTickerRegistry.ticker) + .where( + UniverseTickerRegistry.is_active == True, + UniverseTickerRegistry.sector == None, + ) + .order_by(UniverseTickerRegistry.ticker) + ) + tickers = [row[0] for row in result.fetchall()] + + logger.info("Found %d tickers with NULL sector", len(tickers)) + if dry_run: + logger.info("Dry-run mode — no changes will be written.") + return + + success = 0 + failed = 0 + + for i in range(0, len(tickers), batch_size): + batch = tickers[i:i + batch_size] + logger.info("Batch %d/%d — processing %d tickers", i // batch_size + 1, + (len(tickers) + batch_size - 1) // batch_size, len(batch)) + + async def _enrich(ticker: str) -> bool: + async with Session() as db: + try: + meta = await cms.get_metadata(db, ticker) + sector = meta.get("sector") + logger.info(" %s → sector=%s", ticker, sector or "NULL") + return bool(sector) + except ValueError: + logger.debug(" %s → invalid ticker, skipping", ticker) + return False + except Exception as e: + logger.warning(" %s → error: %s", ticker, e) + return False + + results = await asyncio.gather(*[_enrich(t) for t in batch]) + success += sum(results) + failed += sum(1 for r in results if not r) + + # Brief pause between batches to be kind to yfinance + await asyncio.sleep(1) + + logger.info("Done. success=%d failed/skipped=%d total=%d", success, failed, len(tickers)) + await engine.dispose() + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--batch", type=int, default=50) + parser.add_argument("--dry-run", action="store_true") + args = parser.parse_args() + asyncio.run(backfill(args.batch, args.dry_run)) + + +if __name__ == "__main__": + main()