Files
dodo/backend/today_environment.py
bboysoul ed123b710f
ci / gitleaks (push) Successful in 8s
ci / docker (push) Successful in 3m14s
fix: fall back to latest SGE trading price
2026-09-20 08:33:55 +08:00

369 lines
13 KiB
Python

import asyncio
import logging
import math
import re
import threading
import weakref
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal, InvalidOperation
from html import unescape
from typing import Any
from zoneinfo import ZoneInfo
import httpx
from .lunar_support import solar_to_lunar_text
SHANGHAI_TZ = ZoneInfo("Asia/Shanghai")
WEATHER_URL = "https://api.open-meteo.com/v1/forecast"
GOLD_URL = "https://www.sge.com.cn/sjzx/yshqbg"
GOLD_HISTORY_URL = (
"https://www.sge.com.cn/sjzx/quotation_daily_new"
"?start_date={start_date}&end_date={end_date}"
)
WEATHER_SOURCE = "Open-Meteo"
GOLD_SOURCE = "上海黄金交易所"
WEEKDAYS = ("星期一", "星期二", "星期三", "星期四", "星期五", "星期六", "星期日")
SOURCE_TIMEOUT_SECONDS = 3.0
GOLD_SOURCE_TIMEOUT_SECONDS = 6.0
TOTAL_TIMEOUT_SECONDS = 6.5
WEATHER_TTL = timedelta(minutes=15)
WEATHER_STALE_TTL = timedelta(hours=6)
GOLD_TTL = timedelta(minutes=5)
GOLD_STALE_TTL = timedelta(days=7)
logger = logging.getLogger(__name__)
_TAG_RE = re.compile(r"<[^>]+>")
_DATE_RE = re.compile(r"上海黄金交易所\s*(\d{4})年(\d{2})月(\d{2})日\s*延时行情")
_ROW_RE = re.compile(r"<tr\b[^>]*>(.*?)</tr>", re.IGNORECASE | re.DOTALL)
_CELL_RE = re.compile(r"<t[dh]\b[^>]*>(.*?)</t[dh]>", re.IGNORECASE | re.DOTALL)
@dataclass
class CacheEntry:
value: dict[str, Any]
fetched_at: datetime
@dataclass
class LoopRuntime:
inflight: dict[tuple[int, str], asyncio.Task[dict[str, Any]]]
_cache: dict[str, CacheEntry] = {}
_generation = 0
_state_guard = threading.Lock()
_runtimes: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, LoopRuntime] = (
weakref.WeakKeyDictionary()
)
def _runtime(loop: asyncio.AbstractEventLoop | None = None) -> LoopRuntime:
active_loop = loop or asyncio.get_running_loop()
with _state_guard:
runtime = _runtimes.get(active_loop)
if runtime is None:
runtime = LoopRuntime(inflight={})
_runtimes[active_loop] = runtime
return runtime
def clear_cache() -> None:
global _generation
with _state_guard:
_generation += 1
_cache.clear()
def _text(fragment: str) -> str:
return " ".join(unescape(_TAG_RE.sub(" ", fragment)).split())
def parse_sge_au9999(html: str) -> tuple[Decimal, str]:
date_match = _DATE_RE.search(_text(html))
if date_match is None:
raise ValueError("SGE market date not found")
try:
year, month, day = map(int, date_match.groups())
market_date = date(year, month, day).isoformat()
except ValueError as exc:
raise ValueError("invalid SGE market date") from exc
rows = [[_text(cell) for cell in _CELL_RE.findall(row)] for row in _ROW_RE.findall(html)]
header = next((cells for cells in rows if "合约" in cells and "最新价" in cells), None)
if header is None:
raise ValueError("SGE quote table header not found")
contract_index = header.index("合约")
price_index = header.index("最新价")
required_length = max(contract_index, price_index) + 1
for cells in rows[rows.index(header) + 1 :]:
if len(cells) < required_length or cells[contract_index] != "Au99.99":
continue
try:
price = Decimal(cells[price_index].replace(",", ""))
except InvalidOperation as exc:
raise ValueError("invalid Au99.99 latest price") from exc
if not price.is_finite() or price <= 0:
raise ValueError("invalid Au99.99 latest price")
return price, market_date
raise ValueError("Au99.99 quote not found")
def parse_sge_au9999_history(html: str) -> tuple[Decimal, str]:
rows = [[_text(cell) for cell in _CELL_RE.findall(row)] for row in _ROW_RE.findall(html)]
header = next(
(cells for cells in rows if "日期" in cells and "合约" in cells and "收盘价" in cells),
None,
)
if header is None:
raise ValueError("SGE history table header not found")
date_index = header.index("日期")
contract_index = header.index("合约")
price_index = header.index("收盘价")
required_length = max(date_index, contract_index, price_index) + 1
quotes = []
for cells in rows[rows.index(header) + 1 :]:
if len(cells) < required_length or cells[contract_index] != "Au99.99":
continue
try:
market_date = date.fromisoformat(cells[date_index]).isoformat()
price = Decimal(cells[price_index].replace(",", ""))
except (InvalidOperation, ValueError):
continue
if not price.is_finite() or price <= 0:
continue
quotes.append((market_date, price))
if not quotes:
raise ValueError("Au99.99 history quote not found")
market_date, price = max(quotes)
return price, market_date
async def fetch_weather(request: Callable[..., Awaitable[dict[str, Any]]]) -> dict[str, Any]:
payload = await request(
WEATHER_URL,
params={
"latitude": 29.88,
"longitude": 121.55,
"current": "temperature_2m,apparent_temperature,weather_code",
"timezone": "Asia/Shanghai",
},
)
if payload.get("timezone") != "Asia/Shanghai":
raise ValueError("weather timezone must be Asia/Shanghai")
current = payload["current"]
if not isinstance(current, dict):
raise TypeError("weather current must be an object")
def finite_number(name: str, minimum: float, maximum: float) -> int | float:
value = current[name]
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise TypeError(f"weather {name} must be numeric")
if not math.isfinite(value) or not minimum <= value <= maximum:
raise ValueError(f"weather {name} out of range")
return value
temperature = finite_number("temperature_2m", -100, 100)
apparent_temperature = finite_number("apparent_temperature", -100, 100)
weather_code = current["weather_code"]
if isinstance(weather_code, bool) or not isinstance(weather_code, int):
raise TypeError("weather weather_code must be an integer")
if not 0 <= weather_code <= 99:
raise ValueError("weather weather_code out of range")
raw_time = current["time"]
if not isinstance(raw_time, str):
raise TypeError("weather time must be a string")
observed_at = datetime.fromisoformat(raw_time)
if observed_at.tzinfo is not None:
if observed_at.utcoffset() != SHANGHAI_TZ.utcoffset(observed_at):
raise ValueError("weather time must use Asia/Shanghai")
observed_at = observed_at.astimezone(SHANGHAI_TZ)
else:
observed_at = observed_at.replace(tzinfo=SHANGHAI_TZ)
return {
"temperature_c": temperature,
"apparent_temperature_c": apparent_temperature,
"weather_code": weather_code,
"observed_at": observed_at.isoformat(),
"source": WEATHER_SOURCE,
}
async def fetch_gold(request: Callable[..., Awaitable[str]]) -> dict[str, Any]:
html = await request(GOLD_URL)
try:
price, market_date = parse_sge_au9999(html)
except ValueError:
date_match = _DATE_RE.search(_text(html))
if date_match is None:
raise
try:
quote_date = date(*map(int, date_match.groups()))
except ValueError as exc:
raise ValueError("invalid SGE market date") from exc
start_date = quote_date - timedelta(days=31)
history_url = GOLD_HISTORY_URL.format(
start_date=start_date.isoformat(), end_date=quote_date.isoformat()
)
price, market_date = parse_sge_au9999_history(await request(history_url))
return {
"contract": "Au99.99",
"latest_price": price,
"currency": "CNY",
"unit": "gram",
"market_date": market_date,
"delayed": True,
"source": GOLD_SOURCE,
}
async def _request_json(client: httpx.AsyncClient, url: str, **kwargs: Any) -> dict[str, Any]:
response = await client.get(url, **kwargs)
response.raise_for_status()
return response.json()
async def _request_text(client: httpx.AsyncClient, url: str, **kwargs: Any) -> str:
response = await client.get(url, **kwargs)
response.raise_for_status()
return response.text
def _serializable(value: dict[str, Any]) -> dict[str, Any]:
return {key: str(item) if isinstance(item, Decimal) else item for key, item in value.items()}
async def _refresh_source(
name: str,
fetcher: Callable[[], Awaitable[dict[str, Any]]],
fetched_at: datetime,
generation: int,
) -> dict[str, Any]:
value = await asyncio.wait_for(
fetcher(), timeout=GOLD_SOURCE_TIMEOUT_SECONDS if name == "gold" else SOURCE_TIMEOUT_SECONDS
)
with _state_guard:
if generation == _generation:
newest = _cache.get(name)
if newest is None or newest.fetched_at <= fetched_at:
_cache[name] = CacheEntry(value=value, fetched_at=fetched_at)
return value
def _source_error_code(exc: BaseException) -> str:
if isinstance(exc, TimeoutError):
return "timeout"
if isinstance(exc, (KeyError, TypeError, ValueError, InvalidOperation)):
return "invalid_upstream_response"
return "upstream_unavailable"
def _finish_inflight(
runtime: LoopRuntime,
key: tuple[int, str],
task: asyncio.Task[dict[str, Any]],
) -> None:
# Calling exception() marks failures as retrieved even when every shielded waiter
# was cancelled. This synchronous callback cannot create an untracked cleanup task.
if not task.cancelled():
try:
task.exception()
except asyncio.CancelledError:
pass
if runtime.inflight.get(key) is task:
runtime.inflight.pop(key, None)
async def _cached_source(
name: str,
fetcher: Callable[[], Awaitable[dict[str, Any]]],
now: datetime,
fresh_for: timedelta,
stale_for: timedelta,
) -> tuple[dict[str, Any] | None, str | None]:
runtime = _runtime()
with _state_guard:
generation = _generation
key = (generation, name)
with _state_guard:
cached = _cache.get(name) if generation == _generation else None
if cached is not None and now - cached.fetched_at <= fresh_for:
return {**_serializable(cached.value), "stale": False}, None
task = runtime.inflight.get(key)
if task is None:
task = asyncio.create_task(_refresh_source(name, fetcher, now, generation))
runtime.inflight[key] = task
task.add_done_callback(
lambda completed, active=runtime, task_key=key: _finish_inflight(
active, task_key, completed
)
)
try:
value = await asyncio.shield(task)
except Exception as exc:
error_code = _source_error_code(exc)
logger.warning("Today environment source %s failed (%s)", name, error_code, exc_info=exc)
with _state_guard:
generation_is_current = generation == _generation
if generation_is_current and cached is not None and now - cached.fetched_at <= stale_for:
return {**_serializable(cached.value), "stale": True}, error_code
return None, error_code
return {**_serializable(value), "stale": False}, None
async def _default_weather_fetcher() -> dict[str, Any]:
async with httpx.AsyncClient(
timeout=httpx.Timeout(SOURCE_TIMEOUT_SECONDS),
headers={"User-Agent": "dodo/0.1 (+https://dodo.bboy.app)"},
) as client:
return await fetch_weather(lambda url, **kwargs: _request_json(client, url, **kwargs))
async def _default_gold_fetcher() -> dict[str, Any]:
async with httpx.AsyncClient(
timeout=httpx.Timeout(SOURCE_TIMEOUT_SECONDS),
headers={"User-Agent": "dodo/0.1 (+https://dodo.bboy.app)"},
) as client:
return await fetch_gold(lambda url, **kwargs: _request_text(client, url, **kwargs))
async def get_environment(
*,
now: datetime | None = None,
weather_fetcher: Callable[[], Awaitable[dict[str, Any]]] | None = None,
gold_fetcher: Callable[[], Awaitable[dict[str, Any]]] | None = None,
) -> dict[str, Any]:
current = now or datetime.now(UTC)
local_date = current.astimezone(SHANGHAI_TZ).date()
weather_task = _cached_source(
"weather", weather_fetcher or _default_weather_fetcher, current, WEATHER_TTL, WEATHER_STALE_TTL
)
gold_task = _cached_source(
"gold", gold_fetcher or _default_gold_fetcher, current, GOLD_TTL, GOLD_STALE_TTL
)
weather_result, gold_result = await asyncio.wait_for(
asyncio.gather(weather_task, gold_task), timeout=TOTAL_TIMEOUT_SECONDS
)
weather, weather_error = weather_result
gold, gold_error = gold_result
errors = {}
if weather_error:
errors["weather"] = weather_error
if gold_error:
errors["gold"] = gold_error
return {
"date": {
"solar_date": local_date.isoformat(),
"weekday": WEEKDAYS[local_date.weekday()],
"lunar": solar_to_lunar_text(local_date),
"timezone": "Asia/Shanghai",
},
"weather": weather,
"gold": gold,
"errors": errors,
}