"""客户端与服务端共享 wire 的严格值校验。""" from __future__ import annotations import calendar from datetime import datetime, timezone import json import re from typing import Any, Iterable, Mapping from urllib.parse import quote from .errors import ValidationError UUID4_RE = re.compile( r"[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}" ) LOWER_HEX_64_RE = re.compile(r"[0-9a-f]{64}") RFC3339_Z_RE = re.compile( r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d{1,9})?Z" ) MONEY_RE = re.compile(r"(?:0|[1-9][0-9]*)\.[0-9]{2}") GOODS_ID_RE = re.compile(r"[0-9]+") MAX_TITLE_CODE_POINTS = 120 MAX_SKU_TEXT_CODE_POINTS = 80 MAX_GOODS_ID_ASCII_CHARACTERS = 32 MAX_MONEY_ASCII_CHARACTERS = 32 # Go strings.TrimSpace uses Unicode White_Space plus the six ASCII space # characters below, but unlike Python str.strip it does not include U+001C-- # U+001F. Keep the wire contract independent of either runtime's defaults. GO_UNICODE_WHITE_SPACE = "\t\n\v\f\r \u0085\u00a0\u1680\u2000\u2001\u2002\u2003\u2004\u2005\u2006\u2007\u2008\u2009\u200a\u2028\u2029\u202f\u205f\u3000" def require_string(value: object, reason: str, *, maximum: int = 4096) -> str: if not isinstance(value, str) or not value or len(value) > maximum: raise ValidationError(reason) if any(0xD800 <= ord(character) <= 0xDFFF for character in value): raise ValidationError(reason) return value def require_persisted_text(value: object, reason: str, *, maximum: int) -> str: """Validate text stored by Go after TrimSpace, without Python trim drift.""" text = require_string(value, reason, maximum=maximum) if text.strip(GO_UNICODE_WHITE_SPACE) != text: raise ValidationError(reason) # Python str.strip treats these C0 separators as whitespace while Go does # not. Reject them anywhere on both ends instead of assigning them two # runtime-dependent meanings. if any(0x1C <= ord(character) <= 0x1F for character in text): raise ValidationError(reason) return text def require_uuid4(value: object, reason: str = "invalid_uuid") -> str: text = require_string(value, reason, maximum=36) if UUID4_RE.fullmatch(text) is None: raise ValidationError(reason) return text def require_lower_hex_64(value: object, reason: str = "invalid_hex") -> str: text = require_string(value, reason, maximum=64) if LOWER_HEX_64_RE.fullmatch(text) is None: raise ValidationError(reason) return text def require_rfc3339_z(value: object, reason: str = "invalid_timestamp") -> str: text = require_string(value, reason, maximum=40) if RFC3339_Z_RE.fullmatch(text) is None: raise ValidationError(reason) parsed: datetime | None = None try: parsed = datetime.fromisoformat(text[:-1] + "+00:00") except ValueError: pass if parsed is None: raise ValidationError(reason) if parsed.utcoffset() is None or parsed.utcoffset().total_seconds() != 0: raise ValidationError(reason) return text def rfc3339_z_nanoseconds(value: object, reason: str = "invalid_timestamp") -> int: """无浮点、无微秒截断地把 UTC RFC3339Nano 转成纳秒时间轴。""" text = require_rfc3339_z(value, reason) base: datetime | None = None try: base = datetime.strptime(text[:19], "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc) except ValueError: pass if base is None: raise ValidationError(reason) fraction = "" if len(text) == 20 else text[20:-1] nanoseconds = int(fraction.ljust(9, "0")) if fraction else 0 return calendar.timegm(base.utctimetuple()) * 1_000_000_000 + nanoseconds def datetime_nanoseconds(value: datetime, reason: str = "invalid_timestamp") -> int: if not isinstance(value, datetime) or value.utcoffset() is None: raise ValidationError(reason) utc = value.astimezone(timezone.utc) return calendar.timegm(utc.utctimetuple()) * 1_000_000_000 + utc.microsecond * 1_000 def require_positive_int(value: object, reason: str = "invalid_integer") -> int: # bool 是 int 的子类;wire 中必须显式拒绝 true/false。 if type(value) is not int or value <= 0 or value > 9_223_372_036_854_775_807: raise ValidationError(reason) return value def require_money(value: object, reason: str = "invalid_money") -> str: text = require_string(value, reason, maximum=MAX_MONEY_ASCII_CHARACTERS) if MONEY_RE.fullmatch(text) is None or text == "0.00": raise ValidationError(reason) return text def require_goods_id(value: object) -> str: text = require_string(value, "invalid_goods_id", maximum=MAX_GOODS_ID_ASCII_CHARACTERS) if GOODS_ID_RE.fullmatch(text) is None: raise ValidationError("invalid_goods_id") return text def canonical_product_url(goods_id: str) -> str: require_goods_id(goods_id) return "https://mobile.yangkeduo.com/goods.html?goods_id=" + quote(goods_id, safe="") def require_exact_fields( value: object, required: Iterable[str], reason: str = "invalid_schema", ) -> Mapping[str, Any]: if not isinstance(value, dict): raise ValidationError(reason) expected = frozenset(required) if frozenset(value) != expected: raise ValidationError(reason) return value def strict_json_loads(raw: bytes, *, maximum: int) -> object: if not isinstance(raw, bytes) or len(raw) == 0 or len(raw) > maximum: raise ValidationError("invalid_json_size") text: str | None = None try: text = raw.decode("utf-8") except UnicodeDecodeError: pass if text is None: raise ValidationError("invalid_json_utf8") if text.startswith("\ufeff"): raise ValidationError("invalid_json_bom") def pairs_hook(pairs: list[tuple[str, Any]]) -> dict[str, Any]: result: dict[str, Any] = {} for key, value in pairs: if key in result: raise ValidationError("duplicate_json_key") result[key] = value return result def reject_number(_: str) -> object: raise ValidationError("invalid_json_number") def parse_integer(value: str) -> int: digits = value[1:] if value.startswith("-") else value if len(digits) > 19: raise ValidationError("invalid_json_integer") parsed = int(value) if parsed < -9_223_372_036_854_775_808 or parsed > 9_223_372_036_854_775_807: raise ValidationError("invalid_json_integer") return parsed parsed_json: object | None = None failed = False try: parsed_json = json.loads( text, object_pairs_hook=pairs_hook, parse_int=parse_integer, parse_float=reject_number, parse_constant=reject_number, ) except ValidationError: raise except (json.JSONDecodeError, UnicodeError, ValueError, RecursionError): failed = True if failed: raise ValidationError("invalid_json") return parsed_json