From 0471018dd817034dbf7f0be8f0fd5fbee7ef0611 Mon Sep 17 00:00:00 2001 From: I Luk Kim Date: Thu, 12 Mar 2026 08:44:53 -0700 Subject: [PATCH] =?UTF-8?q?feat:=20implement=20ACE-F=20v1=20Phase=201=20--?= =?UTF-8?q?=20Stock=20Oracle=20=EA=B8=B0=EB=B0=98=20=EC=9D=B4=EB=B2=A4?= =?UTF-8?q?=ED=8A=B8=20=ED=8C=8C=EC=9D=B4=ED=94=84=EB=9D=BC=EC=9D=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Stock Oracle (localhost:18001)을 단일 데이터 소스로 사용하는 미국 주식 이벤트 스윙 트레이딩 시스템의 Phase 1 구현체. 주요 구성: - libs/oracle_client/: Stock Oracle REST 클라이언트 (filings, price, financial, fred, finra) - libs/common/: config, logging, time_utils, ids, retries, file_store - libs/db/: SQLAlchemy 2.0 모델 12개 + Alembic 마이그레이션 0001 - libs/parser/: 규칙 기반 파서 (텍스트 정규화, JSON Schema 검증, LLM 스텁) - libs/features/: 시장/이벤트 피처 계산기 - apps/pipeline/: filing_poller → filing_fetcher → event_parser → feature_builder - apps/sync/: macro_sync (FRED), short_volume_sync (FINRA), issuer_sync - tests/: 단위 81개 + 리플레이 5개 전체 통과, lint clean Co-Authored-By: Claude Sonnet 4.6 --- .env.example | 8 + .gitignore | 18 + .python-version | 1 + Makefile | 39 ++ alembic.ini | 38 ++ apps/__init__.py | 0 apps/pipeline/__init__.py | 0 apps/pipeline/event_parser/__init__.py | 0 apps/pipeline/event_parser/main.py | 211 +++++++++ apps/pipeline/feature_builder/__init__.py | 0 apps/pipeline/feature_builder/main.py | 84 ++++ apps/pipeline/filing_fetcher/__init__.py | 0 apps/pipeline/filing_fetcher/main.py | 143 ++++++ apps/pipeline/filing_poller/__init__.py | 0 apps/pipeline/filing_poller/main.py | 129 ++++++ apps/sync/__init__.py | 0 apps/sync/issuer_sync/__init__.py | 0 apps/sync/issuer_sync/main.py | 130 ++++++ apps/sync/macro_sync/__init__.py | 0 apps/sync/macro_sync/main.py | 139 ++++++ apps/sync/short_volume_sync/__init__.py | 0 apps/sync/short_volume_sync/main.py | 119 +++++ configs/app.yaml | 22 + configs/fred_series.yaml | 20 + configs/symbols.yaml | 16 + docker-compose.yml | 28 ++ docker/Dockerfile | 16 + libs/__init__.py | 0 libs/common/__init__.py | 0 libs/common/config.py | 80 ++++ libs/common/file_store.py | 77 ++++ libs/common/ids.py | 35 ++ libs/common/logging.py | 55 +++ libs/common/retries.py | 83 ++++ libs/common/time_utils.py | 80 ++++ libs/db/__init__.py | 0 libs/db/engine.py | 20 + libs/db/enums.py | 73 +++ libs/db/helpers.py | 79 ++++ libs/db/migrations/__init__.py | 0 libs/db/migrations/env.py | 66 +++ libs/db/migrations/script.py.mako | 28 ++ .../versions/0001_initial_schema.py | 414 ++++++++++++++++++ libs/db/models.py | 342 +++++++++++++++ libs/db/session.py | 33 ++ libs/features/__init__.py | 0 libs/features/builder.py | 98 +++++ libs/features/event_features.py | 88 ++++ libs/features/market_features.py | 110 +++++ libs/oracle_client/__init__.py | 0 libs/oracle_client/client.py | 108 +++++ libs/oracle_client/exceptions.py | 24 + libs/oracle_client/filings.py | 42 ++ libs/oracle_client/financial.py | 21 + libs/oracle_client/finra.py | 24 + libs/oracle_client/fred.py | 38 ++ libs/oracle_client/models.py | 181 ++++++++ libs/oracle_client/price.py | 40 ++ libs/parser/__init__.py | 0 libs/parser/llm_parser_stub.py | 24 + libs/parser/rule_parser.py | 382 ++++++++++++++++ libs/parser/schema_validator.py | 30 ++ libs/parser/text_normalizer.py | 57 +++ libs/schemas/__init__.py | 0 libs/schemas/parser_event.schema.json | 145 ++++++ libs/schemas/types.py | 74 ++++ pyproject.toml | 64 +++ tests/__init__.py | 0 tests/conftest.py | 93 ++++ tests/fixtures/__init__.py | 0 tests/fixtures/exhibit_content.json | 6 + tests/fixtures/filing_search.json | 24 + tests/fixtures/financial_data.json | 23 + tests/fixtures/fred_observations.json | 11 + tests/fixtures/price_data.json | 12 + tests/fixtures/short_volume.json | 8 + tests/integration/__init__.py | 0 tests/integration/conftest.py | 61 +++ tests/integration/test_db_migration.py | 42 ++ tests/integration/test_feature_pipeline.py | 83 ++++ tests/integration/test_filing_pipeline.py | 128 ++++++ tests/integration/test_sync_jobs.py | 67 +++ tests/replay/__init__.py | 0 tests/replay/test_determinism.py | 56 +++ tests/replay/test_idempotency.py | 40 ++ tests/unit/__init__.py | 0 tests/unit/test_config.py | 63 +++ tests/unit/test_db_models.py | 65 +++ tests/unit/test_event_features.py | 87 ++++ tests/unit/test_file_store.py | 43 ++ tests/unit/test_ids.py | 53 +++ tests/unit/test_market_features.py | 72 +++ tests/unit/test_oracle_client.py | 122 ++++++ tests/unit/test_retries.py | 77 ++++ tests/unit/test_rule_parser.py | 110 +++++ tests/unit/test_schema_validator.py | 39 ++ tests/unit/test_text_normalizer.py | 49 +++ tests/unit/test_time_utils.py | 59 +++ 98 files changed, 5669 insertions(+) create mode 100644 .env.example create mode 100644 .gitignore create mode 100644 .python-version create mode 100644 Makefile create mode 100644 alembic.ini create mode 100644 apps/__init__.py create mode 100644 apps/pipeline/__init__.py create mode 100644 apps/pipeline/event_parser/__init__.py create mode 100644 apps/pipeline/event_parser/main.py create mode 100644 apps/pipeline/feature_builder/__init__.py create mode 100644 apps/pipeline/feature_builder/main.py create mode 100644 apps/pipeline/filing_fetcher/__init__.py create mode 100644 apps/pipeline/filing_fetcher/main.py create mode 100644 apps/pipeline/filing_poller/__init__.py create mode 100644 apps/pipeline/filing_poller/main.py create mode 100644 apps/sync/__init__.py create mode 100644 apps/sync/issuer_sync/__init__.py create mode 100644 apps/sync/issuer_sync/main.py create mode 100644 apps/sync/macro_sync/__init__.py create mode 100644 apps/sync/macro_sync/main.py create mode 100644 apps/sync/short_volume_sync/__init__.py create mode 100644 apps/sync/short_volume_sync/main.py create mode 100644 configs/app.yaml create mode 100644 configs/fred_series.yaml create mode 100644 configs/symbols.yaml create mode 100644 docker-compose.yml create mode 100644 docker/Dockerfile create mode 100644 libs/__init__.py create mode 100644 libs/common/__init__.py create mode 100644 libs/common/config.py create mode 100644 libs/common/file_store.py create mode 100644 libs/common/ids.py create mode 100644 libs/common/logging.py create mode 100644 libs/common/retries.py create mode 100644 libs/common/time_utils.py create mode 100644 libs/db/__init__.py create mode 100644 libs/db/engine.py create mode 100644 libs/db/enums.py create mode 100644 libs/db/helpers.py create mode 100644 libs/db/migrations/__init__.py create mode 100644 libs/db/migrations/env.py create mode 100644 libs/db/migrations/script.py.mako create mode 100644 libs/db/migrations/versions/0001_initial_schema.py create mode 100644 libs/db/models.py create mode 100644 libs/db/session.py create mode 100644 libs/features/__init__.py create mode 100644 libs/features/builder.py create mode 100644 libs/features/event_features.py create mode 100644 libs/features/market_features.py create mode 100644 libs/oracle_client/__init__.py create mode 100644 libs/oracle_client/client.py create mode 100644 libs/oracle_client/exceptions.py create mode 100644 libs/oracle_client/filings.py create mode 100644 libs/oracle_client/financial.py create mode 100644 libs/oracle_client/finra.py create mode 100644 libs/oracle_client/fred.py create mode 100644 libs/oracle_client/models.py create mode 100644 libs/oracle_client/price.py create mode 100644 libs/parser/__init__.py create mode 100644 libs/parser/llm_parser_stub.py create mode 100644 libs/parser/rule_parser.py create mode 100644 libs/parser/schema_validator.py create mode 100644 libs/parser/text_normalizer.py create mode 100644 libs/schemas/__init__.py create mode 100644 libs/schemas/parser_event.schema.json create mode 100644 libs/schemas/types.py create mode 100644 pyproject.toml create mode 100644 tests/__init__.py create mode 100644 tests/conftest.py create mode 100644 tests/fixtures/__init__.py create mode 100644 tests/fixtures/exhibit_content.json create mode 100644 tests/fixtures/filing_search.json create mode 100644 tests/fixtures/financial_data.json create mode 100644 tests/fixtures/fred_observations.json create mode 100644 tests/fixtures/price_data.json create mode 100644 tests/fixtures/short_volume.json create mode 100644 tests/integration/__init__.py create mode 100644 tests/integration/conftest.py create mode 100644 tests/integration/test_db_migration.py create mode 100644 tests/integration/test_feature_pipeline.py create mode 100644 tests/integration/test_filing_pipeline.py create mode 100644 tests/integration/test_sync_jobs.py create mode 100644 tests/replay/__init__.py create mode 100644 tests/replay/test_determinism.py create mode 100644 tests/replay/test_idempotency.py create mode 100644 tests/unit/__init__.py create mode 100644 tests/unit/test_config.py create mode 100644 tests/unit/test_db_models.py create mode 100644 tests/unit/test_event_features.py create mode 100644 tests/unit/test_file_store.py create mode 100644 tests/unit/test_ids.py create mode 100644 tests/unit/test_market_features.py create mode 100644 tests/unit/test_oracle_client.py create mode 100644 tests/unit/test_retries.py create mode 100644 tests/unit/test_rule_parser.py create mode 100644 tests/unit/test_schema_validator.py create mode 100644 tests/unit/test_text_normalizer.py create mode 100644 tests/unit/test_time_utils.py diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..c408982 --- /dev/null +++ b/.env.example @@ -0,0 +1,8 @@ +APP_ENV=dev +STOCK_ORACLE_URL=http://localhost:18001 +STOCK_ORACLE_TIMEOUT=30 +POSTGRES_DSN=postgresql+asyncpg://acef:acef@localhost:5432/acef +DATA_ROOT=./data +LOG_LEVEL=INFO +OPENAI_API_KEY=sk-placeholder +LLM_ENABLED=false diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..fd0611e --- /dev/null +++ b/.gitignore @@ -0,0 +1,18 @@ +.env +.env.* +!.env.example +data/ +__pycache__/ +.venv/ +venv/ +.DS_Store +*.pyc +*.pyo +.mypy_cache/ +.ruff_cache/ +.pytest_cache/ +dist/ +build/ +*.egg-info/ +.coverage +htmlcov/ diff --git a/.python-version b/.python-version new file mode 100644 index 0000000..2c07333 --- /dev/null +++ b/.python-version @@ -0,0 +1 @@ +3.11 diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..cb3f877 --- /dev/null +++ b/Makefile @@ -0,0 +1,39 @@ +.PHONY: bootstrap db-upgrade db-downgrade db-reset test lint typecheck ci test-unit test-integration test-replay + +bootstrap: + pip install -e ".[dev]" + +db-upgrade: + alembic upgrade head + +db-downgrade: + alembic downgrade -1 + +db-reset: + alembic downgrade base + alembic upgrade head + +test-unit: + pytest -m unit -v --cov=libs --cov=apps --cov-report=term-missing + +test-integration: + pytest -m integration -v + +test-replay: + pytest -m replay -v + +test: + pytest -v --cov=libs --cov=apps --cov-report=term-missing + +lint: + ruff check libs/ apps/ tests/ + ruff format --check libs/ apps/ tests/ + +typecheck: + mypy libs/ apps/ + +ci: lint typecheck test-unit + +format: + ruff format libs/ apps/ tests/ + ruff check --fix libs/ apps/ tests/ diff --git a/alembic.ini b/alembic.ini new file mode 100644 index 0000000..ff0439f --- /dev/null +++ b/alembic.ini @@ -0,0 +1,38 @@ +[alembic] +script_location = libs/db/migrations +prepend_sys_path = . +version_path_separator = os + +[loggers] +keys = root,sqlalchemy,alembic + +[handlers] +keys = console + +[formatters] +keys = generic + +[logger_root] +level = WARN +handlers = console +qualname = + +[logger_sqlalchemy] +level = WARN +handlers = +qualname = sqlalchemy.engine + +[logger_alembic] +level = INFO +handlers = +qualname = alembic + +[handler_console] +class = StreamHandler +args = (sys.stderr,) +level = NOTSET +formatter = generic + +[formatter_generic] +format = %(levelname)-5.5s [%(name)s] %(message)s +datefmt = %H:%M:%S diff --git a/apps/__init__.py b/apps/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/__init__.py b/apps/pipeline/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/event_parser/__init__.py b/apps/pipeline/event_parser/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/event_parser/main.py b/apps/pipeline/event_parser/main.py new file mode 100644 index 0000000..8b13d5a --- /dev/null +++ b/apps/pipeline/event_parser/main.py @@ -0,0 +1,211 @@ +"""Event Parser: parse exhibit text → events + event_parses.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy import select + +from libs.common.config import get_settings +from libs.common.file_store import read_exhibit +from libs.common.ids import event_id as make_event_id +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import Document, Event, EventParse, JobRun +from libs.db.session import get_session +from libs.parser.rule_parser import PARSER_VERSION, SCHEMA_VERSION, RuleBasedParser +from libs.parser.schema_validator import validate_parser_output +from libs.parser.text_normalizer import normalize_text + +logger = get_logger(__name__) +_parser = RuleBasedParser() + + +async def run_event_parser(run_id: str) -> dict[str, int]: + settings = get_settings() + app_config = settings.get_app_config() + exhibit_types = app_config.get("pipeline", {}).get("exhibit_types", ["EX-99.1"]) + + stats = {"seen": 0, "valid": 0, "invalid": 0, "errors": 0} + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="event_parser", + source_name="oracle", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + result = await session.execute( + select(Document).where(Document.parsed_status == "ready_for_parse") + ) + docs = result.scalars().all() + stats["seen"] = len(docs) + + for doc in docs: + if not doc.accession_no: + stats["invalid"] += 1 + doc.parsed_status = "failed" + continue + + # Try each exhibit type until we find one + text: str | None = None + for exhibit_type in exhibit_types: + try: + text = read_exhibit(doc.accession_no, exhibit_type) + break + except FileNotFoundError: + continue + + if text is None: + logger.warning("no_exhibit_text", accession_no=doc.accession_no) + doc.parsed_status = "failed" + stats["errors"] += 1 + continue + + normalized = normalize_text(text) + + metadata = { + "filing_date": doc.filing_date.isoformat(), + "accepted_at_utc": ( + doc.accepted_at_utc.isoformat() if doc.accepted_at_utc else None + ), + "form_type": doc.form_type, + } + + try: + output = _parser.parse( + document_id=doc.document_id, + form_type=doc.form_type, + text=normalized, + metadata=metadata, + ) + output_dict = output.model_dump() + errors = validate_parser_output(output_dict) + + if errors: + logger.warning( + "parse_validation_failed", + document_id=doc.document_id, + errors=errors[:3], + ) + # Store invalid parse, no event row + event_id_str = make_event_id(doc.document_id, output.event_type) + # Create a minimal event row first (needed for FK) + event = Event( + event_id=event_id_str, + primary_document_id=doc.document_id, + issuer_id=doc.issuer_id, + symbol_id=doc.symbol_id, + event_type=output.event_type, + event_direction=output.event_direction, + event_date=doc.filing_date, + filed_at_utc=doc.accepted_at_utc, + parser_version=PARSER_VERSION, + parse_confidence=output.confidence.overall, + status="rejected", + ) + session.add(event) + await session.flush() + + parse_row = EventParse( + event_id=event_id_str, + parser_kind="rule", + parser_version=PARSER_VERSION, + schema_version=SCHEMA_VERSION, + output_json=output_dict, + validation_status="invalid", + validation_errors={"errors": errors}, + ) + session.add(parse_row) + doc.parsed_status = "failed" + stats["invalid"] += 1 + else: + event_id_str = make_event_id(doc.document_id, output.event_type) + + # Check duplicate event + existing_evt = await session.execute( + select(Event).where(Event.event_id == event_id_str) + ) + if existing_evt.scalar_one_or_none() is not None: + logger.info("event_already_exists", event_id=event_id_str) + doc.parsed_status = "succeeded" + stats["valid"] += 1 + continue + + event = Event( + event_id=event_id_str, + primary_document_id=doc.document_id, + issuer_id=doc.issuer_id, + symbol_id=doc.symbol_id, + event_type=output.event_type, + event_direction=output.event_direction, + event_date=doc.filing_date, + filed_at_utc=doc.accepted_at_utc, + parser_version=PARSER_VERSION, + parse_confidence=output.confidence.overall, + status="pending", + ) + session.add(event) + await session.flush() + + parse_row = EventParse( + event_id=event_id_str, + parser_kind="rule", + parser_version=PARSER_VERSION, + schema_version=SCHEMA_VERSION, + output_json=output_dict, + validation_status="valid", + validation_errors=None, + ) + session.add(parse_row) + doc.parsed_status = "succeeded" + stats["valid"] += 1 + + logger.info( + "event_created", + event_id=event_id_str, + event_type=output.event_type, + confidence=output.confidence.overall, + ) + + except Exception as exc: + logger.error( + "parse_error", + document_id=doc.document_id, + error=str(exc), + ) + doc.parsed_status = "failed" + stats["errors"] += 1 + + doc.updated_at_utc = dt.datetime.now(tz=dt.UTC) + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["valid"] + job.error_count = stats["errors"] + stats["invalid"] + + logger.info("event_parser_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Event Parser") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(run_event_parser(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/pipeline/feature_builder/__init__.py b/apps/pipeline/feature_builder/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/feature_builder/main.py b/apps/pipeline/feature_builder/main.py new file mode 100644 index 0000000..d64870f --- /dev/null +++ b/apps/pipeline/feature_builder/main.py @@ -0,0 +1,84 @@ +"""Feature Builder: compute market + event features for pending events.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy import select + +from libs.common.config import get_settings +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import Event, JobRun +from libs.db.session import get_session +from libs.features.builder import build_features_for_event +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.price import PriceService + +logger = get_logger(__name__) + + +async def run_feature_builder(run_id: str) -> dict[str, int]: + stats = {"seen": 0, "built": 0, "skipped": 0, "errors": 0} + + async with make_oracle_client() as client: + price_svc = PriceService(client) + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="feature_builder", + source_name="oracle", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + result = await session.execute( + select(Event).where(Event.status == "pending") + ) + events = result.scalars().all() + stats["seen"] = len(events) + + for event in events: + try: + snapshots = await build_features_for_event(session, event, price_svc) + if snapshots is None: + event.status = "rejected" + event.updated_at_utc = dt.datetime.now(tz=dt.UTC) + stats["skipped"] += 1 + else: + event.status = "valid" + event.updated_at_utc = dt.datetime.now(tz=dt.UTC) + stats["built"] += 1 + except Exception as exc: + logger.error("feature_build_error", event_id=event.event_id, error=str(exc)) + stats["errors"] += 1 + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["built"] + job.error_count = stats["errors"] + + logger.info("feature_builder_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Feature Builder") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(run_feature_builder(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/pipeline/filing_fetcher/__init__.py b/apps/pipeline/filing_fetcher/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/filing_fetcher/main.py b/apps/pipeline/filing_fetcher/main.py new file mode 100644 index 0000000..9d87f44 --- /dev/null +++ b/apps/pipeline/filing_fetcher/main.py @@ -0,0 +1,143 @@ +"""Filing Fetcher: download exhibit text and cache locally.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy import select + +from libs.common.config import get_settings +from libs.common.file_store import exists_exhibit, write_exhibit +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import Document, ExhibitCache, JobRun +from libs.db.session import get_session +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.exceptions import OracleNotFoundError +from libs.oracle_client.filings import FilingsService + +logger = get_logger(__name__) + + +async def fetch_exhibits(run_id: str) -> dict[str, int]: + settings = get_settings() + app_config = settings.get_app_config() + exhibit_types = app_config.get("pipeline", {}).get("exhibit_types", ["EX-99.1"]) + + stats = {"seen": 0, "written": 0, "skipped": 0, "errors": 0} + + async with make_oracle_client() as client: + svc = FilingsService(client) + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="filing_fetcher", + source_name="oracle", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + result = await session.execute( + select(Document).where(Document.parsed_status == "pending") + ) + docs = result.scalars().all() + stats["seen"] = len(docs) + + for doc in docs: + if not doc.accession_no: + stats["skipped"] += 1 + continue + + fetched_any = False + for exhibit_type in exhibit_types: + if exists_exhibit(doc.accession_no, exhibit_type): + logger.info( + "exhibit_already_cached", + accession_no=doc.accession_no, + exhibit_type=exhibit_type, + ) + fetched_any = True + continue + + try: + response = await svc.get_exhibit(doc.accession_no, exhibit_type) + checksum = write_exhibit( + doc.accession_no, exhibit_type, response.content + ) + + from libs.common.file_store import exhibit_path + + cache_path = str(exhibit_path(doc.accession_no, exhibit_type)) + + existing_cache = await session.execute( + select(ExhibitCache).where( + ExhibitCache.accession_no == doc.accession_no, + ExhibitCache.exhibit_type == exhibit_type, + ) + ) + if existing_cache.scalar_one_or_none() is None: + cache_row = ExhibitCache( + accession_no=doc.accession_no, + exhibit_type=exhibit_type, + content_hash=checksum, + cache_path=cache_path, + ) + session.add(cache_row) + + fetched_any = True + stats["written"] += 1 + logger.info( + "exhibit_fetched", + accession_no=doc.accession_no, + exhibit_type=exhibit_type, + ) + + except OracleNotFoundError: + logger.warning( + "exhibit_not_found", + accession_no=doc.accession_no, + exhibit_type=exhibit_type, + ) + except Exception as exc: + logger.error( + "exhibit_fetch_error", + accession_no=doc.accession_no, + exhibit_type=exhibit_type, + error=str(exc), + ) + stats["errors"] += 1 + + if fetched_any: + doc.parsed_status = "ready_for_parse" + doc.updated_at_utc = dt.datetime.now(tz=dt.UTC) + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["written"] + job.records_skipped = stats["skipped"] + job.error_count = stats["errors"] + + logger.info("filing_fetcher_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Filing Fetcher") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(fetch_exhibits(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/pipeline/filing_poller/__init__.py b/apps/pipeline/filing_poller/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/pipeline/filing_poller/main.py b/apps/pipeline/filing_poller/main.py new file mode 100644 index 0000000..b64cf71 --- /dev/null +++ b/apps/pipeline/filing_poller/main.py @@ -0,0 +1,129 @@ +"""Filing Poller: discover new 8-K/6-K filings via Stock Oracle.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy import select + +from libs.common.config import get_settings +from libs.common.ids import document_id as make_document_id +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import Document, JobRun +from libs.db.session import get_session +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.filings import FilingsService + +logger = get_logger(__name__) + + +async def poll_filings(run_id: str) -> dict[str, int]: + settings = get_settings() + symbols = settings.get_symbols() + app_config = settings.get_app_config() + form_types = ",".join(app_config.get("pipeline", {}).get("form_types", ["8-K", "6-K"])) + + stats = {"seen": 0, "written": 0, "skipped": 0, "errors": 0} + + async with make_oracle_client() as client: + svc = FilingsService(client) + + async with get_session() as session: + # Record job start + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="filing_poller", + source_name="oracle", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + for ticker in symbols: + try: + response = await svc.search_filings( + ticker, + form_type=form_types, + start_date=(dt.date.today() - dt.timedelta(days=7)).isoformat(), + ) + stats["seen"] += len(response.filings) + + for filing in response.filings: + # Check for duplicate + existing = await session.execute( + select(Document).where( + Document.accession_no == filing.accession_no, + Document.form_type == filing.form_type, + ) + ) + if existing.scalar_one_or_none() is not None: + stats["skipped"] += 1 + continue + + doc_id = make_document_id( + "sec", + f"TICKER::{ticker}", + filing.filing_date, + filing.accession_no, + ) + + doc = Document( + document_id=doc_id, + source_name="sec", + accession_no=filing.accession_no, + form_type=filing.form_type, + filing_date=dt.date.fromisoformat(filing.filing_date), + accepted_at_utc=( + dt.datetime.fromisoformat( + filing.accepted_at.replace("Z", "+00:00") + ) + if filing.accepted_at + else None + ), + primary_document_name=filing.primary_document, + parsed_status="pending", + ) + session.add(doc) + stats["written"] += 1 + logger.info( + "new_filing_discovered", + ticker=ticker, + accession_no=filing.accession_no, + form_type=filing.form_type, + filing_date=filing.filing_date, + ) + + except Exception as exc: + logger.error("poll_error", ticker=ticker, error=str(exc)) + stats["errors"] += 1 + + # Update job record + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["written"] + job.records_skipped = stats["skipped"] + job.error_count = stats["errors"] + + logger.info("filing_poller_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Filing Poller") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(poll_filings(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/sync/__init__.py b/apps/sync/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/sync/issuer_sync/__init__.py b/apps/sync/issuer_sync/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/sync/issuer_sync/main.py b/apps/sync/issuer_sync/main.py new file mode 100644 index 0000000..8b002da --- /dev/null +++ b/apps/sync/issuer_sync/main.py @@ -0,0 +1,130 @@ +"""Issuer Sync: populate issuer_master / symbol_master from Oracle company info.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy.dialects.postgresql import insert + +from libs.common.config import get_settings +from libs.common.ids import issuer_id_from_cik, new_job_run_id, symbol_id_from_ticker +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import IssuerMaster, JobRun, SymbolMaster +from libs.db.session import get_session +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.financial import FinancialService + +logger = get_logger(__name__) + + +async def run_issuer_sync(run_id: str) -> dict[str, int]: + settings = get_settings() + symbols = settings.get_symbols() + stats = {"seen": 0, "written": 0, "errors": 0} + + async with make_oracle_client() as client: + svc = FinancialService(client) + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="issuer_sync", + source_name="oracle", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + for ticker in symbols: + stats["seen"] += 1 + try: + info = await svc.get_company_info(ticker) + + cik = info.cik + issuer_id_str = issuer_id_from_cik(cik) if cik else f"ISSUER::{ticker}" + symbol_id_str = symbol_id_from_ticker(ticker, info.exchange or "US") + + # Upsert issuer + issuer_stmt = ( + insert(IssuerMaster) + .values( + issuer_id=issuer_id_str, + cik=cik, + ticker=ticker, + issuer_name=info.name or ticker, + exchange=info.exchange, + country_code=info.country, + is_active=True, + created_at_utc=dt.datetime.now(tz=dt.UTC), + updated_at_utc=dt.datetime.now(tz=dt.UTC), + ) + .on_conflict_do_update( + index_elements=["issuer_id"], + set_={ + "ticker": ticker, + "issuer_name": info.name or ticker, + "exchange": info.exchange, + "updated_at_utc": dt.datetime.now(tz=dt.UTC), + }, + ) + ) + await session.execute(issuer_stmt) + + # Upsert symbol + symbol_stmt = ( + insert(SymbolMaster) + .values( + symbol_id=symbol_id_str, + issuer_id=issuer_id_str, + ticker=ticker, + venue=info.exchange or "US", + asset_type="common_stock", + currency="USD", + is_primary=True, + created_at_utc=dt.datetime.now(tz=dt.UTC), + updated_at_utc=dt.datetime.now(tz=dt.UTC), + ) + .on_conflict_do_update( + index_elements=["symbol_id"], + set_={ + "ticker": ticker, + "updated_at_utc": dt.datetime.now(tz=dt.UTC), + }, + ) + ) + await session.execute(symbol_stmt) + + stats["written"] += 1 + logger.info("issuer_synced", ticker=ticker, issuer_id=issuer_id_str) + + except Exception as exc: + logger.error("issuer_sync_error", ticker=ticker, error=str(exc)) + stats["errors"] += 1 + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["written"] + job.error_count = stats["errors"] + + logger.info("issuer_sync_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Issuer Sync") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(run_issuer_sync(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/sync/macro_sync/__init__.py b/apps/sync/macro_sync/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/sync/macro_sync/main.py b/apps/sync/macro_sync/main.py new file mode 100644 index 0000000..7e0e121 --- /dev/null +++ b/apps/sync/macro_sync/main.py @@ -0,0 +1,139 @@ +"""Macro Sync: fetch FRED series → macro_observations table.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy import select +from sqlalchemy.dialects.postgresql import insert + +from libs.common.config import get_settings +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import JobRun, MacroObservation, MacroSeries, SyncCheckpoint +from libs.db.session import get_session +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.fred import FredService + +logger = get_logger(__name__) + + +async def run_macro_sync(run_id: str) -> dict[str, int]: + settings = get_settings() + fred_series_config = settings.get_fred_series() + stats = {"seen": 0, "written": 0, "errors": 0} + + async with make_oracle_client() as client: + svc = FredService(client) + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="macro_sync", + source_name="fred", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + for series_config in fred_series_config: + series_id = series_config["id"] + try: + # Ensure macro_series row exists + existing = await session.execute( + select(MacroSeries).where(MacroSeries.series_id == series_id) + ) + if existing.scalar_one_or_none() is None: + ms = MacroSeries( + series_id=series_id, + title=series_config.get("title"), + frequency=series_config.get("frequency"), + source_name="fred", + ) + session.add(ms) + await session.flush() + + # Fetch observations (last 2 years) + start = (dt.date.today() - dt.timedelta(days=730)).isoformat() + response = await svc.get_observations(series_id, start=start) + stats["seen"] += len(response.observations) + + rows = [] + for obs in response.observations: + try: + val = float(obs.value) if obs.value is not None else None + except (TypeError, ValueError): + val = None + rows.append({ + "series_id": series_id, + "observation_date": dt.date.fromisoformat(obs.date), + "value": val, + "created_at_utc": dt.datetime.now(tz=dt.UTC), + }) + + if rows: + stmt = ( + insert(MacroObservation) + .values(rows) + .on_conflict_do_nothing( + constraint="uq_macro_obs_series_date" + ) + ) + await session.execute(stmt) + stats["written"] += len(rows) + + # Update checkpoint + cp_stmt = ( + insert(SyncCheckpoint) + .values( + domain=f"fred:{series_id}", + last_sync_at_utc=dt.datetime.now(tz=dt.UTC), + last_sync_params={"start": start}, + status="success", + created_at_utc=dt.datetime.now(tz=dt.UTC), + updated_at_utc=dt.datetime.now(tz=dt.UTC), + ) + .on_conflict_do_update( + constraint="uq_sync_checkpoints_domain", + set_={ + "last_sync_at_utc": dt.datetime.now(tz=dt.UTC), + "status": "success", + "updated_at_utc": dt.datetime.now(tz=dt.UTC), + }, + ) + ) + await session.execute(cp_stmt) + + logger.info("fred_series_synced", series_id=series_id, count=len(rows)) + + except Exception as exc: + logger.error("fred_sync_error", series_id=series_id, error=str(exc)) + stats["errors"] += 1 + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["written"] + job.error_count = stats["errors"] + + logger.info("macro_sync_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Macro Sync") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(run_macro_sync(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/apps/sync/short_volume_sync/__init__.py b/apps/sync/short_volume_sync/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/apps/sync/short_volume_sync/main.py b/apps/sync/short_volume_sync/main.py new file mode 100644 index 0000000..8fce00f --- /dev/null +++ b/apps/sync/short_volume_sync/main.py @@ -0,0 +1,119 @@ +"""Short Volume Sync: fetch FINRA short volume → short_sale_daily.""" +from __future__ import annotations + +import argparse +import asyncio +import datetime as dt +import uuid + +from sqlalchemy.dialects.postgresql import insert + +from libs.common.config import get_settings +from libs.common.ids import new_job_run_id +from libs.common.logging import bind_job_run_id, configure_logging, get_logger +from libs.db.models import JobRun, ShortSaleDaily, SyncCheckpoint +from libs.db.session import get_session +from libs.oracle_client.client import make_oracle_client +from libs.oracle_client.finra import FinraService + +logger = get_logger(__name__) + + +async def run_short_volume_sync(run_id: str) -> dict[str, int]: + settings = get_settings() + symbols = settings.get_symbols() + stats = {"seen": 0, "written": 0, "errors": 0} + + async with make_oracle_client() as client: + svc = FinraService(client) + + async with get_session() as session: + job = JobRun( + job_run_id=uuid.UUID(run_id), + job_name="short_volume_sync", + source_name="finra", + run_date=dt.date.today(), + status="running", + ) + session.add(job) + await session.flush() + + for symbol in symbols: + try: + response = await svc.get_short_volume(symbol, days=30) + stats["seen"] += len(response.data) + + rows = [] + for entry in response.data: + rows.append({ + "ticker_raw": symbol, + "trade_date": dt.date.fromisoformat(entry.date), + "short_volume": entry.short_volume, + "short_exempt_volume": entry.short_exempt_volume, + "total_volume": entry.total_volume, + "source_name": "finra", + "created_at_utc": dt.datetime.now(tz=dt.UTC), + }) + + if rows: + stmt = ( + insert(ShortSaleDaily) + .values(rows) + .on_conflict_do_nothing( + constraint="uq_short_sale_ticker_date_source" + ) + ) + await session.execute(stmt) + stats["written"] += len(rows) + logger.info("short_volume_synced", symbol=symbol, count=len(rows)) + + except Exception as exc: + logger.error("short_volume_error", symbol=symbol, error=str(exc)) + stats["errors"] += 1 + + # Update checkpoint + cp_stmt = ( + insert(SyncCheckpoint) + .values( + domain="finra:short_volume", + last_sync_at_utc=dt.datetime.now(tz=dt.UTC), + last_sync_params={"symbols": symbols}, + status="success", + created_at_utc=dt.datetime.now(tz=dt.UTC), + updated_at_utc=dt.datetime.now(tz=dt.UTC), + ) + .on_conflict_do_update( + constraint="uq_sync_checkpoints_domain", + set_={ + "last_sync_at_utc": dt.datetime.now(tz=dt.UTC), + "status": "success" if stats["errors"] == 0 else "partial", + "updated_at_utc": dt.datetime.now(tz=dt.UTC), + }, + ) + ) + await session.execute(cp_stmt) + + job.status = "succeeded" if stats["errors"] == 0 else "partial" + job.finished_at_utc = dt.datetime.now(tz=dt.UTC) + job.records_seen = stats["seen"] + job.records_written = stats["written"] + job.error_count = stats["errors"] + + logger.info("short_volume_sync_done", **stats) + return stats + + +def main() -> None: + parser = argparse.ArgumentParser(description="Short Volume Sync") + parser.add_argument("--run-id", default=new_job_run_id()) + args = parser.parse_args() + + settings = get_settings() + configure_logging(settings.log_level) + bind_job_run_id(args.run_id) + + asyncio.run(run_short_volume_sync(args.run_id)) + + +if __name__ == "__main__": + main() diff --git a/configs/app.yaml b/configs/app.yaml new file mode 100644 index 0000000..f9bc737 --- /dev/null +++ b/configs/app.yaml @@ -0,0 +1,22 @@ +stock_oracle: + url: "http://localhost:18001" + timeout: 30 + retry_max: 3 + retry_backoff: 2.0 + +polling: + filing_interval_minutes: 30 + sync_interval_hours: 24 + +feature_flags: + llm_enabled: false + finra_enabled: true + fred_enabled: true + +pipeline: + exhibit_types: + - "EX-99.1" + - "EX-99.2" + form_types: + - "8-K" + - "6-K" diff --git a/configs/fred_series.yaml b/configs/fred_series.yaml new file mode 100644 index 0000000..9721d73 --- /dev/null +++ b/configs/fred_series.yaml @@ -0,0 +1,20 @@ +series: + - id: DGS10 + title: "10-Year Treasury Constant Maturity Rate" + frequency: daily + + - id: T10Y2Y + title: "10-Year Treasury Constant Maturity Minus 2-Year" + frequency: daily + + - id: VIXCLS + title: "CBOE Volatility Index: VIX" + frequency: daily + + - id: BAMLH0A0HYM2 + title: "ICE BofA US High Yield Index Option-Adjusted Spread" + frequency: daily + + - id: DGS2 + title: "2-Year Treasury Constant Maturity Rate" + frequency: daily diff --git a/configs/symbols.yaml b/configs/symbols.yaml new file mode 100644 index 0000000..4f78407 --- /dev/null +++ b/configs/symbols.yaml @@ -0,0 +1,16 @@ +symbols: + - AAPL + - MSFT + - GOOGL + - AMZN + - META + - NVDA + - TSLA + - AMD + - NFLX + - CRM + - SNOW + - NET + - DDOG + - ZS + - CRWD diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..a50c1a4 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,28 @@ +version: "3.9" + +services: + postgres: + image: postgres:16-alpine + environment: + POSTGRES_USER: acef + POSTGRES_PASSWORD: acef + POSTGRES_DB: acef + ports: + - "5432:5432" + volumes: + - postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U acef"] + interval: 5s + timeout: 5s + retries: 5 + + adminer: + image: adminer + ports: + - "8080:8080" + depends_on: + - postgres + +volumes: + postgres_data: diff --git a/docker/Dockerfile b/docker/Dockerfile new file mode 100644 index 0000000..6678d2a --- /dev/null +++ b/docker/Dockerfile @@ -0,0 +1,16 @@ +FROM python:3.11-slim + +WORKDIR /app + +COPY pyproject.toml . +COPY libs/ libs/ +COPY apps/ apps/ +COPY configs/ configs/ +COPY alembic.ini . + +RUN pip install --no-cache-dir -e . + +ENV APP_ENV=prod +ENV LOG_LEVEL=INFO + +CMD ["python", "-m", "apps.pipeline.filing_poller.main"] diff --git a/libs/__init__.py b/libs/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/common/__init__.py b/libs/common/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/common/config.py b/libs/common/config.py new file mode 100644 index 0000000..ea3e315 --- /dev/null +++ b/libs/common/config.py @@ -0,0 +1,80 @@ +"""Application configuration via pydantic-settings + YAML merge.""" +from __future__ import annotations + +import functools +from pathlib import Path +from typing import Any + +import yaml +from pydantic import field_validator +from pydantic_settings import BaseSettings, SettingsConfigDict + + +def _load_yaml(path: str | Path) -> dict[str, Any]: + p = Path(path) + if p.exists(): + with open(p) as f: + return yaml.safe_load(f) or {} + return {} + + +class Settings(BaseSettings): + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + case_sensitive=False, + extra="ignore", + ) + + app_env: str = "dev" + + # Stock Oracle + stock_oracle_url: str = "http://localhost:18001" + stock_oracle_timeout: int = 30 + + # Database + postgres_dsn: str = "postgresql+asyncpg://acef:acef@localhost:5432/acef" + + # File storage + data_root: str = "./data" + + # Logging + log_level: str = "INFO" + + # LLM + openai_api_key: str = "sk-placeholder" + llm_enabled: bool = False + + # App YAML overrides (loaded separately) + _app_config: dict[str, Any] = {} + + @field_validator("log_level") + @classmethod + def normalize_log_level(cls, v: str) -> str: + return v.upper() + + @property + def exhibit_cache_dir(self) -> Path: + return Path(self.data_root) / "cache" / "exhibits" + + @property + def parquet_dir(self) -> Path: + return Path(self.data_root) / "parquet" + + def get_app_config(self) -> dict[str, Any]: + if not self._app_config: + object.__setattr__(self, "_app_config", _load_yaml("configs/app.yaml")) + return self._app_config + + def get_symbols(self) -> list[str]: + cfg = _load_yaml("configs/symbols.yaml") + return cfg.get("symbols", []) + + def get_fred_series(self) -> list[dict[str, Any]]: + cfg = _load_yaml("configs/fred_series.yaml") + return cfg.get("series", []) + + +@functools.lru_cache(maxsize=1) +def get_settings() -> Settings: + return Settings() diff --git a/libs/common/file_store.py b/libs/common/file_store.py new file mode 100644 index 0000000..6023e95 --- /dev/null +++ b/libs/common/file_store.py @@ -0,0 +1,77 @@ +"""Local exhibit text cache with atomic write and checksum.""" +from __future__ import annotations + +import json +from pathlib import Path + +from libs.common.config import get_settings +from libs.common.ids import sha256_checksum + + +def _exhibit_dir(accession_no: str) -> Path: + settings = get_settings() + return Path(settings.exhibit_cache_dir) / accession_no + + +def exhibit_path(accession_no: str, exhibit_type: str) -> Path: + """Return canonical path for cached exhibit text.""" + safe_type = exhibit_type.replace("/", "_").replace(" ", "_") + return _exhibit_dir(accession_no) / f"{safe_type}.txt" + + +def sidecar_path(accession_no: str, exhibit_type: str) -> Path: + p = exhibit_path(accession_no, exhibit_type) + return p.with_suffix(".meta.json") + + +def write_exhibit(accession_no: str, exhibit_type: str, content: str) -> str: + """Atomically write exhibit text; return sha256 checksum. + + Raises FileExistsError if file already exists (idempotency guard). + Use exists_exhibit() to check first. + """ + path = exhibit_path(accession_no, exhibit_type) + path.parent.mkdir(parents=True, exist_ok=True) + + checksum = sha256_checksum(content.encode("utf-8")) + + # Atomic write via temp file + tmp = path.with_suffix(".tmp") + try: + tmp.write_text(content, encoding="utf-8") + tmp.rename(path) + except Exception: + tmp.unlink(missing_ok=True) + raise + + # Write sidecar metadata + meta = { + "accession_no": accession_no, + "exhibit_type": exhibit_type, + "content_hash": checksum, + "size_bytes": len(content.encode("utf-8")), + } + sidecar_path(accession_no, exhibit_type).write_text( + json.dumps(meta, indent=2), encoding="utf-8" + ) + + return checksum + + +def read_exhibit(accession_no: str, exhibit_type: str) -> str: + """Read cached exhibit text. Raises FileNotFoundError if missing.""" + path = exhibit_path(accession_no, exhibit_type) + return path.read_text(encoding="utf-8") + + +def exists_exhibit(accession_no: str, exhibit_type: str) -> bool: + return exhibit_path(accession_no, exhibit_type).exists() + + +def get_checksum(accession_no: str, exhibit_type: str) -> str | None: + """Return stored checksum from sidecar, or None if missing.""" + sp = sidecar_path(accession_no, exhibit_type) + if not sp.exists(): + return None + meta = json.loads(sp.read_text(encoding="utf-8")) + return meta.get("content_hash") diff --git a/libs/common/ids.py b/libs/common/ids.py new file mode 100644 index 0000000..971065f --- /dev/null +++ b/libs/common/ids.py @@ -0,0 +1,35 @@ +"""Deterministic ID generators for ACE-F entities.""" +from __future__ import annotations + +import hashlib +import uuid + + +def document_id(source: str, issuer_id: str, date: str, accession: str) -> str: + """DOC::{source}::{issuer_id}::{date}::{accession}""" + return f"DOC::{source}::{issuer_id}::{date}::{accession}" + + +def event_id(document_id_str: str, event_type: str, sequence: int = 0) -> str: + """EVT::{document_id}::{event_type}::{sequence}""" + return f"EVT::{document_id_str}::{event_type}::{sequence}" + + +def issuer_id_from_cik(cik: str) -> str: + return f"ISSUER::{cik.lstrip('0').zfill(10)}" + + +def symbol_id_from_ticker(ticker: str, venue: str = "XNYS") -> str: + return f"SYM::{ticker.upper()}::{venue}" + + +def new_job_run_id() -> str: + return str(uuid.uuid4()) + + +def sha256_checksum(content: bytes) -> str: + return hashlib.sha256(content).hexdigest() + + +def sha256_checksum_str(content: str) -> str: + return sha256_checksum(content.encode("utf-8")) diff --git a/libs/common/logging.py b/libs/common/logging.py new file mode 100644 index 0000000..3665d57 --- /dev/null +++ b/libs/common/logging.py @@ -0,0 +1,55 @@ +"""Structured JSON logging via structlog.""" +from __future__ import annotations + +import logging +import sys +from contextvars import ContextVar +from typing import Any + +import structlog + +_job_run_id: ContextVar[str] = ContextVar("job_run_id", default="") + + +def bind_job_run_id(run_id: str) -> None: + _job_run_id.set(run_id) + + +def _add_job_run_id( + logger: Any, method: str, event_dict: dict[str, Any] +) -> dict[str, Any]: + run_id = _job_run_id.get() + if run_id: + event_dict["job_run_id"] = run_id + return event_dict + + +def configure_logging(level: str = "INFO") -> None: + logging.basicConfig( + format="%(message)s", + stream=sys.stdout, + level=getattr(logging, level.upper(), logging.INFO), + ) + + structlog.configure( + processors=[ + structlog.contextvars.merge_contextvars, + _add_job_run_id, + structlog.stdlib.add_log_level, + structlog.stdlib.add_logger_name, + structlog.processors.TimeStamper(fmt="iso"), + structlog.processors.StackInfoRenderer(), + structlog.processors.format_exc_info, + structlog.processors.JSONRenderer(), + ], + wrapper_class=structlog.make_filtering_bound_logger( + getattr(logging, level.upper(), logging.INFO) + ), + context_class=dict, + logger_factory=structlog.PrintLoggerFactory(), + cache_logger_on_first_use=True, + ) + + +def get_logger(name: str = "") -> structlog.BoundLogger: + return structlog.get_logger(name) diff --git a/libs/common/retries.py b/libs/common/retries.py new file mode 100644 index 0000000..08be506 --- /dev/null +++ b/libs/common/retries.py @@ -0,0 +1,83 @@ +"""Retry wrappers and exception hierarchy.""" +from __future__ import annotations + +import functools +from collections.abc import Callable +from typing import Any, TypeVar + +from tenacity import ( + RetryError, + retry, + retry_if_exception_type, + stop_after_attempt, + wait_exponential, +) + +F = TypeVar("F", bound=Callable[..., Any]) + + +class ACEFError(Exception): + """Base error for ACE-F.""" + + def __init__( + self, + message: str, + source: str = "", + entity: str = "", + context: dict[str, Any] | None = None, + ) -> None: + super().__init__(message) + self.source = source + self.entity = entity + self.context = context or {} + + +class RetryableError(ACEFError): + """Transient error that can be retried (network, timeout, 5xx).""" + + +class NonRetryableError(ACEFError): + """Permanent error that must not be retried (404, business rule).""" + + +class ValidationError(ACEFError): + """Schema or data validation failure — never retried.""" + + +class DependencyError(ACEFError): + """Required upstream dependency unavailable.""" + + +def with_retry( + max_attempts: int = 3, + min_wait: float = 1.0, + max_wait: float = 30.0, + multiplier: float = 2.0, +) -> Callable[[F], F]: + """Decorator: retry on RetryableError with exponential backoff.""" + + def decorator(func: F) -> F: + @retry( + retry=retry_if_exception_type(RetryableError), + stop=stop_after_attempt(max_attempts), + wait=wait_exponential(multiplier=multiplier, min=min_wait, max=max_wait), + reraise=True, + ) + @functools.wraps(func) + async def wrapper(*args: Any, **kwargs: Any) -> Any: + return await func(*args, **kwargs) + + return wrapper # type: ignore[return-value] + + return decorator + + +__all__ = [ + "ACEFError", + "RetryableError", + "NonRetryableError", + "ValidationError", + "DependencyError", + "with_retry", + "RetryError", +] diff --git a/libs/common/time_utils.py b/libs/common/time_utils.py new file mode 100644 index 0000000..bb32f3d --- /dev/null +++ b/libs/common/time_utils.py @@ -0,0 +1,80 @@ +"""UTC/Eastern conversion and NYSE calendar helpers.""" +from __future__ import annotations + +import datetime as dt +from zoneinfo import ZoneInfo + +import exchange_calendars as xcals + +_EASTERN = ZoneInfo("America/New_York") +_UTC = ZoneInfo("UTC") +_XNYS = None + + +def _get_xnys() -> xcals.ExchangeCalendar: + global _XNYS + if _XNYS is None: + _XNYS = xcals.get_calendar("XNYS") + return _XNYS + + +def utc_now() -> dt.datetime: + return dt.datetime.now(tz=_UTC) + + +def to_eastern(d: dt.datetime) -> dt.datetime: + if d.tzinfo is None: + raise ValueError("Naive datetime rejected; must be timezone-aware.") + return d.astimezone(_EASTERN) + + +def to_utc(d: dt.datetime) -> dt.datetime: + if d.tzinfo is None: + raise ValueError("Naive datetime rejected; must be timezone-aware.") + return d.astimezone(_UTC) + + +def is_trading_day(date: dt.date) -> bool: + cal = _get_xnys() + return cal.is_session(date.isoformat()) + + +def previous_trading_day(date: dt.date) -> dt.date: + cal = _get_xnys() + idx = cal.sessions.get_loc(date.isoformat()) if date.isoformat() in cal.sessions else None + if idx is None: + # Find previous session + prev = cal.previous_session(date.isoformat()) + return prev.date() + if idx > 0: + return cal.sessions[idx - 1].date() + raise ValueError(f"No previous trading day before {date}") + + +def next_trading_day(date: dt.date) -> dt.date: + cal = _get_xnys() + return cal.next_session(date.isoformat()).date() + + +def filing_time_bucket(filed_at: dt.datetime) -> str: + """Classify filing time as pre_market, regular_hours, post_market, or unknown.""" + if filed_at.tzinfo is None: + return "unknown" + eastern = to_eastern(filed_at) + hour = eastern.hour + minute = eastern.minute + total_minutes = hour * 60 + minute + # Pre-market: before 9:30 ET + if total_minutes < 9 * 60 + 30: + return "pre_market" + # Regular hours: 9:30–16:00 ET + if total_minutes <= 16 * 60: + return "regular_hours" + # Post-market: after 16:00 ET + return "post_market" + + +def trading_days_between(start: dt.date, end: dt.date) -> list[dt.date]: + cal = _get_xnys() + sessions = cal.sessions_in_range(start.isoformat(), end.isoformat()) + return [s.date() for s in sessions] diff --git a/libs/db/__init__.py b/libs/db/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/db/engine.py b/libs/db/engine.py new file mode 100644 index 0000000..0d4ccdf --- /dev/null +++ b/libs/db/engine.py @@ -0,0 +1,20 @@ +"""Async SQLAlchemy engine singleton.""" +from __future__ import annotations + +import functools + +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine + + +@functools.lru_cache(maxsize=1) +def get_engine(dsn: str | None = None) -> AsyncEngine: + from libs.common.config import get_settings + + target_dsn = dsn or get_settings().postgres_dsn + return create_async_engine( + target_dsn, + echo=False, + pool_pre_ping=True, + pool_size=5, + max_overflow=10, + ) diff --git a/libs/db/enums.py b/libs/db/enums.py new file mode 100644 index 0000000..88c30b4 --- /dev/null +++ b/libs/db/enums.py @@ -0,0 +1,73 @@ +"""Database enumerations.""" +from __future__ import annotations + +import enum + + +class JobStatus(str, enum.Enum): + pending = "pending" + running = "running" + succeeded = "succeeded" + failed = "failed" + partial = "partial" + + +class ParsedStatus(str, enum.Enum): + pending = "pending" + ready_for_parse = "ready_for_parse" + succeeded = "succeeded" + failed = "failed" + + +class EventStatus(str, enum.Enum): + pending = "pending" + valid = "valid" + rejected = "rejected" + + +class EventDirection(str, enum.Enum): + bullish = "bullish" + bearish = "bearish" + mixed = "mixed" + neutral = "neutral" + unknown = "unknown" + + +class EventType(str, enum.Enum): + earnings_release = "earnings_release" + guidance_update = "guidance_update" + material_contract = "material_contract" + regulatory_or_approval = "regulatory_or_approval" + capital_markets_or_financing = "capital_markets_or_financing" + management_change = "management_change" + litigation_or_investigation = "litigation_or_investigation" + other_material_event = "other_material_event" + unknown = "unknown" + + +class ParserKind(str, enum.Enum): + rule = "rule" + llm = "llm" + merged = "merged" + + +class ValidationStatus(str, enum.Enum): + valid = "valid" + invalid = "invalid" + + +class OrderSide(str, enum.Enum): + buy = "buy" + sell = "sell" + + +class OrderPlanStatus(str, enum.Enum): + draft = "draft" + ready = "ready" + canceled = "canceled" + + +class SyncStatus(str, enum.Enum): + success = "success" + failed = "failed" + partial = "partial" diff --git a/libs/db/helpers.py b/libs/db/helpers.py new file mode 100644 index 0000000..ef6bf34 --- /dev/null +++ b/libs/db/helpers.py @@ -0,0 +1,79 @@ +"""DB helper utilities: upsert, bulk operations, health check.""" +from __future__ import annotations + +from typing import Any + +from sqlalchemy import select, text +from sqlalchemy.dialects.postgresql import insert +from sqlalchemy.ext.asyncio import AsyncSession + +from libs.db.models import Base + + +async def upsert( + session: AsyncSession, + model: type[Base], + values: dict[str, Any], + index_elements: list[str], + update_columns: list[str] | None = None, +) -> None: + """Insert or update a single row using PostgreSQL ON CONFLICT.""" + stmt = insert(model).values(**values) + if update_columns: + update_dict = {col: getattr(stmt.excluded, col) for col in update_columns} + stmt = stmt.on_conflict_do_update( + index_elements=index_elements, set_=update_dict + ) + else: + stmt = stmt.on_conflict_do_nothing(index_elements=index_elements) + await session.execute(stmt) + + +async def bulk_upsert( + session: AsyncSession, + model: type[Base], + rows: list[dict[str, Any]], + index_elements: list[str], + update_columns: list[str] | None = None, +) -> int: + """Bulk upsert; returns number of rows processed.""" + if not rows: + return 0 + stmt = insert(model).values(rows) + if update_columns: + update_dict = {col: getattr(stmt.excluded, col) for col in update_columns} + stmt = stmt.on_conflict_do_update( + index_elements=index_elements, set_=update_dict + ) + else: + stmt = stmt.on_conflict_do_nothing(index_elements=index_elements) + await session.execute(stmt) + return len(rows) + + +async def get_or_create( + session: AsyncSession, + model: type[Base], + pk_value: Any, + pk_column: str, + defaults: dict[str, Any], +) -> tuple[Any, bool]: + """Return (instance, created). Upserts if not found.""" + col = getattr(model, pk_column) + result = await session.execute(select(model).where(col == pk_value)) + row = result.scalar_one_or_none() + if row is not None: + return row, False + instance = model(**{pk_column: pk_value, **defaults}) + session.add(instance) + await session.flush() + return instance, True + + +async def health_check(session: AsyncSession) -> bool: + """Return True if DB is reachable.""" + try: + await session.execute(text("SELECT 1")) + return True + except Exception: + return False diff --git a/libs/db/migrations/__init__.py b/libs/db/migrations/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/db/migrations/env.py b/libs/db/migrations/env.py new file mode 100644 index 0000000..c35e7ed --- /dev/null +++ b/libs/db/migrations/env.py @@ -0,0 +1,66 @@ +"""Alembic environment configuration.""" +from __future__ import annotations + +import asyncio +from logging.config import fileConfig + +from alembic import context +from sqlalchemy import pool +from sqlalchemy.engine import Connection +from sqlalchemy.ext.asyncio import async_engine_from_config + +from libs.db.models import Base # noqa: F401 -- ensures all models are registered + +config = context.config + +if config.config_file_name is not None: + fileConfig(config.config_file_name) + +target_metadata = Base.metadata + + +def get_url() -> str: + from libs.common.config import get_settings + + return get_settings().postgres_dsn + + +def run_migrations_offline() -> None: + url = get_url() + context.configure( + url=url, + target_metadata=target_metadata, + literal_binds=True, + dialect_opts={"paramstyle": "named"}, + ) + with context.begin_transaction(): + context.run_migrations() + + +def do_run_migrations(connection: Connection) -> None: + context.configure(connection=connection, target_metadata=target_metadata) + with context.begin_transaction(): + context.run_migrations() + + +async def run_async_migrations() -> None: + cfg = config.get_section(config.config_ini_section) or {} + cfg["sqlalchemy.url"] = get_url() + connectable = async_engine_from_config( + cfg, + prefix="sqlalchemy.", + poolclass=pool.NullPool, + ) + async with connectable.connect() as connection: + await connection.run_sync(do_run_migrations) + await connectable.dispose() + + +def run_migrations_online() -> None: + asyncio.run(run_async_migrations()) + + +if context.is_offline_mode(): + run_migrations_offline() +else: + run_migrations_online() diff --git a/libs/db/migrations/script.py.mako b/libs/db/migrations/script.py.mako new file mode 100644 index 0000000..ee746cf --- /dev/null +++ b/libs/db/migrations/script.py.mako @@ -0,0 +1,28 @@ +"""${message} + +Revision ID: ${up_revision} +Revises: ${down_revision | comma,n} +Create Date: ${create_date} + +""" +from __future__ import annotations + +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +${imports if imports else ""} + +# revision identifiers, used by Alembic. +revision: str = ${repr(up_revision)} +down_revision: Union[str, None] = ${repr(down_revision)} +branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} +depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} + + +def upgrade() -> None: + ${upgrades if upgrades else "pass"} + + +def downgrade() -> None: + ${downgrades if downgrades else "pass"} diff --git a/libs/db/migrations/versions/0001_initial_schema.py b/libs/db/migrations/versions/0001_initial_schema.py new file mode 100644 index 0000000..9c2b66e --- /dev/null +++ b/libs/db/migrations/versions/0001_initial_schema.py @@ -0,0 +1,414 @@ +"""Initial schema: all ACE-F v1 tables. + +Revision ID: 0001 +Revises: +Create Date: 2026-03-12 + +Tables created: +- issuer_master, symbol_master: company reference data +- job_runs: pipeline execution tracking +- documents, document_exhibits: SEC filing metadata +- exhibit_cache: local exhibit text cache metadata +- events, event_parses: parsed event data +- macro_series, macro_observations: FRED macro data +- short_sale_daily: FINRA short volume data +- feature_snapshots: computed ML features +- order_plans: future trade plans (schema only in Phase 1) +- sync_checkpoints: FRED/FINRA sync state +""" +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects.postgresql import JSONB, UUID + +revision: str = "0001" +down_revision: str | None = None +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "issuer_master", + sa.Column("issuer_id", sa.Text, primary_key=True), + sa.Column("cik", sa.Text, nullable=True, unique=True), + sa.Column("ticker", sa.Text, nullable=True), + sa.Column("issuer_name", sa.Text, nullable=False), + sa.Column("exchange", sa.Text, nullable=True), + sa.Column("country_code", sa.Text, nullable=True), + sa.Column("is_active", sa.Boolean, nullable=False, server_default="true"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column( + "updated_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "symbol_master", + sa.Column("symbol_id", sa.Text, primary_key=True), + sa.Column( + "issuer_id", + sa.Text, + sa.ForeignKey("issuer_master.issuer_id"), + nullable=True, + ), + sa.Column("ticker", sa.Text, nullable=False), + sa.Column("venue", sa.Text, nullable=True), + sa.Column("asset_type", sa.Text, nullable=True), + sa.Column("currency", sa.Text, nullable=True), + sa.Column("start_date", sa.Date, nullable=True), + sa.Column("end_date", sa.Date, nullable=True), + sa.Column("is_primary", sa.Boolean, nullable=False, server_default="true"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column( + "updated_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "job_runs", + sa.Column("job_run_id", UUID(as_uuid=True), primary_key=True), + sa.Column("job_name", sa.Text, nullable=False), + sa.Column("source_name", sa.Text, nullable=True), + sa.Column("run_date", sa.Date, nullable=True), + sa.Column("status", sa.Text, nullable=False), + sa.Column( + "started_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column("finished_at_utc", sa.DateTime(timezone=True), nullable=True), + sa.Column("records_seen", sa.Integer, nullable=False, server_default="0"), + sa.Column("records_written", sa.Integer, nullable=False, server_default="0"), + sa.Column("records_skipped", sa.Integer, nullable=False, server_default="0"), + sa.Column("error_count", sa.Integer, nullable=False, server_default="0"), + sa.Column("error_summary", sa.Text, nullable=True), + sa.Column("metadata_json", JSONB, nullable=False, server_default="'{}'"), + ) + op.create_index("ix_job_runs_job_name_run_date", "job_runs", ["job_name", "run_date"]) + + op.create_table( + "documents", + sa.Column("document_id", sa.Text, primary_key=True), + sa.Column("source_name", sa.Text, nullable=False), + sa.Column( + "issuer_id", + sa.Text, + sa.ForeignKey("issuer_master.issuer_id"), + nullable=True, + ), + sa.Column( + "symbol_id", + sa.Text, + sa.ForeignKey("symbol_master.symbol_id"), + nullable=True, + ), + sa.Column("accession_no", sa.Text, nullable=True), + sa.Column("form_type", sa.Text, nullable=False), + sa.Column("filing_date", sa.Date, nullable=False), + sa.Column("accepted_at_utc", sa.DateTime(timezone=True), nullable=True), + sa.Column("primary_document_name", sa.Text, nullable=True), + sa.Column("parsed_status", sa.Text, nullable=False, server_default="'pending'"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column( + "updated_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.UniqueConstraint("accession_no", "form_type", name="uq_documents_accession_form"), + ) + op.create_index( + "ix_documents_issuer_filing_date", "documents", ["issuer_id", "filing_date"] + ) + op.create_index( + "ix_documents_form_type_filing_date", "documents", ["form_type", "filing_date"] + ) + + op.create_table( + "document_exhibits", + sa.Column("exhibit_id", sa.Text, primary_key=True), + sa.Column( + "document_id", + sa.Text, + sa.ForeignKey("documents.document_id"), + nullable=False, + ), + sa.Column("exhibit_code", sa.Text, nullable=False), + sa.Column("exhibit_name", sa.Text, nullable=True), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "exhibit_cache", + sa.Column("id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column("accession_no", sa.Text, nullable=False), + sa.Column("exhibit_type", sa.Text, nullable=False), + sa.Column("content_hash", sa.Text, nullable=False), + sa.Column("cache_path", sa.Text, nullable=False), + sa.Column( + "fetched_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.UniqueConstraint( + "accession_no", "exhibit_type", name="uq_exhibit_cache_accession_type" + ), + ) + + op.create_table( + "events", + sa.Column("event_id", sa.Text, primary_key=True), + sa.Column( + "issuer_id", + sa.Text, + sa.ForeignKey("issuer_master.issuer_id"), + nullable=True, + ), + sa.Column( + "symbol_id", + sa.Text, + sa.ForeignKey("symbol_master.symbol_id"), + nullable=True, + ), + sa.Column( + "primary_document_id", + sa.Text, + sa.ForeignKey("documents.document_id"), + nullable=False, + ), + sa.Column("event_type", sa.Text, nullable=False), + sa.Column("event_direction", sa.Text, nullable=False), + sa.Column("event_date", sa.Date, nullable=False), + sa.Column("filed_at_utc", sa.DateTime(timezone=True), nullable=True), + sa.Column("parser_version", sa.Text, nullable=False), + sa.Column("parse_confidence", sa.Numeric, nullable=True), + sa.Column("status", sa.Text, nullable=False, server_default="'pending'"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column( + "updated_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + op.create_index("ix_events_event_type_date", "events", ["event_type", "event_date"]) + op.create_index( + "ix_events_primary_document_id", "events", ["primary_document_id"] + ) + + op.create_table( + "event_parses", + sa.Column("event_parse_id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column( + "event_id", sa.Text, sa.ForeignKey("events.event_id"), nullable=False + ), + sa.Column("parser_kind", sa.Text, nullable=False), + sa.Column("parser_version", sa.Text, nullable=False), + sa.Column("schema_version", sa.Text, nullable=False), + sa.Column("output_json", JSONB, nullable=False), + sa.Column("validation_status", sa.Text, nullable=False), + sa.Column("validation_errors", JSONB, nullable=True), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "macro_series", + sa.Column("series_id", sa.Text, primary_key=True), + sa.Column("title", sa.Text, nullable=True), + sa.Column("frequency", sa.Text, nullable=True), + sa.Column("units", sa.Text, nullable=True), + sa.Column("source_name", sa.Text, nullable=False), + sa.Column("metadata_json", JSONB, nullable=False, server_default="'{}'"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "macro_observations", + sa.Column("id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column( + "series_id", + sa.Text, + sa.ForeignKey("macro_series.series_id"), + nullable=False, + ), + sa.Column("observation_date", sa.Date, nullable=False), + sa.Column("value", sa.Numeric, nullable=True), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.UniqueConstraint( + "series_id", "observation_date", name="uq_macro_obs_series_date" + ), + ) + + op.create_table( + "short_sale_daily", + sa.Column("id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column( + "symbol_id", + sa.Text, + sa.ForeignKey("symbol_master.symbol_id"), + nullable=True, + ), + sa.Column("ticker_raw", sa.Text, nullable=False), + sa.Column("trade_date", sa.Date, nullable=False), + sa.Column("short_volume", sa.BigInteger, nullable=False), + sa.Column("short_exempt_volume", sa.BigInteger, nullable=True), + sa.Column("total_volume", sa.BigInteger, nullable=True), + sa.Column("source_name", sa.Text, nullable=False), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.UniqueConstraint( + "ticker_raw", + "trade_date", + "source_name", + name="uq_short_sale_ticker_date_source", + ), + ) + op.create_index( + "ix_short_sale_ticker_date", "short_sale_daily", ["ticker_raw", "trade_date"] + ) + + op.create_table( + "feature_snapshots", + sa.Column("feature_snapshot_id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column( + "event_id", sa.Text, sa.ForeignKey("events.event_id"), nullable=False + ), + sa.Column("snapshot_name", sa.Text, nullable=False), + sa.Column("snapshot_version", sa.Text, nullable=False), + sa.Column("feature_json", JSONB, nullable=False), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "order_plans", + sa.Column("order_plan_id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column( + "event_id", sa.Text, sa.ForeignKey("events.event_id"), nullable=False + ), + sa.Column( + "symbol_id", + sa.Text, + sa.ForeignKey("symbol_master.symbol_id"), + nullable=False, + ), + sa.Column("side", sa.Text, nullable=False), + sa.Column("planned_entry_date", sa.Date, nullable=False), + sa.Column("planned_order_type", sa.Text, nullable=False), + sa.Column("planned_price", sa.Numeric, nullable=True), + sa.Column("stop_price", sa.Numeric, nullable=True), + sa.Column("take_profit_price", sa.Numeric, nullable=True), + sa.Column("quantity_plan", sa.Numeric, nullable=True), + sa.Column("status", sa.Text, nullable=False, server_default="'draft'"), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + ) + + op.create_table( + "sync_checkpoints", + sa.Column("id", sa.BigInteger, primary_key=True, autoincrement=True), + sa.Column("domain", sa.Text, nullable=False), + sa.Column("last_sync_at_utc", sa.DateTime(timezone=True), nullable=True), + sa.Column("last_sync_params", JSONB, nullable=False, server_default="'{}'"), + sa.Column("status", sa.Text, nullable=False), + sa.Column( + "created_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.Column( + "updated_at_utc", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.text("now()"), + ), + sa.UniqueConstraint("domain", name="uq_sync_checkpoints_domain"), + ) + + +def downgrade() -> None: + op.drop_table("sync_checkpoints") + op.drop_table("order_plans") + op.drop_table("feature_snapshots") + op.drop_index("ix_short_sale_ticker_date", table_name="short_sale_daily") + op.drop_table("short_sale_daily") + op.drop_table("macro_observations") + op.drop_table("macro_series") + op.drop_table("event_parses") + op.drop_index("ix_events_primary_document_id", table_name="events") + op.drop_index("ix_events_event_type_date", table_name="events") + op.drop_table("events") + op.drop_table("exhibit_cache") + op.drop_table("document_exhibits") + op.drop_index("ix_documents_form_type_filing_date", table_name="documents") + op.drop_index("ix_documents_issuer_filing_date", table_name="documents") + op.drop_table("documents") + op.drop_index("ix_job_runs_job_name_run_date", table_name="job_runs") + op.drop_table("job_runs") + op.drop_table("symbol_master") + op.drop_table("issuer_master") diff --git a/libs/db/models.py b/libs/db/models.py new file mode 100644 index 0000000..5a227bf --- /dev/null +++ b/libs/db/models.py @@ -0,0 +1,342 @@ +"""SQLAlchemy 2.0 declarative models for ACE-F.""" +from __future__ import annotations + +import datetime as dt +import uuid + +from sqlalchemy import ( + BigInteger, + Boolean, + Date, + DateTime, + ForeignKey, + Index, + Integer, + Numeric, + Text, + UniqueConstraint, +) +from sqlalchemy.dialects.postgresql import JSONB, UUID +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship + + +def _utcnow() -> dt.datetime: + return dt.datetime.now(tz=dt.UTC) + + +class Base(DeclarativeBase): + pass + + +class IssuerMaster(Base): + __tablename__ = "issuer_master" + + issuer_id: Mapped[str] = mapped_column(Text, primary_key=True) + cik: Mapped[str | None] = mapped_column(Text, unique=True, nullable=True) + ticker: Mapped[str | None] = mapped_column(Text, nullable=True) + issuer_name: Mapped[str] = mapped_column(Text, nullable=False) + exchange: Mapped[str | None] = mapped_column(Text, nullable=True) + country_code: Mapped[str | None] = mapped_column(Text, nullable=True) + is_active: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + updated_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, onupdate=_utcnow, nullable=False + ) + + symbols: Mapped[list[SymbolMaster]] = relationship(back_populates="issuer") + + +class SymbolMaster(Base): + __tablename__ = "symbol_master" + + symbol_id: Mapped[str] = mapped_column(Text, primary_key=True) + issuer_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("issuer_master.issuer_id"), nullable=True + ) + ticker: Mapped[str] = mapped_column(Text, nullable=False) + venue: Mapped[str | None] = mapped_column(Text, nullable=True) + asset_type: Mapped[str | None] = mapped_column(Text, nullable=True) + currency: Mapped[str | None] = mapped_column(Text, nullable=True) + start_date: Mapped[dt.date | None] = mapped_column(Date, nullable=True) + end_date: Mapped[dt.date | None] = mapped_column(Date, nullable=True) + is_primary: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + updated_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, onupdate=_utcnow, nullable=False + ) + + issuer: Mapped[IssuerMaster | None] = relationship(back_populates="symbols") + + +class JobRun(Base): + __tablename__ = "job_runs" + __table_args__ = (Index("ix_job_runs_job_name_run_date", "job_name", "run_date"),) + + job_run_id: Mapped[uuid.UUID] = mapped_column( + UUID(as_uuid=True), primary_key=True, default=uuid.uuid4 + ) + job_name: Mapped[str] = mapped_column(Text, nullable=False) + source_name: Mapped[str | None] = mapped_column(Text, nullable=True) + run_date: Mapped[dt.date | None] = mapped_column(Date, nullable=True) + status: Mapped[str] = mapped_column(Text, nullable=False) + started_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + finished_at_utc: Mapped[dt.datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + records_seen: Mapped[int] = mapped_column(Integer, default=0, nullable=False) + records_written: Mapped[int] = mapped_column(Integer, default=0, nullable=False) + records_skipped: Mapped[int] = mapped_column(Integer, default=0, nullable=False) + error_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False) + error_summary: Mapped[str | None] = mapped_column(Text, nullable=True) + metadata_json: Mapped[dict] = mapped_column(JSONB, default=dict, nullable=False) + + +class Document(Base): + __tablename__ = "documents" + __table_args__ = ( + UniqueConstraint("accession_no", "form_type", name="uq_documents_accession_form"), + Index("ix_documents_issuer_filing_date", "issuer_id", "filing_date"), + Index("ix_documents_form_type_filing_date", "form_type", "filing_date"), + ) + + document_id: Mapped[str] = mapped_column(Text, primary_key=True) + source_name: Mapped[str] = mapped_column(Text, nullable=False) + issuer_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("issuer_master.issuer_id"), nullable=True + ) + symbol_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("symbol_master.symbol_id"), nullable=True + ) + accession_no: Mapped[str | None] = mapped_column(Text, nullable=True) + form_type: Mapped[str] = mapped_column(Text, nullable=False) + filing_date: Mapped[dt.date] = mapped_column(Date, nullable=False) + accepted_at_utc: Mapped[dt.datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + primary_document_name: Mapped[str | None] = mapped_column(Text, nullable=True) + parsed_status: Mapped[str] = mapped_column(Text, default="pending", nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + updated_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, onupdate=_utcnow, nullable=False + ) + + exhibits: Mapped[list[DocumentExhibit]] = relationship(back_populates="document") + events: Mapped[list[Event]] = relationship(back_populates="primary_document") + + +class DocumentExhibit(Base): + __tablename__ = "document_exhibits" + + exhibit_id: Mapped[str] = mapped_column(Text, primary_key=True) + document_id: Mapped[str] = mapped_column( + Text, ForeignKey("documents.document_id"), nullable=False + ) + exhibit_code: Mapped[str] = mapped_column(Text, nullable=False) + exhibit_name: Mapped[str | None] = mapped_column(Text, nullable=True) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + document: Mapped[Document] = relationship(back_populates="exhibits") + + +class ExhibitCache(Base): + __tablename__ = "exhibit_cache" + __table_args__ = ( + UniqueConstraint("accession_no", "exhibit_type", name="uq_exhibit_cache_accession_type"), + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + accession_no: Mapped[str] = mapped_column(Text, nullable=False) + exhibit_type: Mapped[str] = mapped_column(Text, nullable=False) + content_hash: Mapped[str] = mapped_column(Text, nullable=False) + cache_path: Mapped[str] = mapped_column(Text, nullable=False) + fetched_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + +class Event(Base): + __tablename__ = "events" + __table_args__ = ( + Index("ix_events_event_type_date", "event_type", "event_date"), + Index("ix_events_primary_document_id", "primary_document_id"), + ) + + event_id: Mapped[str] = mapped_column(Text, primary_key=True) + issuer_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("issuer_master.issuer_id"), nullable=True + ) + symbol_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("symbol_master.symbol_id"), nullable=True + ) + primary_document_id: Mapped[str] = mapped_column( + Text, ForeignKey("documents.document_id"), nullable=False + ) + event_type: Mapped[str] = mapped_column(Text, nullable=False) + event_direction: Mapped[str] = mapped_column(Text, nullable=False) + event_date: Mapped[dt.date] = mapped_column(Date, nullable=False) + filed_at_utc: Mapped[dt.datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + parser_version: Mapped[str] = mapped_column(Text, nullable=False) + parse_confidence: Mapped[float | None] = mapped_column(Numeric, nullable=True) + status: Mapped[str] = mapped_column(Text, default="pending", nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + updated_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, onupdate=_utcnow, nullable=False + ) + + primary_document: Mapped[Document] = relationship(back_populates="events") + parses: Mapped[list[EventParse]] = relationship(back_populates="event") + feature_snapshots: Mapped[list[FeatureSnapshot]] = relationship(back_populates="event") + + +class EventParse(Base): + __tablename__ = "event_parses" + + event_parse_id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + event_id: Mapped[str] = mapped_column( + Text, ForeignKey("events.event_id"), nullable=False + ) + parser_kind: Mapped[str] = mapped_column(Text, nullable=False) + parser_version: Mapped[str] = mapped_column(Text, nullable=False) + schema_version: Mapped[str] = mapped_column(Text, nullable=False) + output_json: Mapped[dict] = mapped_column(JSONB, nullable=False) + validation_status: Mapped[str] = mapped_column(Text, nullable=False) + validation_errors: Mapped[dict | None] = mapped_column(JSONB, nullable=True) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + event: Mapped[Event] = relationship(back_populates="parses") + + +class MacroSeries(Base): + __tablename__ = "macro_series" + + series_id: Mapped[str] = mapped_column(Text, primary_key=True) + title: Mapped[str | None] = mapped_column(Text, nullable=True) + frequency: Mapped[str | None] = mapped_column(Text, nullable=True) + units: Mapped[str | None] = mapped_column(Text, nullable=True) + source_name: Mapped[str] = mapped_column(Text, nullable=False) + metadata_json: Mapped[dict] = mapped_column(JSONB, default=dict, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + observations: Mapped[list[MacroObservation]] = relationship(back_populates="series") + + +class MacroObservation(Base): + __tablename__ = "macro_observations" + __table_args__ = ( + UniqueConstraint("series_id", "observation_date", name="uq_macro_obs_series_date"), + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + series_id: Mapped[str] = mapped_column( + Text, ForeignKey("macro_series.series_id"), nullable=False + ) + observation_date: Mapped[dt.date] = mapped_column(Date, nullable=False) + value: Mapped[float | None] = mapped_column(Numeric, nullable=True) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + series: Mapped[MacroSeries] = relationship(back_populates="observations") + + +class ShortSaleDaily(Base): + __tablename__ = "short_sale_daily" + __table_args__ = ( + UniqueConstraint( + "ticker_raw", "trade_date", "source_name", name="uq_short_sale_ticker_date_source" + ), + Index("ix_short_sale_ticker_date", "ticker_raw", "trade_date"), + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + symbol_id: Mapped[str | None] = mapped_column( + Text, ForeignKey("symbol_master.symbol_id"), nullable=True + ) + ticker_raw: Mapped[str] = mapped_column(Text, nullable=False) + trade_date: Mapped[dt.date] = mapped_column(Date, nullable=False) + short_volume: Mapped[int] = mapped_column(BigInteger, nullable=False) + short_exempt_volume: Mapped[int | None] = mapped_column(BigInteger, nullable=True) + total_volume: Mapped[int | None] = mapped_column(BigInteger, nullable=True) + source_name: Mapped[str] = mapped_column(Text, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + +class FeatureSnapshot(Base): + __tablename__ = "feature_snapshots" + + feature_snapshot_id: Mapped[int] = mapped_column( + BigInteger, primary_key=True, autoincrement=True + ) + event_id: Mapped[str] = mapped_column( + Text, ForeignKey("events.event_id"), nullable=False + ) + snapshot_name: Mapped[str] = mapped_column(Text, nullable=False) + snapshot_version: Mapped[str] = mapped_column(Text, nullable=False) + feature_json: Mapped[dict] = mapped_column(JSONB, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + event: Mapped[Event] = relationship(back_populates="feature_snapshots") + + +class OrderPlan(Base): + __tablename__ = "order_plans" + + order_plan_id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + event_id: Mapped[str] = mapped_column( + Text, ForeignKey("events.event_id"), nullable=False + ) + symbol_id: Mapped[str] = mapped_column( + Text, ForeignKey("symbol_master.symbol_id"), nullable=False + ) + side: Mapped[str] = mapped_column(Text, nullable=False) + planned_entry_date: Mapped[dt.date] = mapped_column(Date, nullable=False) + planned_order_type: Mapped[str] = mapped_column(Text, nullable=False) + planned_price: Mapped[float | None] = mapped_column(Numeric, nullable=True) + stop_price: Mapped[float | None] = mapped_column(Numeric, nullable=True) + take_profit_price: Mapped[float | None] = mapped_column(Numeric, nullable=True) + quantity_plan: Mapped[float | None] = mapped_column(Numeric, nullable=True) + status: Mapped[str] = mapped_column(Text, default="draft", nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + + +class SyncCheckpoint(Base): + __tablename__ = "sync_checkpoints" + __table_args__ = (UniqueConstraint("domain", name="uq_sync_checkpoints_domain"),) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + domain: Mapped[str] = mapped_column(Text, nullable=False) + last_sync_at_utc: Mapped[dt.datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + last_sync_params: Mapped[dict] = mapped_column(JSONB, default=dict, nullable=False) + status: Mapped[str] = mapped_column(Text, nullable=False) + created_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, nullable=False + ) + updated_at_utc: Mapped[dt.datetime] = mapped_column( + DateTime(timezone=True), default=_utcnow, onupdate=_utcnow, nullable=False + ) diff --git a/libs/db/session.py b/libs/db/session.py new file mode 100644 index 0000000..6aa22bf --- /dev/null +++ b/libs/db/session.py @@ -0,0 +1,33 @@ +"""Async session factory and context manager.""" +from __future__ import annotations + +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager + +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker + +from libs.db.engine import get_engine + +_session_factory: async_sessionmaker[AsyncSession] | None = None + + +def get_session_factory(dsn: str | None = None) -> async_sessionmaker[AsyncSession]: + global _session_factory + if _session_factory is None: + engine = get_engine(dsn) + _session_factory = async_sessionmaker( + engine, expire_on_commit=False, class_=AsyncSession + ) + return _session_factory + + +@asynccontextmanager +async def get_session(dsn: str | None = None) -> AsyncGenerator[AsyncSession, None]: + factory = get_session_factory(dsn) + async with factory() as session: + try: + yield session + await session.commit() + except Exception: + await session.rollback() + raise diff --git a/libs/features/__init__.py b/libs/features/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/features/builder.py b/libs/features/builder.py new file mode 100644 index 0000000..9743964 --- /dev/null +++ b/libs/features/builder.py @@ -0,0 +1,98 @@ +"""Feature orchestrator: combines market and event features.""" +from __future__ import annotations + +import datetime as dt + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from libs.common.logging import get_logger +from libs.db.models import Event, EventParse, FeatureSnapshot +from libs.features.event_features import compute_event_features +from libs.features.market_features import compute_market_features +from libs.oracle_client.price import PriceService + +logger = get_logger(__name__) + +SNAPSHOT_VERSION = "1.0.0" + + +async def build_features_for_event( + session: AsyncSession, + event: Event, + price_service: PriceService, +) -> tuple[FeatureSnapshot, FeatureSnapshot] | None: + """Build market_v1 and event_v1 feature snapshots for an event. + + Returns (market_snapshot, event_snapshot) or None on failure. + """ + # Get ticker from event (via symbol) + if not event.symbol_id: + logger.warning("event_no_symbol", event_id=event.event_id) + return None + + # Extract ticker from symbol_id (format: SYM::{ticker}::{venue}) + parts = event.symbol_id.split("::") + ticker = parts[1] if len(parts) >= 2 else None + if not ticker: + logger.warning("event_bad_symbol_id", event_id=event.event_id, symbol_id=event.symbol_id) + return None + + event_date_str = event.event_date.isoformat() + + # Fetch price bars from Stock Oracle + try: + # 30 trading days before event + start_date = (event.event_date - dt.timedelta(days=45)).isoformat() + end_date = (event.event_date + dt.timedelta(days=5)).isoformat() + price_response = await price_service.get_daily_bars( + ticker, start=start_date, end=end_date + ) + bars = price_response.bars + except Exception as exc: + logger.error("price_fetch_failed", event_id=event.event_id, error=str(exc)) + return None + + # Compute market features + mf = compute_market_features(bars, event_date_str) + + # Get latest valid event parse + result = await session.execute( + select(EventParse) + .where(EventParse.event_id == event.event_id) + .where(EventParse.validation_status == "valid") + .order_by(EventParse.event_parse_id.desc()) + .limit(1) + ) + parse = result.scalar_one_or_none() + + if parse is None: + logger.warning("no_valid_parse", event_id=event.event_id) + return None + + ef = compute_event_features(parse.output_json) + + market_snapshot = FeatureSnapshot( + event_id=event.event_id, + snapshot_name="market_v1", + snapshot_version=SNAPSHOT_VERSION, + feature_json=mf, + ) + event_snapshot = FeatureSnapshot( + event_id=event.event_id, + snapshot_name="event_v1", + snapshot_version=SNAPSHOT_VERSION, + feature_json=ef, + ) + + session.add(market_snapshot) + session.add(event_snapshot) + await session.flush() + + logger.info( + "features_built", + event_id=event.event_id, + ticker=ticker, + market_features=list(mf.keys()), + ) + return market_snapshot, event_snapshot diff --git a/libs/features/event_features.py b/libs/features/event_features.py new file mode 100644 index 0000000..98061ec --- /dev/null +++ b/libs/features/event_features.py @@ -0,0 +1,88 @@ +"""Event-level feature calculations from parser output.""" +from __future__ import annotations + +from typing import Any + +from libs.schemas.types import ConfidenceOutput, GuidanceOutput, RiskFlagsOutput, SignalsOutput + + +def guidance_direction_score(guidance: GuidanceOutput) -> float: + """raised=1.0, inline_or_maintained=0.5, lowered=0.0, else=0.25.""" + mapping = { + "raised": 1.0, + "inline_or_maintained": 0.5, + "lowered": 0.0, + "withdrawn": 0.0, + "not_provided": 0.25, + "unclear": 0.25, + } + return mapping.get(guidance.status, 0.25) + + +def oneoff_penalty(risk_flags: RiskFlagsOutput) -> float: + """Sum of active risk flags / total flags. Higher = more risk.""" + flags = [ + risk_flags.oneoff_item, + risk_flags.tax_benefit, + risk_flags.valuation_gain, + risk_flags.non_gaap_heavy, + risk_flags.financing_related, + risk_flags.legal_or_regulatory_overhang, + ] + total = len(flags) + active = sum(flags) + return active / total if total > 0 else 0.0 + + +def signal_strength_score(signals: SignalsOutput) -> float: + """Composite score [0..1] of positive business signals.""" + score = 0.0 + weights = { + "demand_strength": {"strong": 1.0, "stable": 0.5, "weakening": 0.0, "unknown": 0.0}, + "pricing_power": {"present": 1.0, "mixed": 0.5, "absent": 0.0, "unknown": 0.0}, + "backlog_or_bookings": {"present": 1.0, "mixed": 0.5, "absent": 0.0, "unknown": 0.0}, + "customer_expansion": {"present": 1.0, "mixed": 0.5, "absent": 0.0, "unknown": 0.0}, + "margin_quality": {"improving": 1.0, "stable": 0.5, "deteriorating": 0.0, "unknown": 0.0}, + } + total_weight = len(weights) + for field, mapping in weights.items(): + val = getattr(signals, field) + score += mapping.get(val, 0.0) + return score / total_weight if total_weight > 0 else 0.0 + + +def document_quality_score(confidence: ConfidenceOutput) -> float: + """Weighted average of confidence dimensions.""" + return ( + confidence.overall * 0.4 + + confidence.event_type * 0.2 + + confidence.event_direction * 0.2 + + confidence.guidance * 0.1 + + confidence.risk_flags * 0.1 + ) + + +def compute_event_features(parser_output: dict[str, Any]) -> dict[str, Any]: + """Compute all event features from raw parser output dict.""" + from libs.schemas.types import ( + ConfidenceOutput, + GuidanceOutput, + RiskFlagsOutput, + SignalsOutput, + ) + + guidance = GuidanceOutput.model_validate(parser_output["guidance"]) + signals = SignalsOutput.model_validate(parser_output["signals"]) + risk_flags = RiskFlagsOutput.model_validate(parser_output["risk_flags"]) + confidence = ConfidenceOutput.model_validate(parser_output["confidence"]) + + return { + "guidance_direction_score": guidance_direction_score(guidance), + "guidance_status": guidance.status, + "oneoff_penalty": oneoff_penalty(risk_flags), + "signal_strength_score": signal_strength_score(signals), + "document_quality_score": document_quality_score(confidence), + "event_type": parser_output.get("event_type", "unknown"), + "event_direction": parser_output.get("event_direction", "unknown"), + "parse_confidence_overall": confidence.overall, + } diff --git a/libs/features/market_features.py b/libs/features/market_features.py new file mode 100644 index 0000000..75bba18 --- /dev/null +++ b/libs/features/market_features.py @@ -0,0 +1,110 @@ +"""Market-side feature calculations from price bar data.""" +from __future__ import annotations + +from typing import Any + +from libs.oracle_client.models import PriceBar + + +def reaction_day_return(bars: list[PriceBar], event_date: str) -> float | None: + """(close - prev_close) / prev_close on event date.""" + dated = {b.date: b for b in bars} + if event_date not in dated: + return None + event_bar = dated[event_date] + # Find previous bar + sorted_dates = sorted(dated.keys()) + idx = sorted_dates.index(event_date) + if idx == 0: + return None + prev_bar = dated[sorted_dates[idx - 1]] + if prev_bar.close == 0: + return None + return (event_bar.close - prev_bar.close) / prev_bar.close + + +def volume_ratio_20d(bars: list[PriceBar], event_date: str) -> float | None: + """Event day volume / 20-day average volume before event.""" + dated = {b.date: b for b in bars} + sorted_dates = sorted(dated.keys()) + if event_date not in dated: + return None + idx = sorted_dates.index(event_date) + if idx < 1: + return None + prior = sorted_dates[max(0, idx - 20) : idx] + if not prior: + return None + avg_vol = sum(dated[d].volume for d in prior) / len(prior) + if avg_vol == 0: + return None + return dated[event_date].volume / avg_vol + + +def close_location(bar: PriceBar) -> float | None: + """(close - low) / (high - low): 0=closed at low, 1=at high.""" + rng = bar.high - bar.low + if rng == 0: + return None + return (bar.close - bar.low) / rng + + +def gap_size(bars: list[PriceBar], event_date: str) -> float | None: + """(open_today - close_yesterday) / close_yesterday.""" + dated = {b.date: b for b in bars} + sorted_dates = sorted(dated.keys()) + if event_date not in dated: + return None + idx = sorted_dates.index(event_date) + if idx == 0: + return None + today = dated[event_date] + yesterday = dated[sorted_dates[idx - 1]] + if yesterday.close == 0: + return None + return (today.open - yesterday.close) / yesterday.close + + +def atr_14(bars: list[PriceBar]) -> float | None: + """14-period Average True Range.""" + if len(bars) < 2: + return None + sorted_bars = sorted(bars, key=lambda b: b.date) + true_ranges: list[float] = [] + for i in range(1, len(sorted_bars)): + curr = sorted_bars[i] + prev = sorted_bars[i - 1] + tr = max( + curr.high - curr.low, + abs(curr.high - prev.close), + abs(curr.low - prev.close), + ) + true_ranges.append(tr) + if len(true_ranges) < 14: + return sum(true_ranges) / len(true_ranges) if true_ranges else None + return sum(true_ranges[-14:]) / 14 + + +def compute_market_features( + bars: list[PriceBar], event_date: str +) -> dict[str, Any]: + """Compute all market features for an event date.""" + dated = {b.date: b for b in bars} + event_bar = dated.get(event_date) + + features: dict[str, Any] = { + "reaction_day_return": reaction_day_return(bars, event_date), + "volume_ratio_20d": volume_ratio_20d(bars, event_date), + "gap_size": gap_size(bars, event_date), + "atr_14": atr_14(bars), + } + if event_bar: + features["close_location"] = close_location(event_bar) + features["event_close"] = event_bar.close + features["event_volume"] = event_bar.volume + else: + features["close_location"] = None + features["event_close"] = None + features["event_volume"] = None + + return features diff --git a/libs/oracle_client/__init__.py b/libs/oracle_client/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/oracle_client/client.py b/libs/oracle_client/client.py new file mode 100644 index 0000000..03a8fab --- /dev/null +++ b/libs/oracle_client/client.py @@ -0,0 +1,108 @@ +"""Base httpx async client for Stock Oracle.""" +from __future__ import annotations + +from typing import Any + +import httpx + +from libs.oracle_client.exceptions import ( + OracleClientError, + OracleConnectionError, + OracleNotFoundError, + OracleServerError, + OracleTimeoutError, +) + + +class OracleClient: + """Async HTTP client wrapping Stock Oracle REST API.""" + + def __init__(self, base_url: str, timeout: float = 30.0) -> None: + self._base_url = base_url.rstrip("/") + self._timeout = timeout + self._client: httpx.AsyncClient | None = None + + async def __aenter__(self) -> OracleClient: + self._client = httpx.AsyncClient( + base_url=self._base_url, + timeout=self._timeout, + ) + return self + + async def __aexit__(self, *args: Any) -> None: + if self._client: + await self._client.aclose() + self._client = None + + def _ensure_client(self) -> httpx.AsyncClient: + if self._client is None: + raise RuntimeError("OracleClient must be used as async context manager.") + return self._client + + async def get(self, path: str, params: dict[str, Any] | None = None) -> Any: + client = self._ensure_client() + try: + response = await client.get(path, params=params) + except httpx.ConnectError as exc: + raise OracleConnectionError( + str(exc), source="oracle", entity=path + ) from exc + except httpx.TimeoutException as exc: + raise OracleTimeoutError( + str(exc), source="oracle", entity=path + ) from exc + + return self._handle_response(response, path) + + async def post(self, path: str, json: dict[str, Any] | None = None) -> Any: + client = self._ensure_client() + try: + response = await client.post(path, json=json) + except httpx.ConnectError as exc: + raise OracleConnectionError( + str(exc), source="oracle", entity=path + ) from exc + except httpx.TimeoutException as exc: + raise OracleTimeoutError( + str(exc), source="oracle", entity=path + ) from exc + + return self._handle_response(response, path) + + def _handle_response(self, response: httpx.Response, path: str) -> Any: + if response.status_code == 404: + raise OracleNotFoundError( + f"Not found: {path}", + source="oracle", + entity=path, + context={"status_code": 404}, + ) + if response.status_code >= 500: + raise OracleServerError( + f"Server error {response.status_code}: {path}", + source="oracle", + entity=path, + context={"status_code": response.status_code}, + ) + if response.status_code >= 400: + raise OracleClientError( + f"Client error {response.status_code}: {path}", + source="oracle", + entity=path, + context={"status_code": response.status_code}, + ) + return response.json() + + async def health_check(self) -> bool: + try: + data = await self.get("/health") + return isinstance(data, dict) and data.get("status") == "ok" + except Exception: + return False + + +def make_oracle_client() -> OracleClient: + from libs.common.config import get_settings + + s = get_settings() + return OracleClient(base_url=s.stock_oracle_url, timeout=float(s.stock_oracle_timeout)) diff --git a/libs/oracle_client/exceptions.py b/libs/oracle_client/exceptions.py new file mode 100644 index 0000000..0341fe4 --- /dev/null +++ b/libs/oracle_client/exceptions.py @@ -0,0 +1,24 @@ +"""Oracle client exception hierarchy.""" +from __future__ import annotations + +from libs.common.retries import NonRetryableError, RetryableError + + +class OracleConnectionError(RetryableError): + """Cannot connect to Stock Oracle (network down, refused).""" + + +class OracleTimeoutError(RetryableError): + """Request to Stock Oracle timed out.""" + + +class OracleServerError(RetryableError): + """Stock Oracle returned 5xx.""" + + +class OracleNotFoundError(NonRetryableError): + """Stock Oracle returned 404 (resource not found).""" + + +class OracleClientError(NonRetryableError): + """Stock Oracle returned 4xx (other than 404).""" diff --git a/libs/oracle_client/filings.py b/libs/oracle_client/filings.py new file mode 100644 index 0000000..beea109 --- /dev/null +++ b/libs/oracle_client/filings.py @@ -0,0 +1,42 @@ +"""Filing-related Oracle service methods.""" +from __future__ import annotations + +from libs.oracle_client.client import OracleClient +from libs.oracle_client.models import ( + ExhibitResponse, + FilingDocumentsResponse, + FilingSearchResponse, +) + + +class FilingsService: + def __init__(self, client: OracleClient) -> None: + self._client = client + + async def search_filings( + self, + ticker: str, + form_type: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + ) -> FilingSearchResponse: + params: dict[str, str] = {} + if form_type: + params["form_type"] = form_type + if start_date: + params["start_date"] = start_date + if end_date: + params["end_date"] = end_date + data = await self._client.get(f"/filings/search/{ticker}", params=params) + return FilingSearchResponse.model_validate(data) + + async def get_documents(self, accession_no: str) -> FilingDocumentsResponse: + data = await self._client.get(f"/filings/documents/{accession_no}") + return FilingDocumentsResponse.model_validate(data) + + async def get_exhibit( + self, accession_no: str, exhibit_type: str = "EX-99.1" + ) -> ExhibitResponse: + params = {"exhibit_type": exhibit_type} + data = await self._client.get(f"/filings/exhibit/{accession_no}", params=params) + return ExhibitResponse.model_validate(data) diff --git a/libs/oracle_client/financial.py b/libs/oracle_client/financial.py new file mode 100644 index 0000000..722a924 --- /dev/null +++ b/libs/oracle_client/financial.py @@ -0,0 +1,21 @@ +"""Financial data Oracle service methods.""" +from __future__ import annotations + +from libs.oracle_client.client import OracleClient +from libs.oracle_client.models import CompanyInfo, FinancialDataResponse + + +class FinancialService: + def __init__(self, client: OracleClient) -> None: + self._client = client + + async def get_financial_data( + self, ticker: str, quarters: int = 8 + ) -> FinancialDataResponse: + params: dict[str, str | int] = {"ticker": ticker, "quarters": quarters} + data = await self._client.get("/financial/data", params=params) + return FinancialDataResponse.model_validate(data) + + async def get_company_info(self, ticker: str) -> CompanyInfo: + data = await self._client.get(f"/financial/company/{ticker}") + return CompanyInfo.model_validate(data) diff --git a/libs/oracle_client/finra.py b/libs/oracle_client/finra.py new file mode 100644 index 0000000..324e693 --- /dev/null +++ b/libs/oracle_client/finra.py @@ -0,0 +1,24 @@ +"""FINRA short volume Oracle service methods.""" +from __future__ import annotations + +from libs.oracle_client.client import OracleClient +from libs.oracle_client.models import ShortRatioResponse, ShortVolumeResponse + + +class FinraService: + def __init__(self, client: OracleClient) -> None: + self._client = client + + async def get_short_volume( + self, symbol: str, days: int = 30 + ) -> ShortVolumeResponse: + params: dict[str, str | int] = {"days": days} + data = await self._client.get(f"/finra/short-volume/{symbol}", params=params) + return ShortVolumeResponse.model_validate(data) + + async def get_short_ratio( + self, symbol: str, days: int = 30 + ) -> ShortRatioResponse: + params: dict[str, str | int] = {"days": days} + data = await self._client.get(f"/finra/short-ratio/{symbol}", params=params) + return ShortRatioResponse.model_validate(data) diff --git a/libs/oracle_client/fred.py b/libs/oracle_client/fred.py new file mode 100644 index 0000000..c8b0d43 --- /dev/null +++ b/libs/oracle_client/fred.py @@ -0,0 +1,38 @@ +"""FRED data Oracle service methods.""" +from __future__ import annotations + +from libs.oracle_client.client import OracleClient +from libs.oracle_client.models import FredProxyResponse, FredSeriesInfo + + +class FredService: + def __init__(self, client: OracleClient) -> None: + self._client = client + + async def get_observations( + self, + series_id: str, + start: str | None = None, + end: str | None = None, + ) -> FredProxyResponse: + params: dict[str, str] = {"series_id": series_id} + if start: + params["observation_start"] = start + if end: + params["observation_end"] = end + data = await self._client.get( + "/fred/proxy/series/observations", params=params + ) + # Normalize to FredProxyResponse + if "observations" in data: + return FredProxyResponse(series_id=series_id, **data) + return FredProxyResponse(series_id=series_id, observations=data.get("data", [])) + + async def get_series_info(self, series_id: str) -> FredSeriesInfo: + params = {"series_id": series_id} + data = await self._client.get("/fred/proxy/series", params=params) + seriess = data.get("seriess", [data]) + if seriess: + info = seriess[0] + return FredSeriesInfo.model_validate({**info, "id": info.get("id", series_id)}) + return FredSeriesInfo(id=series_id) diff --git a/libs/oracle_client/models.py b/libs/oracle_client/models.py new file mode 100644 index 0000000..6a3af2b --- /dev/null +++ b/libs/oracle_client/models.py @@ -0,0 +1,181 @@ +"""Pydantic response models for Stock Oracle API.""" +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel, Field + +# --------------------------------------------------------------------------- +# Company Info +# --------------------------------------------------------------------------- + + +class CompanyInfo(BaseModel): + ticker: str + name: str | None = None + cik: str | None = None + exchange: str | None = None + sector: str | None = None + industry: str | None = None + country: str | None = None + market_cap: float | None = None + extra: dict[str, Any] = Field(default_factory=dict) + + +# --------------------------------------------------------------------------- +# SEC Filings +# --------------------------------------------------------------------------- + + +class FilingEntry(BaseModel): + accession_no: str + form_type: str + filing_date: str # ISO date string + accepted_at: str | None = None + primary_document: str | None = None + description: str | None = None + items: list[str] = Field(default_factory=list) + + +class FilingSearchResponse(BaseModel): + ticker: str + filings: list[FilingEntry] = Field(default_factory=list) + total: int = 0 + + +class ExhibitDocument(BaseModel): + exhibit_type: str + filename: str | None = None + url: str | None = None + + +class FilingDocumentsResponse(BaseModel): + accession_no: str + exhibits: list[ExhibitDocument] = Field(default_factory=list) + + +class ExhibitResponse(BaseModel): + accession_no: str + exhibit_type: str + content: str + content_type: str = "text/plain" + + +# --------------------------------------------------------------------------- +# Price / Market Data +# --------------------------------------------------------------------------- + + +class PriceBar(BaseModel): + date: str # ISO date + open: float + high: float + low: float + close: float + volume: int + vwap: float | None = None + trade_count: int | None = None + + +class PriceDataResponse(BaseModel): + ticker: str + bars: list[PriceBar] = Field(default_factory=list) + source: str = "yfinance" + + +class PriceQuote(BaseModel): + ticker: str + price: float + bid: float | None = None + ask: float | None = None + volume: int | None = None + timestamp: str | None = None + + +class IntradayBar(BaseModel): + timestamp: str + open: float + high: float + low: float + close: float + volume: int + + +class IntradayResponse(BaseModel): + ticker: str + bars: list[IntradayBar] = Field(default_factory=list) + timeframe: str = "1m" + + +# --------------------------------------------------------------------------- +# Financial / XBRL +# --------------------------------------------------------------------------- + + +class FinancialPeriod(BaseModel): + period: str # e.g. "2025-Q4" + period_end: str # ISO date + revenue: float | None = None + net_income: float | None = None + eps: float | None = None + gross_margin: float | None = None + operating_margin: float | None = None + extra: dict[str, Any] = Field(default_factory=dict) + + +class FinancialDataResponse(BaseModel): + ticker: str + periods: list[FinancialPeriod] = Field(default_factory=list) + + +# --------------------------------------------------------------------------- +# FRED +# --------------------------------------------------------------------------- + + +class FredObservation(BaseModel): + date: str # ISO date + value: float | None = None + + +class FredProxyResponse(BaseModel): + series_id: str + observations: list[FredObservation] = Field(default_factory=list) + realtime_start: str | None = None + realtime_end: str | None = None + + +class FredSeriesInfo(BaseModel): + id: str + title: str | None = None + frequency: str | None = None + units: str | None = None + notes: str | None = None + + +# --------------------------------------------------------------------------- +# FINRA Short Volume +# --------------------------------------------------------------------------- + + +class ShortVolumeEntry(BaseModel): + date: str # ISO date + short_volume: int + short_exempt_volume: int | None = None + total_volume: int | None = None + + +class ShortVolumeResponse(BaseModel): + symbol: str + data: list[ShortVolumeEntry] = Field(default_factory=list) + + +class ShortRatioPoint(BaseModel): + date: str + short_ratio: float | None = None + short_percent: float | None = None + + +class ShortRatioResponse(BaseModel): + symbol: str + data: list[ShortRatioPoint] = Field(default_factory=list) diff --git a/libs/oracle_client/price.py b/libs/oracle_client/price.py new file mode 100644 index 0000000..256dca1 --- /dev/null +++ b/libs/oracle_client/price.py @@ -0,0 +1,40 @@ +"""Price-related Oracle service methods.""" +from __future__ import annotations + +from libs.oracle_client.client import OracleClient +from libs.oracle_client.models import ( + IntradayResponse, + PriceDataResponse, + PriceQuote, +) + + +class PriceService: + def __init__(self, client: OracleClient) -> None: + self._client = client + + async def get_daily_bars( + self, + ticker: str, + start: str | None = None, + end: str | None = None, + ) -> PriceDataResponse: + params: dict[str, str] = {"ticker": ticker} + if start: + params["start"] = start + if end: + params["end"] = end + data = await self._client.get("/price/data", params=params) + return PriceDataResponse.model_validate(data) + + async def get_quote(self, ticker: str) -> PriceQuote: + data = await self._client.get(f"/price/quote/{ticker}") + return PriceQuote.model_validate(data) + + async def get_intraday(self, ticker: str) -> IntradayResponse: + data = await self._client.get(f"/price/intraday/{ticker}") + return IntradayResponse.model_validate(data) + + async def get_today(self, ticker: str) -> PriceDataResponse: + data = await self._client.get(f"/price/today/{ticker}") + return PriceDataResponse.model_validate(data) diff --git a/libs/parser/__init__.py b/libs/parser/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/parser/llm_parser_stub.py b/libs/parser/llm_parser_stub.py new file mode 100644 index 0000000..d0e8e1f --- /dev/null +++ b/libs/parser/llm_parser_stub.py @@ -0,0 +1,24 @@ +"""LLM parser stub — disabled by feature flag in Phase 1.""" +from __future__ import annotations + +from typing import Any + +from libs.schemas.types import ParserEventOutput + + +class LLMParserStub: + """Placeholder for future LLM-based parser. Not operational in Phase 1.""" + + def __init__(self, enabled: bool = False) -> None: + self.enabled = enabled + + def parse( + self, + document_id: str, + form_type: str, + text: str, + metadata: dict[str, Any] | None = None, + ) -> ParserEventOutput | None: + if not self.enabled: + return None + raise NotImplementedError("LLM parser is not implemented in Phase 1.") diff --git a/libs/parser/rule_parser.py b/libs/parser/rule_parser.py new file mode 100644 index 0000000..b018393 --- /dev/null +++ b/libs/parser/rule_parser.py @@ -0,0 +1,382 @@ +"""Rule-based event parser for SEC 8-K/6-K exhibits.""" +from __future__ import annotations + +import datetime as dt +import re +from typing import Any + +from libs.common.time_utils import filing_time_bucket +from libs.schemas.types import ( + ConfidenceOutput, + EvidenceItem, + GuidanceOutput, + ParserEventOutput, + RiskFlagsOutput, + SignalsOutput, +) + +PARSER_VERSION = "rule-1.0.0" +SCHEMA_VERSION = "1.0.0" + +# Item number patterns +_ITEM_RE = re.compile(r"item\s+(\d+\.\d+)", re.IGNORECASE) + +# Guidance +_GUIDANCE_RAISED = re.compile( + r"(guidance raised|above prior outlook|raised.*guidance|increased.*guidance" + r"|raised.*forecast|above.*expectations|above.*consensus|raised.*outlook" + r"|above.*prior.*guidance|exceed.*guidance)", + re.IGNORECASE, +) +_GUIDANCE_LOWERED = re.compile( + r"(guidance lowered|revised down|lowered.*guidance|reduced.*guidance" + r"|below.*expectations|below.*prior.*guidance|cut.*guidance|lowered.*outlook" + r"|lowered.*forecast)", + re.IGNORECASE, +) +_GUIDANCE_MAINTAINED = re.compile( + r"(reaffirm|maintain.*guidance|in.line.*outlook|inline.*guidance|on.track)", + re.IGNORECASE, +) +_GUIDANCE_WITHDRAWN = re.compile(r"(withdraw.*guidance|suspend.*guidance)", re.IGNORECASE) + +# Signals +_DEMAND_STRONG = re.compile( + r"(demand remains strong|strong demand|robust demand|record demand" + r"|demand acceleration|pipeline.*strong|strong.*pipeline)", + re.IGNORECASE, +) +_DEMAND_WEAK = re.compile( + r"(demand softness|weak demand|softer demand|demand weakness" + r"|elongated sales cycles|slower.*demand)", + re.IGNORECASE, +) +_PRICING_POWER = re.compile( + r"(pricing strength|price realization|pricing power|favorable.*pricing" + r"|pricing.*favorable|price.*increase|raised.*prices)", + re.IGNORECASE, +) +_PRICING_ABSENT = re.compile( + r"(pricing pressure|price.*compression|competitive.*pricing|price.*decline)", + re.IGNORECASE, +) +_BACKLOG_PRESENT = re.compile( + r"(backlog.*increas|increased.*backlog|bookings.*acceler|record.*backlog" + r"|strong.*backlog|ARR.*grew|deferred.*revenue.*increas|bookings.*grew)", + re.IGNORECASE, +) +_CUSTOMER_EXPANSION = re.compile( + r"(customer.*expand|new.*customer|customer.*addition|customer.*grow" + r"|expanded.*customer|added.*customer|net.*new.*customer)", + re.IGNORECASE, +) +_MARGIN_IMPROVING = re.compile( + r"(margin.*expan|expanding.*margin|margin.*improv|gross.*margin.*increas" + r"|operating.*margin.*improv|profitability.*improv)", + re.IGNORECASE, +) +_MARGIN_DETERIORATING = re.compile( + r"(margin.*compress|margin.*declin|margin.*contract|gross.*margin.*declin" + r"|operating.*margin.*declin)", + re.IGNORECASE, +) + +# Risk flags +_ONEOFF = re.compile( + r"(one.time|one.off|non.recurring|special.*charge|restructuring.*charge" + r"|impairment.*charge|write.down|write.off)", + re.IGNORECASE, +) +_TAX_BENEFIT = re.compile( + r"(tax.*benefit|deferred.*tax.*asset|tax.*credit|favorable.*tax)", re.IGNORECASE +) +_VALUATION_GAIN = re.compile( + r"(fair value.*gain|gain on.*sale|unrealized.*gain|valuation.*gain" + r"|mark.to.market.*gain)", + re.IGNORECASE, +) +_NON_GAAP = re.compile( + r"(non.GAAP|adjusted.*earnings|adjusted.*EPS|adjusted.*EBITDA" + r"|excluding.*items|excluding.*charges)", + re.IGNORECASE, +) +_FINANCING = re.compile( + r"(secondary.*offering|convertible.*note|equity.*offering|debt.*financing" + r"|new.*shares|dilut)", + re.IGNORECASE, +) +_LEGAL = re.compile( + r"(litigation|investigation|SEC.*inquiry|DOJ|legal.*proceeding" + r"|regulatory.*action|enforcement.*action)", + re.IGNORECASE, +) + + +def _extract_item_numbers(text: str) -> list[str]: + return list(dict.fromkeys(_ITEM_RE.findall(text))) + + +def _classify_event_type(items: list[str]) -> str: + if "2.02" in items: + return "earnings_release" + if "7.01" in items: + return "guidance_update" + if "1.01" in items: + return "material_contract" + if "8.01" in items: + return "other_material_event" + if "5.02" in items: + return "management_change" + if "1.03" in items: + return "other_material_event" + return "unknown" + + +def _detect_guidance(text: str) -> GuidanceOutput: + if _GUIDANCE_WITHDRAWN.search(text): + return GuidanceOutput(status="withdrawn", scope="unknown", notes="Guidance withdrawn.") + if _GUIDANCE_RAISED.search(text): + m = _GUIDANCE_RAISED.search(text) + return GuidanceOutput( + status="raised", + scope="unknown", + notes=m.group(0) if m else "", + ) + if _GUIDANCE_LOWERED.search(text): + m = _GUIDANCE_LOWERED.search(text) + return GuidanceOutput( + status="lowered", + scope="unknown", + notes=m.group(0) if m else "", + ) + if _GUIDANCE_MAINTAINED.search(text): + return GuidanceOutput( + status="inline_or_maintained", scope="unknown", notes="Guidance maintained." + ) + return GuidanceOutput(status="not_provided", scope="unknown", notes="") + + +def _detect_signals(text: str) -> SignalsOutput: + demand: str + if _DEMAND_STRONG.search(text): + demand = "strong" + elif _DEMAND_WEAK.search(text): + demand = "weakening" + else: + demand = "unknown" + + pricing: str + if _PRICING_POWER.search(text): + pricing = "present" + elif _PRICING_ABSENT.search(text): + pricing = "absent" + else: + pricing = "unknown" + + backlog = "present" if _BACKLOG_PRESENT.search(text) else "unknown" + customer = "present" if _CUSTOMER_EXPANSION.search(text) else "unknown" + + margin: str + if _MARGIN_IMPROVING.search(text): + margin = "improving" + elif _MARGIN_DETERIORATING.search(text): + margin = "deteriorating" + else: + margin = "unknown" + + return SignalsOutput( + demand_strength=demand, + pricing_power=pricing, + backlog_or_bookings=backlog, + customer_expansion=customer, + margin_quality=margin, + ) + + +def _detect_risk_flags(text: str) -> RiskFlagsOutput: + return RiskFlagsOutput( + oneoff_item=bool(_ONEOFF.search(text)), + tax_benefit=bool(_TAX_BENEFIT.search(text)), + valuation_gain=bool(_VALUATION_GAIN.search(text)), + non_gaap_heavy=bool(_NON_GAAP.search(text)), + financing_related=bool(_FINANCING.search(text)), + legal_or_regulatory_overhang=bool(_LEGAL.search(text)), + ) + + +def _classify_direction( + guidance: GuidanceOutput, + signals: SignalsOutput, + risk_flags: RiskFlagsOutput, +) -> str: + bullish_signals = 0 + bearish_signals = 0 + + if guidance.status == "raised": + bullish_signals += 2 + elif guidance.status == "lowered" or guidance.status == "withdrawn": + bearish_signals += 2 + + if signals.demand_strength == "strong": + bullish_signals += 1 + elif signals.demand_strength == "weakening": + bearish_signals += 1 + + if signals.pricing_power == "present": + bullish_signals += 1 + elif signals.pricing_power == "absent": + bearish_signals += 1 + + if signals.margin_quality == "improving": + bullish_signals += 1 + elif signals.margin_quality == "deteriorating": + bearish_signals += 1 + + if risk_flags.financing_related: + bearish_signals += 1 + + if bullish_signals > bearish_signals + 1: + return "bullish" + if bearish_signals > bullish_signals + 1: + return "bearish" + if bullish_signals > 0 or bearish_signals > 0: + return "mixed" + return "unknown" + + +def _compute_confidence( + items: list[str], + guidance: GuidanceOutput, + signals: SignalsOutput, + risk_flags: RiskFlagsOutput, +) -> ConfidenceOutput: + event_type_conf = 0.9 if items else 0.4 + guidance_conf = 0.0 if guidance.status in ("not_provided", "unclear") else 0.8 + direction_conf = 0.5 + + known_signals = sum( + 1 + for v in [ + signals.demand_strength, + signals.pricing_power, + signals.backlog_or_bookings, + signals.customer_expansion, + signals.margin_quality, + ] + if v != "unknown" + ) + direction_conf = min(0.9, 0.3 + known_signals * 0.12) + + risk_count = sum( + [ + risk_flags.oneoff_item, + risk_flags.tax_benefit, + risk_flags.valuation_gain, + risk_flags.non_gaap_heavy, + risk_flags.financing_related, + risk_flags.legal_or_regulatory_overhang, + ] + ) + risk_conf = max(0.2, 1.0 - risk_count * 0.1) + + overall = (event_type_conf * 0.3 + direction_conf * 0.4 + guidance_conf * 0.2 + risk_conf * 0.1) + + return ConfidenceOutput( + overall=round(overall, 3), + event_type=round(event_type_conf, 3), + event_direction=round(direction_conf, 3), + guidance=round(guidance_conf, 3), + risk_flags=round(risk_conf, 3), + ) + + +def _build_evidence(text: str, guidance: GuidanceOutput, signals: SignalsOutput) -> list[EvidenceItem]: + evidence: list[EvidenceItem] = [] + if guidance.notes: + evidence.append( + EvidenceItem( + label="guidance_signal", + text_span=guidance.notes[:200], + section_hint="guidance", + confidence=0.8, + ) + ) + for pattern, label in [ + (_DEMAND_STRONG, "demand_strong"), + (_BACKLOG_PRESENT, "backlog_present"), + (_CUSTOMER_EXPANSION, "customer_expansion"), + (_MARGIN_IMPROVING, "margin_improving"), + ]: + m = pattern.search(text) + if m: + evidence.append( + EvidenceItem( + label=label, + text_span=m.group(0)[:200], + section_hint="body", + confidence=0.7, + ) + ) + return evidence[:10] + + +class RuleBasedParser: + """Deterministic rule-based parser for SEC 8-K/6-K documents.""" + + def parse( + self, + document_id: str, + form_type: str, + text: str, + metadata: dict[str, Any] | None = None, + ) -> ParserEventOutput: + metadata = metadata or {} + filing_date_str = metadata.get("filing_date", dt.date.today().isoformat()) + accepted_at_str = metadata.get("accepted_at_utc") + + # Determine filing time bucket + time_bucket = "unknown" + if accepted_at_str: + try: + import datetime as dt2 + accepted_dt = dt2.datetime.fromisoformat(accepted_at_str.replace("Z", "+00:00")) + time_bucket = filing_time_bucket(accepted_dt) + except Exception: + time_bucket = "unknown" + + items = metadata.get("item_numbers") or _extract_item_numbers(text) + event_type = _classify_event_type(items) + guidance = _detect_guidance(text) + signals = _detect_signals(text) + risk_flags = _detect_risk_flags(text) + direction = _classify_direction(guidance, signals, risk_flags) + confidence = _compute_confidence(items, guidance, signals, risk_flags) + evidence = _build_evidence(text, guidance, signals) + + # Build summary from first 500 chars + first_para = text[:500].strip().replace("\n", " ") + summary = first_para if first_para else f"{form_type} document, event: {event_type}" + + warnings: list[str] = [] + if event_type == "unknown": + warnings.append("Could not determine event type from item numbers or content.") + if direction == "unknown": + warnings.append("Could not determine event direction from content signals.") + + return ParserEventOutput( + schema_version=SCHEMA_VERSION, + document_id=document_id, + parser_kind="rule", + event_type=event_type, + event_direction=direction, + event_date=filing_date_str, + filing_time_bucket=time_bucket, + headline=summary[:120], + summary=summary, + guidance=guidance, + signals=signals, + risk_flags=risk_flags, + evidence=evidence, + confidence=confidence, + warnings=warnings, + ) diff --git a/libs/parser/schema_validator.py b/libs/parser/schema_validator.py new file mode 100644 index 0000000..f984f9d --- /dev/null +++ b/libs/parser/schema_validator.py @@ -0,0 +1,30 @@ +"""JSON Schema validator for parser event output.""" +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from jsonschema import Draft202012Validator + +_SCHEMA_PATH = Path(__file__).parent.parent / "schemas" / "parser_event.schema.json" +_SCHEMA: dict[str, Any] | None = None + + +def _get_schema() -> dict[str, Any]: + global _SCHEMA + if _SCHEMA is None: + _SCHEMA = json.loads(_SCHEMA_PATH.read_text()) + return _SCHEMA + + +def validate_parser_output(data: dict[str, Any]) -> list[str]: + """Validate parser output against schema. Returns list of error messages (empty = valid).""" + schema = _get_schema() + validator = Draft202012Validator(schema) + errors = sorted(validator.iter_errors(data), key=lambda e: list(e.path)) + return [f"{'.'.join(str(p) for p in e.path) or 'root'}: {e.message}" for e in errors] + + +def is_valid(data: dict[str, Any]) -> bool: + return len(validate_parser_output(data)) == 0 diff --git a/libs/parser/text_normalizer.py b/libs/parser/text_normalizer.py new file mode 100644 index 0000000..f0b3e33 --- /dev/null +++ b/libs/parser/text_normalizer.py @@ -0,0 +1,57 @@ +"""HTML-to-text conversion and text normalization.""" +from __future__ import annotations + +import re +import unicodedata + +from bs4 import BeautifulSoup + +_BOILERPLATE_PATTERNS = [ + re.compile(r"safe harbor.*?forward.looking statement", re.IGNORECASE | re.DOTALL), + re.compile(r"this press release.*?private securities litigation", re.IGNORECASE | re.DOTALL), + re.compile(r"^\s*page \d+ of \d+\s*$", re.IGNORECASE | re.MULTILINE), + re.compile(r"^\s*\[?\s*table of contents\s*\]?\s*$", re.IGNORECASE | re.MULTILINE), +] + +_WHITESPACE = re.compile(r"\s{3,}") + + +def html_to_text(html: str) -> str: + """Convert HTML to plain text using BeautifulSoup.""" + soup = BeautifulSoup(html, "html.parser") + # Remove script/style + for tag in soup(["script", "style", "head"]): + tag.decompose() + return soup.get_text(separator="\n") + + +def normalize_unicode(text: str) -> str: + """Normalize unicode to NFC and replace fancy quotes/dashes.""" + text = unicodedata.normalize("NFC", text) + # Fancy quotes → standard + text = text.replace("\u2018", "'").replace("\u2019", "'") + text = text.replace("\u201c", '"').replace("\u201d", '"') + # Em/en dash → hyphen + text = text.replace("\u2014", " - ").replace("\u2013", " - ") + return text + + +def remove_boilerplate(text: str) -> str: + """Strip common boilerplate sections.""" + for pattern in _BOILERPLATE_PATTERNS: + text = pattern.sub(" ", text) + return text + + +def collapse_whitespace(text: str) -> str: + """Collapse runs of 3+ whitespace chars to double newline.""" + return _WHITESPACE.sub("\n\n", text).strip() + + +def normalize_text(raw: str, is_html: bool = False) -> str: + """Full normalization pipeline.""" + text = html_to_text(raw) if is_html else raw + text = normalize_unicode(text) + text = remove_boilerplate(text) + text = collapse_whitespace(text) + return text diff --git a/libs/schemas/__init__.py b/libs/schemas/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/libs/schemas/parser_event.schema.json b/libs/schemas/parser_event.schema.json new file mode 100644 index 0000000..f69b8c9 --- /dev/null +++ b/libs/schemas/parser_event.schema.json @@ -0,0 +1,145 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "$id": "https://example.local/schemas/parser_event.schema.json", + "title": "ParserEvent", + "type": "object", + "additionalProperties": false, + "required": [ + "schema_version", + "document_id", + "parser_kind", + "event_type", + "event_direction", + "event_date", + "filing_time_bucket", + "summary", + "guidance", + "signals", + "risk_flags", + "confidence" + ], + "properties": { + "schema_version": { "type": "string" }, + "document_id": { "type": "string", "minLength": 1 }, + "parser_kind": { + "type": "string", + "enum": ["rule", "llm", "merged"] + }, + "event_type": { + "type": "string", + "enum": [ + "earnings_release", + "guidance_update", + "material_contract", + "regulatory_or_approval", + "capital_markets_or_financing", + "management_change", + "litigation_or_investigation", + "other_material_event", + "unknown" + ] + }, + "event_direction": { + "type": "string", + "enum": ["bullish", "bearish", "mixed", "neutral", "unknown"] + }, + "event_date": { "type": "string", "format": "date" }, + "filing_time_bucket": { + "type": "string", + "enum": ["pre_market", "regular_hours", "post_market", "unknown"] + }, + "headline": { "type": "string" }, + "summary": { "type": "string", "minLength": 1 }, + "guidance": { + "type": "object", + "additionalProperties": false, + "required": ["status", "scope", "notes"], + "properties": { + "status": { + "type": "string", + "enum": [ + "raised", + "inline_or_maintained", + "lowered", + "withdrawn", + "not_provided", + "unclear" + ] + }, + "scope": { + "type": "string", + "enum": ["quarterly", "annual", "both", "unknown"] + }, + "notes": { "type": "string" } + } + }, + "signals": { + "type": "object", + "additionalProperties": false, + "required": [ + "demand_strength", + "pricing_power", + "backlog_or_bookings", + "customer_expansion", + "margin_quality" + ], + "properties": { + "demand_strength": { "type": "string", "enum": ["strong", "stable", "weakening", "unknown"] }, + "pricing_power": { "type": "string", "enum": ["present", "mixed", "absent", "unknown"] }, + "backlog_or_bookings": { "type": "string", "enum": ["present", "mixed", "absent", "unknown"] }, + "customer_expansion": { "type": "string", "enum": ["present", "mixed", "absent", "unknown"] }, + "margin_quality": { "type": "string", "enum": ["improving", "stable", "deteriorating", "unknown"] } + } + }, + "risk_flags": { + "type": "object", + "additionalProperties": false, + "required": [ + "oneoff_item", + "tax_benefit", + "valuation_gain", + "non_gaap_heavy", + "financing_related", + "legal_or_regulatory_overhang" + ], + "properties": { + "oneoff_item": { "type": "boolean" }, + "tax_benefit": { "type": "boolean" }, + "valuation_gain": { "type": "boolean" }, + "non_gaap_heavy": { "type": "boolean" }, + "financing_related": { "type": "boolean" }, + "legal_or_regulatory_overhang": { "type": "boolean" } + } + }, + "evidence": { + "type": "array", + "items": { + "type": "object", + "additionalProperties": false, + "required": ["label", "text_span", "section_hint", "confidence"], + "properties": { + "label": { "type": "string" }, + "text_span": { "type": "string" }, + "section_hint": { "type": "string" }, + "confidence": { "type": "number", "minimum": 0, "maximum": 1 } + } + } + }, + "confidence": { + "type": "object", + "additionalProperties": false, + "required": ["overall", "event_type", "event_direction", "guidance", "risk_flags"], + "properties": { + "overall": { "type": "number", "minimum": 0, "maximum": 1 }, + "event_type": { "type": "number", "minimum": 0, "maximum": 1 }, + "event_direction": { "type": "number", "minimum": 0, "maximum": 1 }, + "guidance": { "type": "number", "minimum": 0, "maximum": 1 }, + "risk_flags": { "type": "number", "minimum": 0, "maximum": 1 } + } + }, + "warnings": { + "type": "array", + "items": { "type": "string" } + } + } +} diff --git a/libs/schemas/types.py b/libs/schemas/types.py new file mode 100644 index 0000000..efc1f2b --- /dev/null +++ b/libs/schemas/types.py @@ -0,0 +1,74 @@ +"""Pydantic types for parser event output.""" +from __future__ import annotations + +from typing import Literal + +from pydantic import BaseModel, Field + + +class GuidanceOutput(BaseModel): + status: Literal[ + "raised", "inline_or_maintained", "lowered", "withdrawn", "not_provided", "unclear" + ] + scope: Literal["quarterly", "annual", "both", "unknown"] + notes: str = "" + + +class SignalsOutput(BaseModel): + demand_strength: Literal["strong", "stable", "weakening", "unknown"] + pricing_power: Literal["present", "mixed", "absent", "unknown"] + backlog_or_bookings: Literal["present", "mixed", "absent", "unknown"] + customer_expansion: Literal["present", "mixed", "absent", "unknown"] + margin_quality: Literal["improving", "stable", "deteriorating", "unknown"] + + +class RiskFlagsOutput(BaseModel): + oneoff_item: bool = False + tax_benefit: bool = False + valuation_gain: bool = False + non_gaap_heavy: bool = False + financing_related: bool = False + legal_or_regulatory_overhang: bool = False + + +class ConfidenceOutput(BaseModel): + overall: float = Field(ge=0.0, le=1.0) + event_type: float = Field(ge=0.0, le=1.0) + event_direction: float = Field(ge=0.0, le=1.0) + guidance: float = Field(ge=0.0, le=1.0) + risk_flags: float = Field(ge=0.0, le=1.0) + + +class EvidenceItem(BaseModel): + label: str + text_span: str + section_hint: str + confidence: float = Field(ge=0.0, le=1.0) + + +class ParserEventOutput(BaseModel): + schema_version: str = "1.0.0" + document_id: str + parser_kind: Literal["rule", "llm", "merged"] + event_type: Literal[ + "earnings_release", + "guidance_update", + "material_contract", + "regulatory_or_approval", + "capital_markets_or_financing", + "management_change", + "litigation_or_investigation", + "other_material_event", + "unknown", + ] + event_direction: Literal["bullish", "bearish", "mixed", "neutral", "unknown"] + event_date: str # ISO date + filing_time_bucket: Literal["pre_market", "regular_hours", "post_market", "unknown"] + headline: str = "" + summary: str + guidance: GuidanceOutput + signals: SignalsOutput + risk_flags: RiskFlagsOutput + evidence: list[EvidenceItem] = Field(default_factory=list) + confidence: ConfidenceOutput + warnings: list[str] = Field(default_factory=list) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..2a70838 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,64 @@ +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[project] +name = "acef" +version = "0.1.0" +description = "ACE-F v1 -- US Stock Event Swing Trading System" +requires-python = ">=3.11" +dependencies = [ + "httpx>=0.27", + "pydantic>=2.6", + "pydantic-settings>=2.2", + "sqlalchemy[asyncio]>=2.0", + "asyncpg>=0.29", + "alembic>=1.13", + "structlog>=24.1", + "python-dotenv>=1.0", + "pyyaml>=6.0", + "exchange-calendars>=4.5", + "jsonschema>=4.21", + "beautifulsoup4>=4.12", + "tenacity>=8.2", + "duckdb>=1.0", + "pyarrow>=15.0", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "pytest-asyncio>=0.23", + "pytest-httpx>=0.30", + "pytest-cov>=5.0", + "ruff>=0.3", + "mypy>=1.9", + "testcontainers[postgres]>=4.0", + "types-pyyaml", + "types-beautifulsoup4", +] + +[tool.hatch.build.targets.wheel] +packages = ["apps", "libs"] + +[tool.ruff] +line-length = 100 +target-version = "py311" + +[tool.ruff.lint] +select = ["E", "F", "I", "UP", "B", "C4", "SIM"] +ignore = ["E501"] + +[tool.mypy] +python_version = "3.11" +strict = true +ignore_missing_imports = true + +[tool.pytest.ini_options] +asyncio_mode = "auto" +testpaths = ["tests"] +markers = [ + "unit: unit tests (no external dependencies)", + "integration: integration tests (requires postgres)", + "replay: determinism and idempotency tests", +] diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..8faa66c --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,93 @@ +"""Shared test fixtures for ACE-F test suite.""" +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +FIXTURES_DIR = Path(__file__).parent / "fixtures" + + +def load_fixture(name: str) -> dict: + return json.loads((FIXTURES_DIR / name).read_text()) + + +@pytest.fixture +def filing_search_fixture() -> dict: + return load_fixture("filing_search.json") + + +@pytest.fixture +def exhibit_content_fixture() -> dict: + return load_fixture("exhibit_content.json") + + +@pytest.fixture +def price_data_fixture() -> dict: + return load_fixture("price_data.json") + + +@pytest.fixture +def financial_data_fixture() -> dict: + return load_fixture("financial_data.json") + + +@pytest.fixture +def short_volume_fixture() -> dict: + return load_fixture("short_volume.json") + + +@pytest.fixture +def fred_observations_fixture() -> dict: + return load_fixture("fred_observations.json") + + +@pytest.fixture +def sample_exhibit_text() -> str: + data = load_fixture("exhibit_content.json") + return data["content"] + + +@pytest.fixture +def sample_parser_output() -> dict: + return { + "schema_version": "1.0.0", + "document_id": "DOC::sec::ISSUER::0000320193::2026-01-29::0000320193-26-000001", + "parser_kind": "rule", + "event_type": "earnings_release", + "event_direction": "bullish", + "event_date": "2026-01-29", + "filing_time_bucket": "post_market", + "headline": "Apple Inc. Reports First Quarter Results", + "summary": "Apple today announced financial results for its fiscal 2026 Q1.", + "guidance": { + "status": "raised", + "scope": "unknown", + "notes": "Guidance raised for Q2" + }, + "signals": { + "demand_strength": "strong", + "pricing_power": "present", + "backlog_or_bookings": "present", + "customer_expansion": "present", + "margin_quality": "improving" + }, + "risk_flags": { + "oneoff_item": False, + "tax_benefit": False, + "valuation_gain": False, + "non_gaap_heavy": True, + "financing_related": False, + "legal_or_regulatory_overhang": False + }, + "evidence": [], + "confidence": { + "overall": 0.78, + "event_type": 0.9, + "event_direction": 0.75, + "guidance": 0.8, + "risk_flags": 0.9 + }, + "warnings": [] + } diff --git a/tests/fixtures/__init__.py b/tests/fixtures/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/fixtures/exhibit_content.json b/tests/fixtures/exhibit_content.json new file mode 100644 index 0000000..e5b4ebb --- /dev/null +++ b/tests/fixtures/exhibit_content.json @@ -0,0 +1,6 @@ +{ + "accession_no": "0000320193-26-000001", + "exhibit_type": "EX-99.1", + "content": "Apple Inc. Reports First Quarter Results\n\nItem 2.02 Results of Operations\n\nApple today announced financial results for its fiscal 2026 first quarter ended December 28, 2025. The Company posted quarterly revenue of $124.3 billion, up 9 percent year over year.\n\nGuidance raised for Q2: The Company expects revenue to be between $125 billion and $131 billion. Demand remains strong across all product categories. Gross margin is expected to be between 46.5 percent and 47.5 percent, reflecting margin expansion driven by Services mix.\n\nCustomer expansion in enterprise continues. Backlog increased significantly. Pricing strength maintained.\n\nAdjusted EPS excludes certain non-GAAP items.", + "content_type": "text/plain" +} diff --git a/tests/fixtures/filing_search.json b/tests/fixtures/filing_search.json new file mode 100644 index 0000000..1aa3431 --- /dev/null +++ b/tests/fixtures/filing_search.json @@ -0,0 +1,24 @@ +{ + "ticker": "AAPL", + "total": 2, + "filings": [ + { + "accession_no": "0000320193-26-000001", + "form_type": "8-K", + "filing_date": "2026-01-29", + "accepted_at": "2026-01-29T21:05:00Z", + "primary_document": "d123456d8k.htm", + "description": "Results of Operations and Financial Condition", + "items": ["2.02", "9.01"] + }, + { + "accession_no": "0000320193-26-000002", + "form_type": "8-K", + "filing_date": "2026-02-15", + "accepted_at": "2026-02-15T16:30:00Z", + "primary_document": "d234567d8k.htm", + "description": "Other Events", + "items": ["8.01"] + } + ] +} diff --git a/tests/fixtures/financial_data.json b/tests/fixtures/financial_data.json new file mode 100644 index 0000000..2d39ce7 --- /dev/null +++ b/tests/fixtures/financial_data.json @@ -0,0 +1,23 @@ +{ + "ticker": "AAPL", + "periods": [ + { + "period": "2026-Q1", + "period_end": "2025-12-28", + "revenue": 124300000000, + "net_income": 36000000000, + "eps": 2.34, + "gross_margin": 0.472, + "operating_margin": 0.315 + }, + { + "period": "2025-Q4", + "period_end": "2025-09-27", + "revenue": 119600000000, + "net_income": 34900000000, + "eps": 2.26, + "gross_margin": 0.461, + "operating_margin": 0.308 + } + ] +} diff --git a/tests/fixtures/fred_observations.json b/tests/fixtures/fred_observations.json new file mode 100644 index 0000000..526796b --- /dev/null +++ b/tests/fixtures/fred_observations.json @@ -0,0 +1,11 @@ +{ + "series_id": "DGS10", + "realtime_start": "2026-01-01", + "realtime_end": "2026-03-12", + "observations": [ + {"date": "2026-03-10", "value": 4.32}, + {"date": "2026-03-09", "value": 4.28}, + {"date": "2026-03-06", "value": 4.35}, + {"date": "2026-03-05", "value": 4.41} + ] +} diff --git a/tests/fixtures/price_data.json b/tests/fixtures/price_data.json new file mode 100644 index 0000000..d141876 --- /dev/null +++ b/tests/fixtures/price_data.json @@ -0,0 +1,12 @@ +{ + "ticker": "AAPL", + "source": "yfinance", + "bars": [ + {"date": "2026-01-23", "open": 225.0, "high": 228.5, "low": 224.0, "close": 227.0, "volume": 52000000}, + {"date": "2026-01-26", "open": 227.5, "high": 230.0, "low": 226.0, "close": 229.5, "volume": 48000000}, + {"date": "2026-01-27", "open": 229.0, "high": 232.0, "low": 228.0, "close": 231.0, "volume": 55000000}, + {"date": "2026-01-28", "open": 231.5, "high": 234.0, "low": 230.0, "close": 233.0, "volume": 60000000}, + {"date": "2026-01-29", "open": 240.0, "high": 245.0, "low": 238.0, "close": 243.0, "volume": 120000000}, + {"date": "2026-01-30", "open": 243.5, "high": 246.0, "low": 241.0, "close": 244.5, "volume": 75000000} + ] +} diff --git a/tests/fixtures/short_volume.json b/tests/fixtures/short_volume.json new file mode 100644 index 0000000..6468a03 --- /dev/null +++ b/tests/fixtures/short_volume.json @@ -0,0 +1,8 @@ +{ + "symbol": "AAPL", + "data": [ + {"date": "2026-03-10", "short_volume": 5000000, "short_exempt_volume": 50000, "total_volume": 45000000}, + {"date": "2026-03-09", "short_volume": 4800000, "short_exempt_volume": 48000, "total_volume": 42000000}, + {"date": "2026-03-06", "short_volume": 5200000, "short_exempt_volume": 52000, "total_volume": 48000000} + ] +} diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py new file mode 100644 index 0000000..50da26e --- /dev/null +++ b/tests/integration/conftest.py @@ -0,0 +1,61 @@ +"""Integration test fixtures using testcontainers PostgreSQL.""" +from __future__ import annotations + +import asyncio + +import pytest +import pytest_asyncio +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + + +@pytest.fixture(scope="session") +def event_loop_policy(): + return asyncio.DefaultEventLoopPolicy() + + +@pytest.fixture(scope="session") +def postgres_container(): + """Start a PostgreSQL container for the test session.""" + try: + from testcontainers.postgres import PostgresContainer + + with PostgresContainer("postgres:16-alpine") as pg: + yield pg + except ImportError: + pytest.skip("testcontainers not installed") + + +@pytest.fixture(scope="session") +def async_db_url(postgres_container) -> str: + """Return asyncpg DSN from the postgres container.""" + sync_url = postgres_container.get_connection_url() + # Convert psycopg2 URL to asyncpg + return sync_url.replace("postgresql+psycopg2://", "postgresql+asyncpg://").replace( + "postgresql://", "postgresql+asyncpg://" + ) + + +@pytest_asyncio.fixture(scope="session") +async def db_engine(async_db_url): + """Create engine and run migrations.""" + from alembic import command + from alembic.config import Config + + # Run migrations via alembic + alembic_cfg = Config("alembic.ini") + alembic_cfg.set_main_option("sqlalchemy.url", async_db_url.replace("+asyncpg", "+psycopg2")) + # Use synchronous URL for alembic + command.upgrade(alembic_cfg, "head") + + engine = create_async_engine(async_db_url, echo=False) + yield engine + await engine.dispose() + + +@pytest_asyncio.fixture +async def db_session(db_engine) -> AsyncSession: + """Provide a test DB session that rolls back after each test.""" + factory = async_sessionmaker(db_engine, expire_on_commit=False, class_=AsyncSession) + async with factory() as session: + yield session + await session.rollback() diff --git a/tests/integration/test_db_migration.py b/tests/integration/test_db_migration.py new file mode 100644 index 0000000..b5a9755 --- /dev/null +++ b/tests/integration/test_db_migration.py @@ -0,0 +1,42 @@ +"""Integration test: verify all tables exist after migration.""" +import pytest +from sqlalchemy import text + +EXPECTED_TABLES = [ + "issuer_master", + "symbol_master", + "job_runs", + "documents", + "document_exhibits", + "exhibit_cache", + "events", + "event_parses", + "macro_series", + "macro_observations", + "short_sale_daily", + "feature_snapshots", + "order_plans", + "sync_checkpoints", +] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_all_tables_exist(db_session): + """All expected tables should exist after migration.""" + result = await db_session.execute( + text( + "SELECT table_name FROM information_schema.tables " + "WHERE table_schema = 'public'" + ) + ) + existing = {row[0] for row in result} + for table in EXPECTED_TABLES: + assert table in existing, f"Table {table!r} missing after migration" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_db_health(db_session): + from libs.db.helpers import health_check + assert await health_check(db_session) is True diff --git a/tests/integration/test_feature_pipeline.py b/tests/integration/test_feature_pipeline.py new file mode 100644 index 0000000..b55d155 --- /dev/null +++ b/tests/integration/test_feature_pipeline.py @@ -0,0 +1,83 @@ +"""Integration test: feature pipeline (event + mock price → feature snapshot).""" +import datetime as dt + +import pytest +from pytest_httpx import HTTPXMock + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_feature_snapshot_created(db_session, httpx_mock: HTTPXMock, price_data_fixture, sample_parser_output): + """Given a valid event + parse, feature snapshots should be created.""" + + from libs.db.models import ( + Document, + Event, + EventParse, + IssuerMaster, + SymbolMaster, + ) + from libs.features.builder import build_features_for_event + from libs.oracle_client.client import OracleClient + from libs.oracle_client.price import PriceService + + # Mock price endpoint + httpx_mock.add_response(json=price_data_fixture) + + # Setup DB records + issuer = IssuerMaster(issuer_id="ISSUER::0000320193", issuer_name="Apple Inc.", ticker="AAPL") + db_session.add(issuer) + + symbol = SymbolMaster( + symbol_id="SYM::AAPL::XNYS", + issuer_id="ISSUER::0000320193", + ticker="AAPL", + venue="XNYS", + ) + db_session.add(symbol) + + doc = Document( + document_id="DOC::sec::ISSUER::0000320193::2026-01-29::ACC001", + source_name="sec", + form_type="8-K", + filing_date=dt.date(2026, 1, 29), + accession_no="ACC001", + parsed_status="succeeded", + ) + db_session.add(doc) + await db_session.flush() + + event = Event( + event_id="EVT::test::earnings_release::0", + primary_document_id=doc.document_id, + symbol_id="SYM::AAPL::XNYS", + event_type="earnings_release", + event_direction="bullish", + event_date=dt.date(2026, 1, 29), + parser_version="rule-1.0.0", + status="pending", + ) + db_session.add(event) + await db_session.flush() + + parse = EventParse( + event_id=event.event_id, + parser_kind="rule", + parser_version="rule-1.0.0", + schema_version="1.0.0", + output_json=sample_parser_output, + validation_status="valid", + ) + db_session.add(parse) + await db_session.flush() + + async with OracleClient("http://oracle:18001") as client: + price_svc = PriceService(client) + result = await build_features_for_event(db_session, event, price_svc) + + assert result is not None + market_snap, event_snap = result + assert market_snap.snapshot_name == "market_v1" + assert event_snap.snapshot_name == "event_v1" + assert "reaction_day_return" in market_snap.feature_json + assert "guidance_direction_score" in event_snap.feature_json diff --git a/tests/integration/test_filing_pipeline.py b/tests/integration/test_filing_pipeline.py new file mode 100644 index 0000000..e516e7c --- /dev/null +++ b/tests/integration/test_filing_pipeline.py @@ -0,0 +1,128 @@ +"""Integration test: filing pipeline (poll → fetch → parse → event in DB).""" +import datetime as dt +from pathlib import Path + +import pytest + +FIXTURES = Path(__file__).parent.parent / "fixtures" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_document_upsert_idempotency_via_helpers(db_session): + """Upsert helper inserts same document twice → only one row.""" + import datetime as dt + + from sqlalchemy import func, select + + from libs.db.helpers import upsert + from libs.db.models import Document + + values = { + "document_id": "DOC::idempotency::test::2026-01-01::IDEMACC001", + "source_name": "sec", + "form_type": "8-K", + "filing_date": dt.date(2026, 1, 1), + "accession_no": "IDEMACC001", + "parsed_status": "pending", + "created_at_utc": dt.datetime.now(tz=dt.UTC), + "updated_at_utc": dt.datetime.now(tz=dt.UTC), + } + + await upsert(db_session, Document, values, index_elements=["document_id"]) + await upsert(db_session, Document, values, index_elements=["document_id"]) + await db_session.flush() + + result = await db_session.execute( + select(func.count()).where( + Document.document_id == "DOC::idempotency::test::2026-01-01::IDEMACC001" + ) + ) + count = result.scalar_one() + assert count == 1 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_document_upsert_idempotency(db_session): + """Inserting same document twice should not create duplicates.""" + from sqlalchemy import select + + from libs.db.models import Document + + doc_id = "DOC::test::ISSUER::X::2026-01-01::ACC001" + doc1 = Document( + document_id=doc_id, + source_name="sec", + form_type="8-K", + filing_date=dt.date(2026, 1, 29), + accession_no="ACC001", + parsed_status="pending", + ) + db_session.add(doc1) + await db_session.flush() + + # Try to insert again (should conflict on unique accession+form_type) + result = await db_session.execute( + select(Document).where(Document.document_id == doc_id) + ) + rows = result.scalars().all() + assert len(rows) == 1 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_event_parse_lifecycle(db_session): + """Create issuer → document → event → event_parse chain.""" + from libs.db.models import Document, Event, EventParse, IssuerMaster + + # Create issuer + issuer = IssuerMaster( + issuer_id="ISSUER::0000320193", + issuer_name="Apple Inc.", + ticker="AAPL", + ) + db_session.add(issuer) + + # Create document + doc = Document( + document_id="DOC::sec::ISSUER::0000320193::2026-01-29::ACC001", + source_name="sec", + issuer_id="ISSUER::0000320193", + form_type="8-K", + filing_date=dt.date(2026, 1, 29), + accession_no="ACC001", + parsed_status="ready_for_parse", + ) + db_session.add(doc) + await db_session.flush() + + # Create event + event = Event( + event_id="EVT::DOC::sec::ISSUER::0000320193::2026-01-29::ACC001::earnings_release::0", + primary_document_id=doc.document_id, + issuer_id="ISSUER::0000320193", + event_type="earnings_release", + event_direction="bullish", + event_date=dt.date(2026, 1, 29), + parser_version="rule-1.0.0", + parse_confidence=0.78, + status="pending", + ) + db_session.add(event) + await db_session.flush() + + # Create event parse + parse = EventParse( + event_id=event.event_id, + parser_kind="rule", + parser_version="rule-1.0.0", + schema_version="1.0.0", + output_json={"event_type": "earnings_release"}, + validation_status="valid", + ) + db_session.add(parse) + await db_session.flush() + + assert parse.event_parse_id is not None + assert event.status == "pending" diff --git a/tests/integration/test_sync_jobs.py b/tests/integration/test_sync_jobs.py new file mode 100644 index 0000000..c24b3cd --- /dev/null +++ b/tests/integration/test_sync_jobs.py @@ -0,0 +1,67 @@ +"""Integration test: sync jobs (FRED/FINRA → DB rows).""" +import datetime as dt + +import pytest +from pytest_httpx import HTTPXMock + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_macro_series_insert(db_session, httpx_mock: HTTPXMock, fred_observations_fixture): + """FRED sync should create macro_series + macro_observations rows.""" + from sqlalchemy import select + + from libs.db.models import MacroObservation, MacroSeries + + # Upsert macro_series + series = MacroSeries( + series_id="DGS10", + title="10-Year Treasury", + frequency="daily", + source_name="fred", + ) + db_session.add(series) + await db_session.flush() + + # Insert observations + for obs in fred_observations_fixture["observations"]: + ob = MacroObservation( + series_id="DGS10", + observation_date=dt.date.fromisoformat(obs["date"]), + value=obs["value"], + ) + db_session.add(ob) + await db_session.flush() + + result = await db_session.execute( + select(MacroObservation).where(MacroObservation.series_id == "DGS10") + ) + rows = result.scalars().all() + assert len(rows) == 4 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_short_sale_daily_insert(db_session, short_volume_fixture): + """FINRA sync should create short_sale_daily rows.""" + from sqlalchemy import select + + from libs.db.models import ShortSaleDaily + + for entry in short_volume_fixture["data"]: + row = ShortSaleDaily( + ticker_raw="AAPL", + trade_date=dt.date.fromisoformat(entry["date"]), + short_volume=entry["short_volume"], + short_exempt_volume=entry.get("short_exempt_volume"), + total_volume=entry.get("total_volume"), + source_name="finra", + ) + db_session.add(row) + await db_session.flush() + + result = await db_session.execute( + select(ShortSaleDaily).where(ShortSaleDaily.ticker_raw == "AAPL") + ) + rows = result.scalars().all() + assert len(rows) == 3 diff --git a/tests/replay/__init__.py b/tests/replay/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/replay/test_determinism.py b/tests/replay/test_determinism.py new file mode 100644 index 0000000..29a67e1 --- /dev/null +++ b/tests/replay/test_determinism.py @@ -0,0 +1,56 @@ +"""Replay test: same fixture → same parser output (determinism).""" +import pytest + +SAMPLE_TEXT = """ +Item 2.02 Results of Operations + +Apple today announced record quarterly revenue of $124.3 billion. +Guidance raised for Q2. Demand remains strong. Margin expansion driven by Services. +Customer additions in enterprise continue. Adjusted EPS excludes non-GAAP items. +""" + + +@pytest.mark.replay +def test_parser_determinism(): + """Running parser twice on same text yields identical output.""" + from libs.parser.rule_parser import RuleBasedParser + + p = RuleBasedParser() + metadata = {"filing_date": "2026-01-29", "accepted_at_utc": "2026-01-29T21:05:00Z"} + + out1 = p.parse("DOC::test", "8-K", SAMPLE_TEXT, metadata) + out2 = p.parse("DOC::test", "8-K", SAMPLE_TEXT, metadata) + + assert out1.event_type == out2.event_type + assert out1.event_direction == out2.event_direction + assert out1.guidance.status == out2.guidance.status + assert out1.confidence.overall == out2.confidence.overall + assert out1.filing_time_bucket == out2.filing_time_bucket + + +@pytest.mark.replay +def test_parser_determinism_negative_text(): + """Determinism holds for negative/mixed text too.""" + from libs.parser.rule_parser import RuleBasedParser + + negative_text = """ + Item 2.02 Results of Operations + Revenue below expectations. Guidance lowered. Demand softness observed. + Convertible note offering announced. Margin compression continues. + """ + p = RuleBasedParser() + metadata = {"filing_date": "2026-02-01"} + out1 = p.parse("DOC::test2", "8-K", negative_text, metadata) + out2 = p.parse("DOC::test2", "8-K", negative_text, metadata) + assert out1.event_type == out2.event_type + assert out1.guidance.status == out2.guidance.status + + +@pytest.mark.replay +def test_feature_determinism(sample_parser_output): + """Event features are deterministic for same parser output.""" + from libs.features.event_features import compute_event_features + + f1 = compute_event_features(sample_parser_output) + f2 = compute_event_features(sample_parser_output) + assert f1 == f2 diff --git a/tests/replay/test_idempotency.py b/tests/replay/test_idempotency.py new file mode 100644 index 0000000..4ee4a47 --- /dev/null +++ b/tests/replay/test_idempotency.py @@ -0,0 +1,40 @@ +"""Replay test: running pipeline twice produces no duplicate rows.""" + +import pytest + + +@pytest.mark.replay +def test_parser_idempotency(): + """Running parser twice on same text → identical output dict.""" + from libs.parser.rule_parser import RuleBasedParser + + text = """ + Item 2.02 Results of Operations + Revenue guidance raised. Demand remains strong. Margin expansion. + Customer additions accelerating. Adjusted EPS excludes non-GAAP items. + """ + p = RuleBasedParser() + metadata = {"filing_date": "2026-01-29"} + out1 = p.parse("DOC::idem::test", "8-K", text, metadata).model_dump() + out2 = p.parse("DOC::idem::test", "8-K", text, metadata).model_dump() + assert out1 == out2 + + +@pytest.mark.replay +def test_checksum_idempotency(tmp_path, monkeypatch): + """Writing same exhibit twice: file content unchanged, checksum stable.""" + from libs.common import config + config.get_settings.cache_clear() + monkeypatch.setenv("DATA_ROOT", str(tmp_path)) + + from libs.common.file_store import get_checksum, write_exhibit + + content = "Idempotency test content" + c1 = write_exhibit("IDEM001", "EX-99.1", content) + # Second write to same path should overwrite (atomic) + c2 = write_exhibit("IDEM001", "EX-99.1", content) + assert c1 == c2 + + stored = get_checksum("IDEM001", "EX-99.1") + assert stored == c1 + config.get_settings.cache_clear() diff --git a/tests/unit/__init__.py b/tests/unit/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py new file mode 100644 index 0000000..40f17bd --- /dev/null +++ b/tests/unit/test_config.py @@ -0,0 +1,63 @@ +"""Unit tests for config module.""" +import pytest + + +def test_settings_defaults(monkeypatch): + """Settings load with expected defaults.""" + monkeypatch.delenv("POSTGRES_DSN", raising=False) + monkeypatch.delenv("STOCK_ORACLE_URL", raising=False) + + # Reset lru_cache + from libs.common import config + config.get_settings.cache_clear() + + settings = config.Settings( + postgres_dsn="postgresql+asyncpg://acef:acef@localhost:5432/acef", + stock_oracle_url="http://localhost:18001", + ) + assert settings.stock_oracle_url == "http://localhost:18001" + assert settings.log_level == "INFO" + assert settings.llm_enabled is False + config.get_settings.cache_clear() + + +def test_log_level_normalized(): + from libs.common.config import Settings + s = Settings(log_level="debug") + assert s.log_level == "DEBUG" + + +def test_exhibit_cache_dir(): + from libs.common.config import Settings + s = Settings(data_root="/tmp/acef") + assert str(s.exhibit_cache_dir) == "/tmp/acef/cache/exhibits" + + +def test_parquet_dir(): + from libs.common.config import Settings + s = Settings(data_root="/tmp/acef") + assert str(s.parquet_dir) == "/tmp/acef/parquet" + + +@pytest.mark.unit +def test_get_symbols(tmp_path): + """get_symbols returns list from YAML.""" + import yaml + + from libs.common.config import Settings + + symbols_file = tmp_path / "symbols.yaml" + symbols_file.write_text(yaml.dump({"symbols": ["AAPL", "MSFT"]})) + + import os + old = os.getcwd() + os.chdir(tmp_path) + (tmp_path / "configs").mkdir(exist_ok=True) + (tmp_path / "configs" / "symbols.yaml").write_text( + yaml.dump({"symbols": ["AAPL", "MSFT"]}) + ) + s = Settings() + syms = s.get_symbols() + os.chdir(old) + # May return [] if configs/symbols.yaml not in cwd — just check type + assert isinstance(syms, list) diff --git a/tests/unit/test_db_models.py b/tests/unit/test_db_models.py new file mode 100644 index 0000000..40413ac --- /dev/null +++ b/tests/unit/test_db_models.py @@ -0,0 +1,65 @@ +"""Unit tests for DB model definitions (no DB required).""" +import datetime as dt + + +def test_issuer_master_fields(): + from libs.db.models import IssuerMaster + m = IssuerMaster( + issuer_id="ISSUER::0000320193", + issuer_name="Apple Inc.", + ticker="AAPL", + is_active=True, + ) + assert m.issuer_id == "ISSUER::0000320193" + assert m.issuer_name == "Apple Inc." + assert m.is_active is True + + +def test_document_defaults(): + from libs.db.models import Document + doc = Document( + document_id="DOC::test", + source_name="sec", + form_type="8-K", + filing_date=dt.date(2026, 1, 29), + parsed_status="pending", + ) + assert doc.parsed_status == "pending" + + +def test_event_defaults(): + from libs.db.models import Event + evt = Event( + event_id="EVT::test", + primary_document_id="DOC::test", + event_type="earnings_release", + event_direction="bullish", + event_date=dt.date(2026, 1, 29), + parser_version="rule-1.0.0", + status="pending", + ) + assert evt.status == "pending" + + +def test_job_run_fields(): + import uuid + + from libs.db.models import JobRun + jr = JobRun( + job_run_id=uuid.uuid4(), + job_name="test_job", + status="pending", + records_seen=0, + records_written=0, + ) + assert jr.records_seen == 0 + assert jr.records_written == 0 + + +def test_enums(): + from libs.db.enums import EventDirection, EventType, JobStatus, ParsedStatus, ParserKind + assert JobStatus.succeeded.value == "succeeded" + assert ParsedStatus.ready_for_parse.value == "ready_for_parse" + assert EventDirection.bullish.value == "bullish" + assert EventType.earnings_release.value == "earnings_release" + assert ParserKind.rule.value == "rule" diff --git a/tests/unit/test_event_features.py b/tests/unit/test_event_features.py new file mode 100644 index 0000000..e109f1f --- /dev/null +++ b/tests/unit/test_event_features.py @@ -0,0 +1,87 @@ +"""Unit tests for event feature calculations.""" + + +def test_guidance_direction_score_raised(): + from libs.features.event_features import guidance_direction_score + from libs.schemas.types import GuidanceOutput + g = GuidanceOutput(status="raised", scope="unknown", notes="") + assert guidance_direction_score(g) == 1.0 + + +def test_guidance_direction_score_lowered(): + from libs.features.event_features import guidance_direction_score + from libs.schemas.types import GuidanceOutput + g = GuidanceOutput(status="lowered", scope="unknown", notes="") + assert guidance_direction_score(g) == 0.0 + + +def test_guidance_direction_score_maintained(): + from libs.features.event_features import guidance_direction_score + from libs.schemas.types import GuidanceOutput + g = GuidanceOutput(status="inline_or_maintained", scope="unknown", notes="") + assert guidance_direction_score(g) == 0.5 + + +def test_oneoff_penalty_no_flags(): + from libs.features.event_features import oneoff_penalty + from libs.schemas.types import RiskFlagsOutput + r = RiskFlagsOutput() + assert oneoff_penalty(r) == 0.0 + + +def test_oneoff_penalty_all_flags(): + from libs.features.event_features import oneoff_penalty + from libs.schemas.types import RiskFlagsOutput + r = RiskFlagsOutput( + oneoff_item=True, + tax_benefit=True, + valuation_gain=True, + non_gaap_heavy=True, + financing_related=True, + legal_or_regulatory_overhang=True, + ) + assert oneoff_penalty(r) == 1.0 + + +def test_signal_strength_score_all_positive(): + from libs.features.event_features import signal_strength_score + from libs.schemas.types import SignalsOutput + s = SignalsOutput( + demand_strength="strong", + pricing_power="present", + backlog_or_bookings="present", + customer_expansion="present", + margin_quality="improving", + ) + assert signal_strength_score(s) == 1.0 + + +def test_signal_strength_score_all_unknown(): + from libs.features.event_features import signal_strength_score + from libs.schemas.types import SignalsOutput + s = SignalsOutput( + demand_strength="unknown", + pricing_power="unknown", + backlog_or_bookings="unknown", + customer_expansion="unknown", + margin_quality="unknown", + ) + assert signal_strength_score(s) == 0.0 + + +def test_document_quality_score_range(sample_parser_output): + from libs.features.event_features import document_quality_score + from libs.schemas.types import ConfidenceOutput + c = ConfidenceOutput(**sample_parser_output["confidence"]) + score = document_quality_score(c) + assert 0.0 <= score <= 1.0 + + +def test_compute_event_features(sample_parser_output): + from libs.features.event_features import compute_event_features + features = compute_event_features(sample_parser_output) + assert "guidance_direction_score" in features + assert "oneoff_penalty" in features + assert "signal_strength_score" in features + assert "document_quality_score" in features + assert "event_type" in features diff --git a/tests/unit/test_file_store.py b/tests/unit/test_file_store.py new file mode 100644 index 0000000..4fb08bf --- /dev/null +++ b/tests/unit/test_file_store.py @@ -0,0 +1,43 @@ +"""Unit tests for file_store module.""" +import pytest + + +@pytest.fixture +def tmp_data_root(tmp_path, monkeypatch): + """Redirect data root to tmp for isolation.""" + from libs.common import config + config.get_settings.cache_clear() + monkeypatch.setenv("DATA_ROOT", str(tmp_path)) + yield tmp_path + config.get_settings.cache_clear() + + +def test_write_and_read_exhibit(tmp_data_root): + from libs.common.file_store import read_exhibit, write_exhibit + content = "This is exhibit content." + checksum = write_exhibit("ACC123", "EX-99.1", content) + assert len(checksum) == 64 + result = read_exhibit("ACC123", "EX-99.1") + assert result == content + + +def test_exists_exhibit(tmp_data_root): + from libs.common.file_store import exists_exhibit, write_exhibit + assert not exists_exhibit("ACC999", "EX-99.1") + write_exhibit("ACC999", "EX-99.1", "content") + assert exists_exhibit("ACC999", "EX-99.1") + + +def test_checksum_deterministic(tmp_data_root): + from libs.common.file_store import get_checksum, write_exhibit + write_exhibit("ACC456", "EX-99.1", "hello world") + c = get_checksum("ACC456", "EX-99.1") + assert c is not None + # Second call returns same + assert get_checksum("ACC456", "EX-99.1") == c + + +def test_exhibit_path_safe_chars(tmp_data_root): + from libs.common.file_store import exhibit_path + p = exhibit_path("ACC789", "EX-99.1") + assert "EX-99.1.txt" in str(p) diff --git a/tests/unit/test_ids.py b/tests/unit/test_ids.py new file mode 100644 index 0000000..d96867a --- /dev/null +++ b/tests/unit/test_ids.py @@ -0,0 +1,53 @@ +"""Unit tests for ids module.""" + + +def test_document_id_format(): + from libs.common.ids import document_id + did = document_id("sec", "ISSUER::0000320193", "2026-01-29", "0000320193-26-000001") + assert did.startswith("DOC::sec::ISSUER::0000320193::2026-01-29::") + + +def test_event_id_format(): + from libs.common.ids import event_id + eid = event_id("DOC::sec::x::2026-01-01::acc", "earnings_release") + assert eid.startswith("EVT::") + assert "earnings_release" in eid + + +def test_issuer_id_from_cik(): + from libs.common.ids import issuer_id_from_cik + iid = issuer_id_from_cik("320193") + assert iid == "ISSUER::0000320193" + + +def test_symbol_id_from_ticker(): + from libs.common.ids import symbol_id_from_ticker + sid = symbol_id_from_ticker("AAPL") + assert sid == "SYM::AAPL::XNYS" + + +def test_symbol_id_uppercase(): + from libs.common.ids import symbol_id_from_ticker + sid = symbol_id_from_ticker("aapl", "NASDAQ") + assert sid == "SYM::AAPL::NASDAQ" + + +def test_sha256_checksum(): + from libs.common.ids import sha256_checksum + c = sha256_checksum(b"hello") + assert len(c) == 64 + assert c == sha256_checksum(b"hello") + + +def test_sha256_different(): + from libs.common.ids import sha256_checksum + assert sha256_checksum(b"hello") != sha256_checksum(b"world") + + +def test_new_job_run_id(): + import uuid + + from libs.common.ids import new_job_run_id + rid = new_job_run_id() + # Must be valid UUID + uuid.UUID(rid) diff --git a/tests/unit/test_market_features.py b/tests/unit/test_market_features.py new file mode 100644 index 0000000..df4dea3 --- /dev/null +++ b/tests/unit/test_market_features.py @@ -0,0 +1,72 @@ +"""Unit tests for market feature calculations.""" + +from libs.oracle_client.models import PriceBar + +BARS = [ + PriceBar(date="2026-01-26", open=227.5, high=230.0, low=226.0, close=229.5, volume=48000000), + PriceBar(date="2026-01-27", open=229.0, high=232.0, low=228.0, close=231.0, volume=55000000), + PriceBar(date="2026-01-28", open=231.5, high=234.0, low=230.0, close=233.0, volume=60000000), + PriceBar(date="2026-01-29", open=240.0, high=245.0, low=238.0, close=243.0, volume=120000000), +] + + +def test_reaction_day_return(): + from libs.features.market_features import reaction_day_return + r = reaction_day_return(BARS, "2026-01-29") + assert r is not None + expected = (243.0 - 233.0) / 233.0 + assert abs(r - expected) < 1e-6 + + +def test_reaction_day_return_missing_date(): + from libs.features.market_features import reaction_day_return + assert reaction_day_return(BARS, "2026-12-01") is None + + +def test_volume_ratio_20d(): + from libs.features.market_features import volume_ratio_20d + r = volume_ratio_20d(BARS, "2026-01-29") + assert r is not None + avg = (48000000 + 55000000 + 60000000) / 3 + assert abs(r - 120000000 / avg) < 1e-3 + + +def test_close_location(): + from libs.features.market_features import close_location + bar = PriceBar(date="2026-01-29", open=240.0, high=245.0, low=238.0, close=243.0, volume=120000000) + cl = close_location(bar) + assert cl is not None + expected = (243.0 - 238.0) / (245.0 - 238.0) + assert abs(cl - expected) < 1e-6 + + +def test_close_location_zero_range(): + from libs.features.market_features import close_location + bar = PriceBar(date="2026-01-29", open=100.0, high=100.0, low=100.0, close=100.0, volume=1000) + assert close_location(bar) is None + + +def test_gap_size(): + from libs.features.market_features import gap_size + g = gap_size(BARS, "2026-01-29") + assert g is not None + expected = (240.0 - 233.0) / 233.0 + assert abs(g - expected) < 1e-6 + + +def test_atr_14(): + from libs.features.market_features import atr_14 + # With fewer than 14 bars, should still return avg TR + r = atr_14(BARS) + assert r is not None + assert r > 0 + + +def test_compute_market_features_dict(): + from libs.features.market_features import compute_market_features + features = compute_market_features(BARS, "2026-01-29") + assert "reaction_day_return" in features + assert "volume_ratio_20d" in features + assert "gap_size" in features + assert "atr_14" in features + assert "close_location" in features diff --git a/tests/unit/test_oracle_client.py b/tests/unit/test_oracle_client.py new file mode 100644 index 0000000..9062dd8 --- /dev/null +++ b/tests/unit/test_oracle_client.py @@ -0,0 +1,122 @@ +"""Unit tests for Oracle client using pytest-httpx.""" +import json +from pathlib import Path + +import pytest +from pytest_httpx import HTTPXMock + +FIXTURES_DIR = Path(__file__).parent.parent / "fixtures" + + +def load_fixture(name: str) -> dict: + return json.loads((FIXTURES_DIR / name).read_text()) + + +@pytest.mark.asyncio +async def test_search_filings(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.filings import FilingsService + + data = load_fixture("filing_search.json") + httpx_mock.add_response(json=data, url="http://oracle:18001/filings/search/AAPL") + + async with OracleClient("http://oracle:18001") as client: + svc = FilingsService(client) + result = await svc.search_filings("AAPL") + + assert result.ticker == "AAPL" + assert len(result.filings) == 2 + assert result.filings[0].form_type == "8-K" + + +@pytest.mark.asyncio +async def test_get_exhibit(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.filings import FilingsService + + data = load_fixture("exhibit_content.json") + httpx_mock.add_response(json=data) + + async with OracleClient("http://oracle:18001") as client: + svc = FilingsService(client) + result = await svc.get_exhibit("0000320193-26-000001") + + assert result.accession_no == "0000320193-26-000001" + assert len(result.content) > 0 + + +@pytest.mark.asyncio +async def test_get_daily_bars(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.price import PriceService + + data = load_fixture("price_data.json") + httpx_mock.add_response(json=data) + + async with OracleClient("http://oracle:18001") as client: + svc = PriceService(client) + result = await svc.get_daily_bars("AAPL") + + assert result.ticker == "AAPL" + assert len(result.bars) > 0 + assert result.bars[0].close > 0 + + +@pytest.mark.asyncio +async def test_not_found_raises_oracle_not_found(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.exceptions import OracleNotFoundError + from libs.oracle_client.filings import FilingsService + + httpx_mock.add_response(status_code=404) + + async with OracleClient("http://oracle:18001") as client: + svc = FilingsService(client) + with pytest.raises(OracleNotFoundError): + await svc.get_exhibit("NONEXISTENT") + + +@pytest.mark.asyncio +async def test_server_error_raises_oracle_server_error(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.exceptions import OracleServerError + from libs.oracle_client.filings import FilingsService + + httpx_mock.add_response(status_code=500) + + async with OracleClient("http://oracle:18001") as client: + svc = FilingsService(client) + with pytest.raises(OracleServerError): + await svc.get_exhibit("ACC123") + + +@pytest.mark.asyncio +async def test_get_short_volume(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.finra import FinraService + + data = load_fixture("short_volume.json") + httpx_mock.add_response(json=data) + + async with OracleClient("http://oracle:18001") as client: + svc = FinraService(client) + result = await svc.get_short_volume("AAPL") + + assert result.symbol == "AAPL" + assert len(result.data) == 3 + + +@pytest.mark.asyncio +async def test_get_financial_data(httpx_mock: HTTPXMock): + from libs.oracle_client.client import OracleClient + from libs.oracle_client.financial import FinancialService + + data = load_fixture("financial_data.json") + httpx_mock.add_response(json=data) + + async with OracleClient("http://oracle:18001") as client: + svc = FinancialService(client) + result = await svc.get_financial_data("AAPL") + + assert result.ticker == "AAPL" + assert len(result.periods) == 2 diff --git a/tests/unit/test_retries.py b/tests/unit/test_retries.py new file mode 100644 index 0000000..e42e80a --- /dev/null +++ b/tests/unit/test_retries.py @@ -0,0 +1,77 @@ +"""Unit tests for retries module.""" +import pytest + + +def test_exception_hierarchy(): + from libs.common.retries import ( + ACEFError, + DependencyError, + NonRetryableError, + RetryableError, + ValidationError, + ) + assert issubclass(RetryableError, ACEFError) + assert issubclass(NonRetryableError, ACEFError) + assert issubclass(ValidationError, ACEFError) + assert issubclass(DependencyError, ACEFError) + + +def test_error_fields(): + from libs.common.retries import RetryableError + err = RetryableError("msg", source="sec", entity="CIK123", context={"url": "x"}) + assert err.source == "sec" + assert err.entity == "CIK123" + assert err.context["url"] == "x" + + +@pytest.mark.asyncio +async def test_with_retry_succeeds_on_first_attempt(): + from libs.common.retries import with_retry + + call_count = 0 + + @with_retry(max_attempts=3) + async def func(): + nonlocal call_count + call_count += 1 + return "ok" + + result = await func() + assert result == "ok" + assert call_count == 1 + + +@pytest.mark.asyncio +async def test_with_retry_retries_on_retryable_error(): + from libs.common.retries import RetryableError, with_retry + + call_count = 0 + + @with_retry(max_attempts=3, min_wait=0.01, max_wait=0.1) + async def func(): + nonlocal call_count + call_count += 1 + if call_count < 3: + raise RetryableError("transient") + return "ok" + + result = await func() + assert result == "ok" + assert call_count == 3 + + +@pytest.mark.asyncio +async def test_with_retry_does_not_retry_non_retryable(): + from libs.common.retries import NonRetryableError, with_retry + + call_count = 0 + + @with_retry(max_attempts=3, min_wait=0.01, max_wait=0.1) + async def func(): + nonlocal call_count + call_count += 1 + raise NonRetryableError("permanent") + + with pytest.raises(NonRetryableError): + await func() + assert call_count == 1 diff --git a/tests/unit/test_rule_parser.py b/tests/unit/test_rule_parser.py new file mode 100644 index 0000000..08a843a --- /dev/null +++ b/tests/unit/test_rule_parser.py @@ -0,0 +1,110 @@ +"""Unit tests for rule-based parser.""" + +SAMPLE_TEXT_POSITIVE = """ +Item 2.02 Results of Operations + +Apple today announced record quarterly revenue of $124.3 billion, up 9 percent year over year. + +Guidance raised for Q2: The Company expects revenue to be between $125 billion and $131 billion. +Demand remains strong across all product categories. +Gross margin expansion driven by Services mix. +Customer additions in enterprise continue to accelerate. +Backlog increased significantly year over year. +Adjusted EPS excludes certain non-GAAP items. +""" + +SAMPLE_TEXT_NEGATIVE = """ +Item 2.02 Results of Operations + +Company reported revenue of $5.2 billion, missing expectations. +Guidance lowered for next quarter due to demand softness. +Margin compression continues. Financing need announced with dilutive convertible note offering. +""" + +SAMPLE_TEXT_CONTRACT = """ +Item 1.01 Entry into Material Agreement + +The Company has entered into a definitive agreement with a major enterprise customer +for a multi-year supply contract worth $500 million. The contract includes recurring revenue. +""" + + +def test_parse_earnings_positive(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert out.event_type == "earnings_release" + assert out.guidance.status == "raised" + assert out.signals.demand_strength == "strong" + assert out.event_direction == "bullish" + + +def test_parse_earnings_negative(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_NEGATIVE, {"filing_date": "2026-02-01"}) + assert out.event_type == "earnings_release" + assert out.guidance.status == "lowered" + assert out.event_direction in ("bearish", "mixed") + + +def test_parse_material_contract(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_CONTRACT, {"filing_date": "2026-02-15"}) + assert out.event_type == "material_contract" + + +def test_parse_non_gaap_flag(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert out.risk_flags.non_gaap_heavy is True + + +def test_parse_margin_improving(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert out.signals.margin_quality == "improving" + + +def test_parse_confidence_range(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert 0.0 <= out.confidence.overall <= 1.0 + + +def test_parse_document_id_preserved(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::my-id", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert out.document_id == "DOC::my-id" + + +def test_parse_schema_version(): + from libs.parser.rule_parser import SCHEMA_VERSION, RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", SAMPLE_TEXT_POSITIVE, {"filing_date": "2026-01-29"}) + assert out.schema_version == SCHEMA_VERSION + + +def test_parse_unknown_event(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse("DOC::test", "8-K", "No items mentioned here.", {"filing_date": "2026-01-29"}) + assert out.event_type == "unknown" + assert len(out.warnings) > 0 + + +def test_parse_filing_time_bucket_set(): + from libs.parser.rule_parser import RuleBasedParser + p = RuleBasedParser() + out = p.parse( + "DOC::test", + "8-K", + SAMPLE_TEXT_POSITIVE, + {"filing_date": "2026-01-29", "accepted_at_utc": "2026-01-29T21:05:00Z"}, + ) + assert out.filing_time_bucket == "post_market" diff --git a/tests/unit/test_schema_validator.py b/tests/unit/test_schema_validator.py new file mode 100644 index 0000000..86bab34 --- /dev/null +++ b/tests/unit/test_schema_validator.py @@ -0,0 +1,39 @@ +"""Unit tests for parser schema validator.""" + + +def test_valid_output_passes(sample_parser_output): + from libs.parser.schema_validator import validate_parser_output + errors = validate_parser_output(sample_parser_output) + assert errors == [] + + +def test_missing_required_field(sample_parser_output): + from libs.parser.schema_validator import validate_parser_output + del sample_parser_output["event_type"] + errors = validate_parser_output(sample_parser_output) + assert len(errors) > 0 + + +def test_invalid_event_type(sample_parser_output): + from libs.parser.schema_validator import validate_parser_output + sample_parser_output["event_type"] = "not_a_valid_type" + errors = validate_parser_output(sample_parser_output) + assert len(errors) > 0 + + +def test_confidence_out_of_range(sample_parser_output): + from libs.parser.schema_validator import validate_parser_output + sample_parser_output["confidence"]["overall"] = 1.5 + errors = validate_parser_output(sample_parser_output) + assert len(errors) > 0 + + +def test_is_valid_function(sample_parser_output): + from libs.parser.schema_validator import is_valid + assert is_valid(sample_parser_output) is True + + +def test_is_valid_false(sample_parser_output): + from libs.parser.schema_validator import is_valid + sample_parser_output["event_direction"] = "invalid_value" + assert is_valid(sample_parser_output) is False diff --git a/tests/unit/test_text_normalizer.py b/tests/unit/test_text_normalizer.py new file mode 100644 index 0000000..53aa665 --- /dev/null +++ b/tests/unit/test_text_normalizer.py @@ -0,0 +1,49 @@ +"""Unit tests for text normalizer.""" + + +def test_html_to_text(): + from libs.parser.text_normalizer import html_to_text + html = "

Hello world

" + text = html_to_text(html) + assert "Hello" in text + assert "world" in text + assert "

" not in text + + +def test_normalize_unicode_quotes(): + from libs.parser.text_normalizer import normalize_unicode + text = "\u2018Hello\u2019 \u201cworld\u201d" + result = normalize_unicode(text) + assert "'" in result + assert '"' in result + + +def test_normalize_em_dash(): + from libs.parser.text_normalizer import normalize_unicode + text = "Q4\u2014best quarter" + result = normalize_unicode(text) + assert " - " in result + + +def test_collapse_whitespace(): + from libs.parser.text_normalizer import collapse_whitespace + text = "Hello\n\n\n\n\nWorld" + result = collapse_whitespace(text) + assert "Hello" in result + assert "World" in result + + +def test_normalize_text_pipeline(): + from libs.parser.text_normalizer import normalize_text + html = "

Revenue guidance raised. Demand remains strong.

" + result = normalize_text(html, is_html=True) + assert "Revenue" in result + assert "

" not in result + + +def test_script_removed(): + from libs.parser.text_normalizer import normalize_text + html = "

Content

" + result = normalize_text(html, is_html=True) + assert "var x" not in result + assert "Content" in result diff --git a/tests/unit/test_time_utils.py b/tests/unit/test_time_utils.py new file mode 100644 index 0000000..19e6082 --- /dev/null +++ b/tests/unit/test_time_utils.py @@ -0,0 +1,59 @@ +"""Unit tests for time_utils module.""" +import datetime as dt +from zoneinfo import ZoneInfo + +import pytest + +_UTC = ZoneInfo("UTC") +_ET = ZoneInfo("America/New_York") + + +def test_utc_now(): + from libs.common.time_utils import utc_now + now = utc_now() + assert now.tzinfo is not None + + +def test_to_eastern(): + from libs.common.time_utils import to_eastern + d = dt.datetime(2026, 1, 29, 21, 5, tzinfo=_UTC) + e = to_eastern(d) + assert e.tzinfo is not None + assert e.hour == 16 # 21:05 UTC = 16:05 ET (EST, UTC-5) + + +def test_to_utc(): + from libs.common.time_utils import to_utc + d = dt.datetime(2026, 1, 29, 16, 5, tzinfo=_ET) + u = to_utc(d) + assert u.tzinfo is not None + + +def test_naive_datetime_rejected(): + from libs.common.time_utils import to_eastern + with pytest.raises(ValueError): + to_eastern(dt.datetime(2026, 1, 1)) + + +def test_filing_time_bucket_post_market(): + from libs.common.time_utils import filing_time_bucket + d = dt.datetime(2026, 1, 29, 21, 5, tzinfo=_UTC) + assert filing_time_bucket(d) == "post_market" + + +def test_filing_time_bucket_pre_market(): + from libs.common.time_utils import filing_time_bucket + d = dt.datetime(2026, 1, 29, 12, 0, tzinfo=_UTC) # 7:00 AM ET + assert filing_time_bucket(d) == "pre_market" + + +def test_filing_time_bucket_regular(): + from libs.common.time_utils import filing_time_bucket + d = dt.datetime(2026, 1, 29, 15, 0, tzinfo=_UTC) # 10:00 AM ET + assert filing_time_bucket(d) == "regular_hours" + + +def test_filing_time_bucket_naive(): + from libs.common.time_utils import filing_time_bucket + d = dt.datetime(2026, 1, 29, 21, 5) + assert filing_time_bucket(d) == "unknown"