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}")