| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225 |
- """MySQL state and audit storage for real-time control."""
- from __future__ import annotations
- import os
- from contextlib import contextmanager
- from datetime import date, datetime
- from pathlib import Path
- from typing import Any, Iterator
- import pymysql
- import pymysql.cursors
- ROOT = Path(__file__).resolve().parent
- def connect() -> pymysql.Connection:
- required = ["DB_HOST", "DB_USER", "DB_NAME"]
- missing = [key for key in required if not os.getenv(key)]
- if missing:
- raise RuntimeError(f"Missing database environment variables: {', '.join(missing)}")
- return pymysql.connect(
- host=os.environ["DB_HOST"],
- port=int(os.getenv("DB_PORT", "3306")),
- user=os.environ["DB_USER"],
- password=os.getenv("DB_PASSWORD", ""),
- database=os.environ["DB_NAME"],
- charset="utf8mb4",
- cursorclass=pymysql.cursors.DictCursor,
- autocommit=True,
- connect_timeout=int(os.getenv("DB_CONNECT_TIMEOUT", "10")),
- read_timeout=int(os.getenv("DB_READ_TIMEOUT", "60")),
- write_timeout=int(os.getenv("DB_WRITE_TIMEOUT", "60")),
- )
- def initialize_schema() -> None:
- statements = [
- statement.strip()
- for statement in (ROOT / "schema.sql").read_text(encoding="utf-8").split(";")
- if statement.strip()
- ]
- connection = connect()
- try:
- with connection.cursor() as cursor:
- for statement in statements:
- cursor.execute(statement)
- finally:
- connection.close()
- @contextmanager
- def advisory_lock(lock_name: str) -> Iterator[bool]:
- connection = connect()
- acquired = False
- try:
- with connection.cursor() as cursor:
- cursor.execute("SELECT GET_LOCK(%s, 0) AS acquired", (lock_name,))
- acquired = bool((cursor.fetchone() or {}).get("acquired"))
- yield acquired
- finally:
- if acquired:
- try:
- with connection.cursor() as cursor:
- cursor.execute("SELECT RELEASE_LOCK(%s)", (lock_name,))
- except Exception:
- pass
- connection.close()
- def load_enabled_accounts() -> list[dict[str, Any]]:
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(
- """
- SELECT c.account_id, c.audience_name, c.bid_scene
- FROM ad_creation_account_config c
- JOIN account_whitelist w ON w.account_id = c.account_id
- WHERE c.enabled = TRUE
- AND w.enabled = TRUE
- ORDER BY c.account_id
- """
- )
- return list(cursor.fetchall())
- finally:
- connection.close()
- def load_daily_state(control_date: date) -> dict[str, Any]:
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(
- "SELECT * FROM realtime_control_daily_state WHERE control_date=%s",
- (control_date,),
- )
- return cursor.fetchone() or {}
- finally:
- connection.close()
- def save_daily_state(control_date: date, **values: Any) -> None:
- allowed = {
- "morning_recovery_done",
- "cutoff_done",
- "last_observed_partition",
- "last_observed_cpm",
- "last_decision",
- "last_inventory_refresh_at",
- "last_evaluated_at",
- }
- unknown = set(values) - allowed
- if unknown:
- raise ValueError(f"Unsupported daily-state fields: {sorted(unknown)}")
- columns = ["control_date", *values]
- params = [control_date, *values.values()]
- updates = ", ".join(f"{column}=VALUES({column})" for column in values)
- placeholders = ", ".join(["%s"] * len(columns))
- sql = (
- f"INSERT INTO realtime_control_daily_state ({', '.join(columns)}) "
- f"VALUES ({placeholders}) ON DUPLICATE KEY UPDATE {updates}"
- )
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(sql, params)
- finally:
- connection.close()
- def load_ad_states(account_id: int) -> dict[int, dict[str, Any]]:
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(
- "SELECT * FROM realtime_control_ad_state WHERE account_id=%s",
- (account_id,),
- )
- return {int(row["adgroup_id"]): row for row in cursor.fetchall()}
- finally:
- connection.close()
- def upsert_ad_state(
- *,
- account_id: int,
- adgroup_id: int,
- adgroup_name: str,
- bid_field: str,
- base_bid_fen: int,
- boosted_date: date | None,
- paused_by_strategy: bool,
- pause_reason: str | None,
- last_action: str,
- action_at: datetime,
- ) -> None:
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(
- """
- INSERT INTO realtime_control_ad_state
- (account_id, adgroup_id, adgroup_name, bid_field, base_bid_fen,
- boosted_date, paused_by_strategy, pause_reason, last_action,
- last_action_at, last_seen_at)
- VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)
- ON DUPLICATE KEY UPDATE
- adgroup_name=VALUES(adgroup_name),
- bid_field=VALUES(bid_field),
- boosted_date=VALUES(boosted_date),
- paused_by_strategy=VALUES(paused_by_strategy),
- pause_reason=VALUES(pause_reason),
- last_action=VALUES(last_action),
- last_action_at=VALUES(last_action_at),
- last_seen_at=VALUES(last_seen_at)
- """,
- (
- account_id,
- adgroup_id,
- adgroup_name,
- bid_field,
- base_bid_fen,
- boosted_date,
- paused_by_strategy,
- pause_reason,
- last_action,
- action_at,
- action_at,
- ),
- )
- finally:
- connection.close()
- def insert_action_log(record: dict[str, Any]) -> None:
- columns = [
- "run_id",
- "control_date",
- "observed_partition",
- "observed_cpm",
- "decision",
- "account_id",
- "adgroup_id",
- "adgroup_name",
- "bid_field",
- "base_bid_fen",
- "before_bid_fen",
- "target_bid_fen",
- "before_status",
- "target_status",
- "apply_mode",
- "execution_status",
- "error_message",
- ]
- connection = connect()
- try:
- with connection.cursor() as cursor:
- cursor.execute(
- f"INSERT INTO realtime_control_action_log ({', '.join(columns)}) "
- f"VALUES ({', '.join(['%s'] * len(columns))})",
- [record.get(column) for column in columns],
- )
- finally:
- connection.close()
|