from __future__ import annotations
|
|
import json
|
import socket
|
import threading
|
import time
|
import urllib.error
|
import urllib.parse
|
import urllib.request
|
from dataclasses import dataclass, field
|
from datetime import datetime
|
from datetime import date, timedelta
|
from pathlib import Path
|
from typing import Any, Iterable
|
|
from .cache import ContentCache, canonical_bytes, request_fingerprint, sha256_bytes
|
from .telemetry import aware_now, iso
|
|
|
class NetworkBudgetExceeded(TimeoutError):
|
pass
|
|
|
class SourceContractError(RuntimeError):
|
pass
|
|
|
QueryValue = str | int | float | bool | list[str] | tuple[str, ...]
|
|
|
@dataclass(frozen=True)
|
class HttpRequest:
|
provider_id: str
|
adapter_version: str
|
data_kind: str
|
method: str
|
url: str
|
ticker: str
|
as_of: str
|
fixture_id: str
|
query: dict[str, QueryValue] | list[tuple[str, str]] = field(default_factory=dict)
|
body: bytes | None = None
|
content_type: str = "application/json"
|
ttl_seconds: int = 21600
|
headers: dict[str, str] = field(default_factory=dict)
|
|
|
class NoRedirect(urllib.request.HTTPRedirectHandler):
|
def redirect_request(self, req, fp, code, msg, headers, newurl): # type: ignore[no-untyped-def]
|
raise SourceContractError(f"未登记重定向:{newurl}")
|
|
|
def _query_items(query: dict[str, QueryValue] | list[tuple[str, str]]) -> list[tuple[str, str]]:
|
if isinstance(query, list):
|
return [(str(key), str(value)) for key, value in query]
|
items: list[tuple[str, str]] = []
|
for key in sorted(query):
|
value = query[key]
|
if isinstance(value, (list, tuple)):
|
items.extend((str(key), str(item)) for item in value)
|
else:
|
items.append((str(key), str(value)))
|
return items
|
|
|
def _canonical_body(body: bytes | None, content_type: str) -> bytes:
|
if not body:
|
return b""
|
base = content_type.split(";", 1)[0].lower()
|
if base == "application/json":
|
try:
|
return canonical_bytes(json.loads(body.decode("utf-8")))
|
except (UnicodeDecodeError, json.JSONDecodeError):
|
return body
|
if base == "application/x-www-form-urlencoded":
|
pairs = urllib.parse.parse_qsl(body.decode("utf-8"), keep_blank_values=True)
|
return urllib.parse.urlencode(sorted(pairs), doseq=True).encode("utf-8")
|
return body
|
|
|
class HttpClient:
|
def __init__(
|
self,
|
registry: dict[str, Any],
|
cache: ContentCache,
|
deadline_mono: float,
|
run_id: str,
|
process_start: datetime,
|
fixture_dir: Path | None = None,
|
baseline_states: dict[str, str] | None = None,
|
baseline_entries: dict[str, dict[str, Any]] | None = None,
|
):
|
self.registry = registry
|
self.cache = cache
|
self.deadline_mono = deadline_mono
|
self.run_id = run_id
|
self.process_start = process_start
|
self.fixture_dir = fixture_dir
|
self.baseline_states = baseline_states
|
self.baseline_entries = baseline_entries or {}
|
self.opener = urllib.request.build_opener(NoRedirect)
|
self.fixture_manifest: dict[str, Any] | None = None
|
self.fixture_identity: str | None = None
|
self.cancel_event = threading.Event()
|
self.network_attempts = 0
|
self.cache_mutations = 0
|
self.telemetry: list[dict[str, Any]] = []
|
self._telemetry_lock = threading.Lock()
|
self._negative: dict[str, tuple[float, str]] = {}
|
if fixture_dir:
|
manifest_path = fixture_dir / "fixture_manifest.json"
|
raw = manifest_path.read_bytes()
|
self.fixture_manifest = json.loads(raw.decode("utf-8"))
|
self.fixture_identity = sha256_bytes(raw)
|
|
def remaining(self) -> float:
|
if self.cancel_event.is_set():
|
return 0.0
|
return max(0.0, self.deadline_mono - time.monotonic())
|
|
def cancel(self) -> None:
|
self.cancel_event.set()
|
|
def _contract(self, request: HttpRequest) -> dict[str, Any]:
|
provider = self.registry.get("providers", {}).get(request.provider_id)
|
if not provider or provider.get("adapter_version") != request.adapter_version:
|
raise SourceContractError(f"未登记 provider:{request.provider_id}")
|
contract = provider.get("contracts", {}).get(request.data_kind)
|
if not contract:
|
raise SourceContractError(f"未登记 data_kind:{request.provider_id}/{request.data_kind}")
|
return {**provider, **contract}
|
|
def _validate(self, request: HttpRequest) -> dict[str, Any]:
|
contract = self._contract(request)
|
parsed = urllib.parse.urlsplit(request.url)
|
if parsed.query or parsed.fragment or parsed.username or parsed.password:
|
raise SourceContractError(f"URL 必须只含登记 scheme/host/path:{request.url}")
|
if parsed.scheme != "https" or parsed.hostname != contract["host"]:
|
raise SourceContractError(f"未登记域名:{request.url}")
|
if parsed.path != contract["path"] or request.method != contract["method"]:
|
raise SourceContractError(f"未登记 endpoint/method:{request.method} {request.url}")
|
items = _query_items(request.query)
|
keys = [key for key, _ in items]
|
allowed = contract.get("query_keys", [])
|
if sorted(keys) != sorted(allowed):
|
raise SourceContractError(f"query 键不匹配:{keys}")
|
values = {key: value for key, value in items}
|
for key, expected in contract.get("fixed_query", {}).items():
|
if values.get(key) != str(expected):
|
raise SourceContractError(f"query 固定值不匹配:{key}")
|
code, market = request.ticker.split(".")
|
secid = f"{1 if market == 'SH' else 0}.{code}"
|
if "secid" in values and values["secid"] != secid:
|
raise SourceContractError("secid 与 ticker 不匹配")
|
if request.data_kind == "market_close":
|
expected_beg = (date.fromisoformat(request.as_of) - timedelta(days=14)).strftime("%Y%m%d")
|
if values.get("beg") != expected_beg or values.get("end") != request.as_of.replace("-", ""):
|
raise SourceContractError("K 线窗口不匹配")
|
if request.data_kind.startswith("finance_") and values.get("filter") != f'(SECUCODE="{request.ticker}")':
|
raise SourceContractError("finance filter 与 ticker 不匹配")
|
if request.data_kind == "forecast_summary" and values.get("filter") != f'(SECURITY_CODE="{code}")':
|
raise SourceContractError("forecast filter 与 ticker 不匹配")
|
if request.data_kind == "forecast_detail" and values.get("code") != f"{market}{code}":
|
raise SourceContractError("forecast detail code 与 ticker 不匹配")
|
if contract.get("body_keys") is not None:
|
if request.content_type.split(";", 1)[0].lower() != "application/x-www-form-urlencoded":
|
raise SourceContractError("登记 body 必须为 form")
|
try:
|
body_items = urllib.parse.parse_qsl(
|
(request.body or b"").decode("utf-8"), keep_blank_values=True
|
)
|
except UnicodeDecodeError as exc:
|
raise SourceContractError("form body 不是 UTF-8") from exc
|
body_values = {key: value for key, value in body_items}
|
if sorted(key for key, _ in body_items) != sorted(contract["body_keys"]):
|
raise SourceContractError("form body 键不匹配")
|
for key, expected in contract.get("fixed_body", {}).items():
|
if body_values.get(key) != str(expected):
|
raise SourceContractError(f"form body 固定值不匹配:{key}")
|
if body_values.get("stock", "").split(",", 1)[0] != code:
|
raise SourceContractError("CNInfo stock body 与 ticker 不匹配")
|
if body_values.get("seDate", "").split("~")[-1] != request.as_of:
|
raise SourceContractError("CNInfo seDate 与 as-of 不匹配")
|
if request.content_type.split(";", 1)[0].lower() not in contract.get(
|
"request_content_types", ["application/json"]
|
):
|
raise SourceContractError(f"未登记 request content-type:{request.content_type}")
|
return contract
|
|
def _headers(self, request: HttpRequest) -> dict[str, str]:
|
headers = {
|
"User-Agent": "project-info-stock-valuation-v2/2.0",
|
"Accept": "application/json,text/html;q=0.8",
|
**request.headers,
|
}
|
if request.body is not None:
|
headers["Content-Type"] = request.content_type
|
return headers
|
|
def _fingerprint(self, request: HttpRequest) -> str:
|
parsed = urllib.parse.urlsplit(request.url)
|
headers = sorted((key.lower().strip(), value.strip()) for key, value in self._headers(request).items())
|
preimage = {
|
"schema_version": 1,
|
"provider_id": request.provider_id,
|
"adapter_version": request.adapter_version,
|
"method": request.method,
|
"canonical_scheme_host_path": f"{parsed.scheme}://{parsed.hostname}{parsed.path}",
|
"sorted_query": _query_items(request.query),
|
"canonical_headers": headers,
|
"canonical_body_sha256": sha256_bytes(_canonical_body(request.body, request.content_type)),
|
"ticker": request.ticker,
|
"as_of": request.as_of,
|
"data_kind": request.data_kind,
|
"fixture_identity": (
|
f"{self.fixture_identity}:{request.fixture_id}" if self.fixture_identity else None
|
),
|
}
|
return request_fingerprint(preimage)
|
|
def _check_active(self, message: str = "network deadline exhausted") -> float:
|
remaining = self.remaining()
|
if remaining <= 0:
|
raise NetworkBudgetExceeded(message)
|
return remaining
|
|
def fetch(self, request: HttpRequest) -> dict[str, Any]:
|
contract = self._validate(request)
|
fingerprint = self._fingerprint(request)
|
negative = self._negative.get(fingerprint)
|
if negative and time.monotonic() <= negative[0]:
|
raise SourceContractError(f"本次运行负结果去重:{negative[1]}")
|
state = self.baseline_states.get(request.data_kind) if self.baseline_states else None
|
if state == "FRESH":
|
baseline_hit = self.cache.load_baseline_blob(
|
self.baseline_entries.get(request.data_kind, {}),
|
request.provider_id,
|
fingerprint,
|
)
|
if baseline_hit is None:
|
raise SourceContractError(
|
f"FRESH baseline blob/fingerprint invalid: {request.data_kind}"
|
)
|
return {
|
"body": baseline_hit["body"],
|
"meta": baseline_hit["meta"],
|
"from_cache": True,
|
"from_baseline": True,
|
"fingerprint": fingerprint,
|
}
|
allow_logical_reuse = (
|
self.baseline_states is None
|
or state == "FRESH"
|
)
|
hit = (
|
self.cache.load_reusable(request.provider_id, fingerprint, self.process_start)
|
if allow_logical_reuse
|
else None
|
)
|
if hit:
|
self._record_telemetry(request, fingerprint, hit["meta"], True)
|
return {
|
"body": hit["body"],
|
"meta": hit["meta"],
|
"from_cache": True,
|
"from_baseline": False,
|
"fingerprint": fingerprint,
|
}
|
remaining_before = self._check_active()
|
started = aware_now()
|
if self.fixture_dir:
|
assert self.fixture_manifest is not None
|
entry = self.fixture_manifest["requests"].get(request.fixture_id)
|
if entry is None:
|
raise SourceContractError(f"fixture request 缺失:{request.fixture_id}")
|
body = (self.fixture_dir / entry["file"]).read_bytes()
|
status = int(entry.get("status", 200))
|
headers = {key.lower(): value for key, value in entry.get(
|
"headers", {"content-type": request.content_type}
|
).items()}
|
attempts: list[dict[str, Any]] = []
|
else:
|
body, status, headers, attempts = self._fetch_network(request, contract)
|
self._check_active("network deadline exhausted before cache archive")
|
content_type = headers.get("content-type", "").split(";", 1)[0].strip().lower()
|
allowed_response = contract.get("response_content_types", [])
|
max_bytes = int(contract["max_response_bytes"])
|
ok = (
|
200 <= status < 300
|
and 0 < len(body) <= max_bytes
|
and content_type in allowed_response
|
)
|
fetched = aware_now()
|
raw_meta = {
|
"run_id": self.run_id,
|
"transport_complete": True,
|
"http_status": status,
|
"content_type": content_type,
|
"parse_ok": False,
|
"schema_ok": False,
|
"as_of_ok": False,
|
"semantic_status": "PENDING_PARSE" if ok else "GAP",
|
"adapter_version": request.adapter_version,
|
"method": request.method,
|
"canonical_url": request.url,
|
"query": _query_items(request.query),
|
"canonical_headers": sorted(
|
(key.lower(), value) for key, value in self._headers(request).items()
|
),
|
"body_sha256": sha256_bytes(_canonical_body(request.body, request.content_type)),
|
"ticker": request.ticker,
|
"as_of": request.as_of,
|
"data_kind": request.data_kind,
|
"fixture_identity": self.fixture_identity,
|
"started_at": iso(started),
|
"finished_at": iso(fetched),
|
"remaining_before": remaining_before,
|
"remaining_after": self.remaining(),
|
"attempts": attempts,
|
}
|
self._check_active("network deadline exhausted before raw archive")
|
meta = self.cache.store(
|
request.provider_id,
|
fingerprint,
|
body,
|
fetched,
|
request.ttl_seconds if ok else 60,
|
reusable=False,
|
raw_meta=raw_meta,
|
)
|
self.cache_mutations += 1
|
self._record_telemetry(request, fingerprint, meta, False)
|
if not ok:
|
self._negative[fingerprint] = (
|
time.monotonic() + min(60.0, self.remaining()),
|
f"status={status} bytes={len(body)} content-type={content_type}",
|
)
|
raise SourceContractError(
|
f"HTTP/响应合同失败:status={status} bytes={len(body)} content-type={content_type}"
|
)
|
return {
|
"body": body,
|
"meta": meta,
|
"from_cache": False,
|
"from_baseline": False,
|
"fingerprint": fingerprint,
|
}
|
|
def _record_telemetry(
|
self,
|
request: HttpRequest,
|
fingerprint: str,
|
meta: dict[str, Any],
|
from_cache: bool,
|
) -> None:
|
item = {
|
"provider_id": request.provider_id,
|
"data_kind": request.data_kind,
|
"fingerprint": fingerprint,
|
"raw_hash": meta.get("blob_hash"),
|
"bytes": meta.get("bytes"),
|
"http_status": meta.get("http_status"),
|
"from_cache": from_cache,
|
"started_at": meta.get("started_at"),
|
"finished_at": meta.get("finished_at"),
|
"remaining_before": meta.get("remaining_before"),
|
"remaining_after": meta.get("remaining_after"),
|
"attempts": meta.get("attempts", []),
|
}
|
with self._telemetry_lock:
|
self.telemetry.append(item)
|
|
def confirm_reusable(self, request: HttpRequest, response: dict[str, Any]) -> None:
|
if response["from_cache"]:
|
return
|
self._check_active("network deadline exhausted before reusable promotion")
|
response["meta"] = self.cache.promote_reusable(
|
request.provider_id, response["fingerprint"], response["meta"]
|
)
|
self.cache_mutations += 1
|
|
@staticmethod
|
def _set_response_timeout(response: Any, timeout: float) -> None:
|
candidates: Iterable[Any] = (
|
getattr(response, "fp", None),
|
getattr(getattr(response, "fp", None), "raw", None),
|
getattr(getattr(getattr(response, "fp", None), "raw", None), "_sock", None),
|
)
|
for candidate in candidates:
|
setter = getattr(candidate, "settimeout", None)
|
if setter:
|
try:
|
setter(timeout)
|
return
|
except OSError:
|
continue
|
|
def _read_chunk_bounded(self, response: Any, size: int) -> bytes:
|
"""Bound even transports that ignore socket timeouts (notably test/error bodies)."""
|
remaining = self._check_active("network deadline exhausted during read")
|
self._set_response_timeout(response, remaining)
|
completed = threading.Event()
|
outcome: dict[str, Any] = {}
|
|
def read_once() -> None:
|
try:
|
outcome["chunk"] = response.read(size)
|
except BaseException as exc: # propagated in the caller thread
|
outcome["error"] = exc
|
finally:
|
completed.set()
|
|
worker = threading.Thread(
|
target=read_once,
|
name="valuation-v2-bounded-read",
|
daemon=True,
|
)
|
worker.start()
|
if not completed.wait(timeout=remaining):
|
closer = getattr(response, "close", None)
|
if closer:
|
try:
|
closer()
|
except OSError:
|
pass
|
raise NetworkBudgetExceeded("network deadline exhausted during read")
|
error = outcome.get("error")
|
if error is not None:
|
raise error
|
return bytes(outcome.get("chunk", b""))
|
|
def _read_response_bounded(self, response: Any, max_bytes: int) -> bytes:
|
chunks: list[bytes] = []
|
total = 0
|
while True:
|
chunk = self._read_chunk_bounded(
|
response, min(65536, max_bytes - total + 1)
|
)
|
if not chunk:
|
break
|
chunks.append(chunk)
|
total += len(chunk)
|
if total > max_bytes:
|
raise SourceContractError(
|
f"响应超过登记上限 {max_bytes} bytes"
|
)
|
return b"".join(chunks)
|
|
def _fetch_network(
|
self, request: HttpRequest, contract: dict[str, Any]
|
) -> tuple[bytes, int, dict[str, str], list[dict[str, Any]]]:
|
query = urllib.parse.urlencode(_query_items(request.query), doseq=True)
|
url = request.url + ("?" + query if query else "")
|
headers = self._headers(request)
|
max_bytes = int(contract["max_response_bytes"])
|
last: Exception | None = None
|
attempts: list[dict[str, Any]] = []
|
for attempt, backoff in enumerate((0.0, 0.25, 0.75), start=1):
|
before = self._check_active()
|
if backoff:
|
time.sleep(min(backoff, before))
|
before = self._check_active("network deadline exhausted after backoff")
|
timeout = min(20.0, before)
|
if timeout <= 0:
|
raise NetworkBudgetExceeded("network deadline exhausted before open")
|
self.network_attempts += 1
|
req = urllib.request.Request(
|
url, data=request.body, headers=headers, method=request.method
|
)
|
attempt_started = aware_now()
|
try:
|
with self.opener.open(req, timeout=timeout) as response:
|
body = self._read_response_bounded(response, max_bytes)
|
total = len(body)
|
attempts.append(
|
{
|
"attempt": attempt,
|
"started_at": iso(attempt_started),
|
"finished_at": iso(aware_now()),
|
"remaining_before": before,
|
"remaining_after": self.remaining(),
|
"http_status": response.status,
|
"bytes": total,
|
"deadline_cause": None,
|
}
|
)
|
return body, response.status, {
|
key.lower(): value for key, value in response.headers.items()
|
}, attempts
|
except urllib.error.HTTPError as exc:
|
try:
|
body = self._read_response_bounded(exc, max_bytes)
|
finally:
|
try:
|
exc.close()
|
except OSError:
|
pass
|
attempts.append(
|
{
|
"attempt": attempt,
|
"started_at": iso(attempt_started),
|
"finished_at": iso(aware_now()),
|
"remaining_before": before,
|
"remaining_after": self.remaining(),
|
"http_status": exc.code,
|
"bytes": len(body),
|
"deadline_cause": None,
|
}
|
)
|
if 500 <= exc.code < 600 and attempt < 3:
|
last = exc
|
continue
|
return body, exc.code, {
|
key.lower(): value for key, value in exc.headers.items()
|
}, attempts
|
except NetworkBudgetExceeded:
|
raise
|
except (urllib.error.URLError, socket.timeout, TimeoutError, OSError) as exc:
|
last = exc
|
attempts.append(
|
{
|
"attempt": attempt,
|
"started_at": iso(attempt_started),
|
"finished_at": iso(aware_now()),
|
"remaining_before": before,
|
"remaining_after": self.remaining(),
|
"http_status": None,
|
"bytes": 0,
|
"deadline_cause": type(exc).__name__,
|
}
|
)
|
if attempt < 3:
|
continue
|
break
|
raise SourceContractError(f"网络请求失败:{last}")
|