storage.py 6.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225
  1. """MySQL state and audit storage for real-time control."""
  2. from __future__ import annotations
  3. import os
  4. from contextlib import contextmanager
  5. from datetime import date, datetime
  6. from pathlib import Path
  7. from typing import Any, Iterator
  8. import pymysql
  9. import pymysql.cursors
  10. ROOT = Path(__file__).resolve().parent
  11. def connect() -> pymysql.Connection:
  12. required = ["DB_HOST", "DB_USER", "DB_NAME"]
  13. missing = [key for key in required if not os.getenv(key)]
  14. if missing:
  15. raise RuntimeError(f"Missing database environment variables: {', '.join(missing)}")
  16. return pymysql.connect(
  17. host=os.environ["DB_HOST"],
  18. port=int(os.getenv("DB_PORT", "3306")),
  19. user=os.environ["DB_USER"],
  20. password=os.getenv("DB_PASSWORD", ""),
  21. database=os.environ["DB_NAME"],
  22. charset="utf8mb4",
  23. cursorclass=pymysql.cursors.DictCursor,
  24. autocommit=True,
  25. connect_timeout=int(os.getenv("DB_CONNECT_TIMEOUT", "10")),
  26. read_timeout=int(os.getenv("DB_READ_TIMEOUT", "60")),
  27. write_timeout=int(os.getenv("DB_WRITE_TIMEOUT", "60")),
  28. )
  29. def initialize_schema() -> None:
  30. statements = [
  31. statement.strip()
  32. for statement in (ROOT / "schema.sql").read_text(encoding="utf-8").split(";")
  33. if statement.strip()
  34. ]
  35. connection = connect()
  36. try:
  37. with connection.cursor() as cursor:
  38. for statement in statements:
  39. cursor.execute(statement)
  40. finally:
  41. connection.close()
  42. @contextmanager
  43. def advisory_lock(lock_name: str) -> Iterator[bool]:
  44. connection = connect()
  45. acquired = False
  46. try:
  47. with connection.cursor() as cursor:
  48. cursor.execute("SELECT GET_LOCK(%s, 0) AS acquired", (lock_name,))
  49. acquired = bool((cursor.fetchone() or {}).get("acquired"))
  50. yield acquired
  51. finally:
  52. if acquired:
  53. try:
  54. with connection.cursor() as cursor:
  55. cursor.execute("SELECT RELEASE_LOCK(%s)", (lock_name,))
  56. except Exception:
  57. pass
  58. connection.close()
  59. def load_enabled_accounts() -> list[dict[str, Any]]:
  60. connection = connect()
  61. try:
  62. with connection.cursor() as cursor:
  63. cursor.execute(
  64. """
  65. SELECT c.account_id, c.audience_name, c.bid_scene
  66. FROM ad_creation_account_config c
  67. JOIN account_whitelist w ON w.account_id = c.account_id
  68. WHERE c.enabled = TRUE
  69. AND w.enabled = TRUE
  70. ORDER BY c.account_id
  71. """
  72. )
  73. return list(cursor.fetchall())
  74. finally:
  75. connection.close()
  76. def load_daily_state(control_date: date) -> dict[str, Any]:
  77. connection = connect()
  78. try:
  79. with connection.cursor() as cursor:
  80. cursor.execute(
  81. "SELECT * FROM realtime_control_daily_state WHERE control_date=%s",
  82. (control_date,),
  83. )
  84. return cursor.fetchone() or {}
  85. finally:
  86. connection.close()
  87. def save_daily_state(control_date: date, **values: Any) -> None:
  88. allowed = {
  89. "morning_recovery_done",
  90. "cutoff_done",
  91. "last_observed_partition",
  92. "last_observed_cpm",
  93. "last_decision",
  94. "last_inventory_refresh_at",
  95. "last_evaluated_at",
  96. }
  97. unknown = set(values) - allowed
  98. if unknown:
  99. raise ValueError(f"Unsupported daily-state fields: {sorted(unknown)}")
  100. columns = ["control_date", *values]
  101. params = [control_date, *values.values()]
  102. updates = ", ".join(f"{column}=VALUES({column})" for column in values)
  103. placeholders = ", ".join(["%s"] * len(columns))
  104. sql = (
  105. f"INSERT INTO realtime_control_daily_state ({', '.join(columns)}) "
  106. f"VALUES ({placeholders}) ON DUPLICATE KEY UPDATE {updates}"
  107. )
  108. connection = connect()
  109. try:
  110. with connection.cursor() as cursor:
  111. cursor.execute(sql, params)
  112. finally:
  113. connection.close()
  114. def load_ad_states(account_id: int) -> dict[int, dict[str, Any]]:
  115. connection = connect()
  116. try:
  117. with connection.cursor() as cursor:
  118. cursor.execute(
  119. "SELECT * FROM realtime_control_ad_state WHERE account_id=%s",
  120. (account_id,),
  121. )
  122. return {int(row["adgroup_id"]): row for row in cursor.fetchall()}
  123. finally:
  124. connection.close()
  125. def upsert_ad_state(
  126. *,
  127. account_id: int,
  128. adgroup_id: int,
  129. adgroup_name: str,
  130. bid_field: str,
  131. base_bid_fen: int,
  132. boosted_date: date | None,
  133. paused_by_strategy: bool,
  134. pause_reason: str | None,
  135. last_action: str,
  136. action_at: datetime,
  137. ) -> None:
  138. connection = connect()
  139. try:
  140. with connection.cursor() as cursor:
  141. cursor.execute(
  142. """
  143. INSERT INTO realtime_control_ad_state
  144. (account_id, adgroup_id, adgroup_name, bid_field, base_bid_fen,
  145. boosted_date, paused_by_strategy, pause_reason, last_action,
  146. last_action_at, last_seen_at)
  147. VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)
  148. ON DUPLICATE KEY UPDATE
  149. adgroup_name=VALUES(adgroup_name),
  150. bid_field=VALUES(bid_field),
  151. boosted_date=VALUES(boosted_date),
  152. paused_by_strategy=VALUES(paused_by_strategy),
  153. pause_reason=VALUES(pause_reason),
  154. last_action=VALUES(last_action),
  155. last_action_at=VALUES(last_action_at),
  156. last_seen_at=VALUES(last_seen_at)
  157. """,
  158. (
  159. account_id,
  160. adgroup_id,
  161. adgroup_name,
  162. bid_field,
  163. base_bid_fen,
  164. boosted_date,
  165. paused_by_strategy,
  166. pause_reason,
  167. last_action,
  168. action_at,
  169. action_at,
  170. ),
  171. )
  172. finally:
  173. connection.close()
  174. def insert_action_log(record: dict[str, Any]) -> None:
  175. columns = [
  176. "run_id",
  177. "control_date",
  178. "observed_partition",
  179. "observed_cpm",
  180. "decision",
  181. "account_id",
  182. "adgroup_id",
  183. "adgroup_name",
  184. "bid_field",
  185. "base_bid_fen",
  186. "before_bid_fen",
  187. "target_bid_fen",
  188. "before_status",
  189. "target_status",
  190. "apply_mode",
  191. "execution_status",
  192. "error_message",
  193. ]
  194. connection = connect()
  195. try:
  196. with connection.cursor() as cursor:
  197. cursor.execute(
  198. f"INSERT INTO realtime_control_action_log ({', '.join(columns)}) "
  199. f"VALUES ({', '.join(['%s'] * len(columns))})",
  200. [record.get(column) for column in columns],
  201. )
  202. finally:
  203. connection.close()