diff --git a/core/huya/cookie_utils.py b/core/huya/cookie_utils.py index 832a176..4f2c917 100644 --- a/core/huya/cookie_utils.py +++ b/core/huya/cookie_utils.py @@ -5,6 +5,7 @@ from __future__ import annotations from collections.abc import Iterable, Mapping import requests +from requests.cookies import RequestsCookieJar def cookie_pairs(cookie: str) -> list[tuple[str, str]]: @@ -36,16 +37,25 @@ def normalize_cookie_pairs(pairs: Iterable[tuple[str, str]]) -> str: return "; ".join(f"{key}={values[key]}" for key in ordered_keys) -def normalize_huya_cookie(cookie: str | Mapping[str, str] | requests.cookies.RequestsCookieJar | None) -> str: +def normalize_huya_cookie( + cookie: str | Mapping[str, str] | RequestsCookieJar | None, +) -> str: """把 Cookie 字符串、dict 或 CookieJar 转成去重后的浏览器 Cookie 字符串。""" if not cookie: return "" if isinstance(cookie, str): return normalize_cookie_pairs(cookie_pairs(cookie)) - return normalize_cookie_pairs(cookie.items()) + return normalize_cookie_pairs( + [ + (str(key), "" if value is None else str(value)) + for key, value in cookie.items() + ] + ) -def cookie_value(cookie: str | Mapping[str, str] | requests.cookies.RequestsCookieJar | None, key: str) -> str: +def cookie_value( + cookie: str | Mapping[str, str] | RequestsCookieJar | None, key: str +) -> str: """从 Cookie 中读取 key;有重复时以最后一次出现为准。""" if not cookie or not key: return "" diff --git a/core/huya/login.py b/core/huya/login.py index ebdbeb6..63d47c7 100644 --- a/core/huya/login.py +++ b/core/huya/login.py @@ -14,6 +14,7 @@ from http.cookies import SimpleCookie from urllib.parse import quote, urlsplit, urlunsplit import requests +from requests.cookies import RequestsCookieJar from loguru import logger from .cookie_utils import normalize_huya_cookie @@ -75,7 +76,7 @@ def password_sha1(password: str) -> str: return hashlib.sha1(password.encode("utf-8")).hexdigest() -def cookie_string(cookies: requests.cookies.RequestsCookieJar | Mapping[str, str]) -> str: +def cookie_string(cookies: RequestsCookieJar | Mapping[str, str]) -> str: """把 CookieJar/dict 转成浏览器 Cookie 字符串。""" return normalize_huya_cookie(cookies) @@ -116,15 +117,19 @@ def encode_behavior(page: str = HUYA_PAGE_URL) -> str: actions.append({"id": action_id, "d": elapsed, "time": now}) now += random.randint(120, 400) elapsed += random.randint(120, 400) - actions.append({ - "id": "11", - "x": random.randint(430, 560), - "y": random.randint(250, 330), - "d": elapsed, - "time": now, - }) + actions.append( + { + "id": "11", + "x": random.randint(430, 560), + "y": random.randint(250, 330), + "d": elapsed, + "time": now, + } + ) value = {"furl": page, "curl": page, "user_action": actions} - return quote(json.dumps(value, separators=(",", ":"), ensure_ascii=False), safe="~()*!.'") + return quote( + json.dumps(value, separators=(",", ":"), ensure_ascii=False), safe="~()*!.'" + ) class HuyaPasswordLogin: @@ -177,23 +182,25 @@ class HuyaPasswordLogin: self._setup_headers() def _setup_headers(self) -> None: - self.session.headers.update({ - "Accept": "*/*", - "Accept-Language": "zh-CN,zh;q=0.9", - "Cache-Control": "no-cache", - "Connection": "keep-alive", - "Content-Type": "application/json;charset=UTF-8", - "Origin": "https://aq.huya.com", - "Pragma": "no-cache", - "Referer": "https://aq.huya.com/", - "Sec-Fetch-Dest": "empty", - "Sec-Fetch-Mode": "cors", - "Sec-Fetch-Site": "same-site", - "User-Agent": self.ua, - "sec-ch-ua": '"Not;A=Brand";v="8", "Chromium";v="150", "Google Chrome";v="150"', - "sec-ch-ua-mobile": "?0", - "sec-ch-ua-platform": '"macOS"', - }) + self.session.headers.update( + { + "Accept": "*/*", + "Accept-Language": "zh-CN,zh;q=0.9", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "Content-Type": "application/json;charset=UTF-8", + "Origin": "https://aq.huya.com", + "Pragma": "no-cache", + "Referer": "https://aq.huya.com/", + "Sec-Fetch-Dest": "empty", + "Sec-Fetch-Mode": "cors", + "Sec-Fetch-Site": "same-site", + "User-Agent": self.ua, + "sec-ch-ua": '"Not;A=Brand";v="8", "Chromium";v="150", "Google Chrome";v="150"', + "sec-ch-ua-mobile": "?0", + "sec-ch-ua-platform": '"macOS"', + } + ) @staticmethod def _safe_url(url: str) -> str: @@ -202,7 +209,9 @@ class HuyaPasswordLogin: def _request_json(self, method: str, url: str, source: str, **kwargs) -> dict: response = self.session.request(method, url, timeout=self.timeout, **kwargs) - logger.debug(f"{method.upper()} {self._safe_url(url)} -> {response.status_code}") + logger.debug( + f"{method.upper()} {self._safe_url(url)} -> {response.status_code}" + ) response.raise_for_status() try: return response.json() @@ -219,10 +228,12 @@ class HuyaPasswordLogin: logger.debug(f"GET {self._safe_url(url)} -> {response.status_code}") except requests.RequestException as exc: logger.debug(f"虎牙UDB middle初始化失败: {host}: {exc}") - self.session.headers.update({ - "Origin": "https://udblgn.huya.com", - "Referer": self.middle_url, - }) + self.session.headers.update( + { + "Origin": "https://udblgn.huya.com", + "Referer": self.middle_url, + } + ) def prepare_device(self) -> str: """获取虎牙风控 sdid(优先 hydevice 高信任指纹,失败降级旧版)。""" @@ -230,8 +241,11 @@ class HuyaPasswordLogin: try: # 一号一设备: 每账号独立 hydevice 状态, sdid/40hex hdid 不跨账号同源 - result = get_huya_sdid(app_id="5008", timeout=self.timeout, - state_dir=account_state_dir(self.username)) + result = get_huya_sdid( + app_id="5008", + timeout=self.timeout, + state_dir=account_state_dir(self.username), + ) if result.sdid: self.sdid = result.sdid self.sdid_source = result.source @@ -242,12 +256,16 @@ class HuyaPasswordLogin: logger.warning("虎牙高信任指纹异常: {},降级旧流程", exc) token_payload = {"encryptVersion": "1.0.1", "fingerprintVersion": "1.2.41"} - token_res = self._request_json("post", DF_TOKEN_URL, "获取虎牙 df token", json=token_payload) + token_res = self._request_json( + "post", DF_TOKEN_URL, "获取虎牙 df token", json=token_payload + ) token = token_res.get("data", {}).get("token") if not token: raise HuyaLoginError(f"获取虎牙 df token 失败: {token_res}") - collect_res = self._request_json("post", DF_COLLECT_URL, "获取虎牙 sdid", json={"token": token}) + collect_res = self._request_json( + "post", DF_COLLECT_URL, "获取虎牙 sdid", json={"token": token} + ) self.sdid = collect_res.get("data", {}).get("sdid", "") if not self.sdid: raise HuyaLoginError(f"获取虎牙 sdid 失败: {collect_res}") @@ -303,7 +321,11 @@ class HuyaPasswordLogin: from .verification import HuyaNoSessionError, HuyaVerificationSolver solver = HuyaVerificationSolver( - cookie=self.session.cookies.get_dict(), + cookie={ + key: value + for key, value in self.session.cookies.get_dict().items() + if value is not None + }, ua=self.ua, sdid=self.sdid, session=self.session, @@ -323,7 +345,10 @@ class HuyaPasswordLogin: @staticmethod def _is_credential_error(payload: dict) -> bool: text = json.dumps(payload, ensure_ascii=False) - return any(word in text for word in ("密码错误", "账号或密码", "账号不存在", "用户不存在")) + return any( + word in text + for word in ("密码错误", "账号或密码", "账号不存在", "用户不存在") + ) @staticmethod def _payload_data(payload: dict) -> dict: @@ -407,13 +432,17 @@ class HuyaPasswordLogin: logger.info(f"虎牙密码登录触发风控: {return_code}") auth_id = self._solve_verification(payload) for index in range(3): - payload = self._password_login_once(auth_id=auth_id, session_data=session_data) + payload = self._password_login_once( + auth_id=auth_id, session_data=session_data + ) return_code = int(payload.get("returnCode") or 0) if return_code == 0: return self._success(payload) if self._is_credential_error(payload): raise HuyaCredentialError(f"虎牙账号或密码错误: {payload}") - session_data = str(self._payload_data(payload).get("sessionData") or session_data) + session_data = str( + self._payload_data(payload).get("sessionData") or session_data + ) if return_code in (10030, 10039): logger.info(f"虎牙密码登录二次风控: {return_code} ({index + 1}/3)") auth_id = self._solve_verification(payload) diff --git a/core/huya/sms_login.py b/core/huya/sms_login.py index 0e6d44c..c9e5648 100644 --- a/core/huya/sms_login.py +++ b/core/huya/sms_login.py @@ -36,108 +36,110 @@ SMS_CODE_URL = "https://udblgn.huya.com/web/v2/smsCode" SMS_LOGIN_URL = "https://udblgn.huya.com/web/v2/smsLogin" DF_TOKEN_URL = "https://df.huya.com/web/df/token" DF_COLLECT_URL = "https://df.huya.com/web/df/collect" -HUYA_COUNTRY_CALLING_CODES = frozenset({ - "1", - "7", - "20", - "27", - "30", - "31", - "32", - "33", - "34", - "36", - "39", - "40", - "41", - "43", - "44", - "45", - "46", - "47", - "48", - "49", - "52", - "55", - "60", - "61", - "62", - "63", - "64", - "65", - "66", - "81", - "82", - "84", - "86", - "90", - "91", - "92", - "93", - "94", - "95", - "98", - "212", - "213", - "216", - "218", - "234", - "254", - "351", - "352", - "353", - "354", - "355", - "356", - "357", - "358", - "359", - "370", - "371", - "372", - "373", - "374", - "375", - "376", - "377", - "378", - "380", - "381", - "382", - "385", - "386", - "387", - "389", - "420", - "421", - "852", - "853", - "855", - "856", - "886", - "960", - "961", - "962", - "963", - "964", - "965", - "966", - "967", - "968", - "971", - "972", - "973", - "974", - "975", - "976", - "977", - "992", - "993", - "994", - "995", - "996", - "998", -}) +HUYA_COUNTRY_CALLING_CODES = frozenset( + { + "1", + "7", + "20", + "27", + "30", + "31", + "32", + "33", + "34", + "36", + "39", + "40", + "41", + "43", + "44", + "45", + "46", + "47", + "48", + "49", + "52", + "55", + "60", + "61", + "62", + "63", + "64", + "65", + "66", + "81", + "82", + "84", + "86", + "90", + "91", + "92", + "93", + "94", + "95", + "98", + "212", + "213", + "216", + "218", + "234", + "254", + "351", + "352", + "353", + "354", + "355", + "356", + "357", + "358", + "359", + "370", + "371", + "372", + "373", + "374", + "375", + "376", + "377", + "378", + "380", + "381", + "382", + "385", + "386", + "387", + "389", + "420", + "421", + "852", + "853", + "855", + "856", + "886", + "960", + "961", + "962", + "963", + "964", + "965", + "966", + "967", + "968", + "971", + "972", + "973", + "974", + "975", + "976", + "977", + "992", + "993", + "994", + "995", + "996", + "998", + } +) @dataclass @@ -154,7 +156,9 @@ class HuyaSmsCodeResult: request_id: str = "" -def _split_country_calling_code(digits: str, *, allow_single_digit: bool = True) -> tuple[str, str]: +def _split_country_calling_code( + digits: str, *, allow_single_digit: bool = True +) -> tuple[str, str]: """从国际号码中拆出国家/地区区号。 allow_single_digit=False 时不匹配 1/7 等单位区号,避免把国内 1 开头手机号误判成美国号。 @@ -238,7 +242,9 @@ def normalize_huya_phone(phone: str) -> str: # 国内号 if digits.startswith("86") and len(digits) == 13: return f"0{digits}" - if len(digits) == 11 and digits.startswith(("13", "14", "15", "16", "17", "18", "19")): + if len(digits) == 11 and digits.startswith( + ("13", "14", "15", "16", "17", "18", "19") + ): return f"086{digits}" if len(digits) == 11: return f"086{digits}" @@ -251,28 +257,34 @@ def encode_sms_behavior(page: str = HUYA_PAGE_URL) -> str: """生成短信登录抓包同款行为轨迹。""" now = int(time.time() * 1000) - random.randint(8000, 18000) elapsed = random.randint(1200, 2400) - actions = [{ - "id": "12", - "x": random.randint(520, 620), - "y": random.randint(45, 75), - "d": elapsed, - "time": now, - }] + actions = [ + { + "id": "12", + "x": random.randint(520, 620), + "y": random.randint(45, 75), + "d": elapsed, + "time": now, + } + ] for action_id in ("17", "17"): now += random.randint(1800, 4200) elapsed += random.randint(1800, 4200) actions.append({"id": action_id, "d": elapsed, "time": now}) now += random.randint(80, 240) elapsed += random.randint(80, 240) - actions.append({ - "id": "18", - "x": random.randint(610, 660), - "y": random.randint(165, 195), - "d": elapsed, - "time": now, - }) + actions.append( + { + "id": "18", + "x": random.randint(610, 660), + "y": random.randint(165, 195), + "d": elapsed, + "time": now, + } + ) value = {"furl": page, "curl": page, "user_action": actions} - return quote(json.dumps(value, separators=(",", ":"), ensure_ascii=False), safe="~()*!.'") + return quote( + json.dumps(value, separators=(",", ":"), ensure_ascii=False), safe="~()*!.'" + ) class HuyaSmsLogin: @@ -298,6 +310,7 @@ class HuyaSmsLogin: self.context = generate_context(self.device_id) self.middle_url = f"https://udblgn.huya.com/web/middle/{APP_VERSION}/{self.exchange}/https/{self.device_id}" self.sdid = "" + self.session_data = "" self.session = session or requests.Session() self.session.trust_env = False @@ -308,23 +321,25 @@ class HuyaSmsLogin: self._setup_headers() def _setup_headers(self) -> None: - self.session.headers.update({ - "Accept": "*/*", - "Accept-Language": "zh-CN,zh;q=0.9", - "Cache-Control": "no-cache", - "Connection": "keep-alive", - "Content-Type": "application/json;charset=UTF-8", - "Origin": "https://aq.huya.com", - "Pragma": "no-cache", - "Referer": "https://aq.huya.com/", - "Sec-Fetch-Dest": "empty", - "Sec-Fetch-Mode": "cors", - "Sec-Fetch-Site": "same-site", - "User-Agent": self.ua, - "sec-ch-ua": '"Not;A=Brand";v="8", "Chromium";v="150", "Google Chrome";v="150"', - "sec-ch-ua-mobile": "?0", - "sec-ch-ua-platform": '"macOS"', - }) + self.session.headers.update( + { + "Accept": "*/*", + "Accept-Language": "zh-CN,zh;q=0.9", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "Content-Type": "application/json;charset=UTF-8", + "Origin": "https://aq.huya.com", + "Pragma": "no-cache", + "Referer": "https://aq.huya.com/", + "Sec-Fetch-Dest": "empty", + "Sec-Fetch-Mode": "cors", + "Sec-Fetch-Site": "same-site", + "User-Agent": self.ua, + "sec-ch-ua": '"Not;A=Brand";v="8", "Chromium";v="150", "Google Chrome";v="150"', + "sec-ch-ua-mobile": "?0", + "sec-ch-ua-platform": '"macOS"', + } + ) @staticmethod def _safe_url(url: str) -> str: @@ -339,7 +354,9 @@ class HuyaSmsLogin: def _request_json(self, method: str, url: str, source: str, **kwargs) -> dict: response = self.session.request(method, url, timeout=self.timeout, **kwargs) - logger.debug(f"{method.upper()} {self._safe_url(url)} -> {response.status_code}") + logger.debug( + f"{method.upper()} {self._safe_url(url)} -> {response.status_code}" + ) response.raise_for_status() try: return response.json() @@ -356,20 +373,26 @@ class HuyaSmsLogin: logger.debug(f"GET {self._safe_url(url)} -> {response.status_code}") except requests.RequestException as exc: logger.debug(f"虎牙UDB middle初始化失败: {host}: {exc}") - self.session.headers.update({ - "Origin": "https://udblgn.huya.com", - "Referer": self.middle_url, - }) + self.session.headers.update( + { + "Origin": "https://udblgn.huya.com", + "Referer": self.middle_url, + } + ) def prepare_device(self) -> str: """获取虎牙风控 sdid。""" token_payload = {"encryptVersion": "1.0.1", "fingerprintVersion": "1.2.41"} - token_res = self._request_json("post", DF_TOKEN_URL, "获取虎牙 df token", json=token_payload) + token_res = self._request_json( + "post", DF_TOKEN_URL, "获取虎牙 df token", json=token_payload + ) token = token_res.get("data", {}).get("token") if not token: raise HuyaLoginError(f"获取虎牙 df token 失败: {token_res}") - collect_res = self._request_json("post", DF_COLLECT_URL, "获取虎牙 sdid", json={"token": token}) + collect_res = self._request_json( + "post", DF_COLLECT_URL, "获取虎牙 sdid", json={"token": token} + ) self.sdid = collect_res.get("data", {}).get("sdid", "") if not self.sdid: raise HuyaLoginError(f"获取虎牙 sdid 失败: {collect_res}") @@ -380,7 +403,11 @@ class HuyaSmsLogin: from .verification import HuyaVerificationSolver solver = HuyaVerificationSolver( - cookie=self.session.cookies.get_dict(), + cookie={ + key: value + for key, value in self.session.cookies.get_dict().items() + if value is not None + }, ua=self.ua, sdid=self.sdid, session=self.session, @@ -511,7 +538,9 @@ class HuyaSmsLogin: return_code = int(payload.get("returnCode") or 0) if return_code == 0: cookie = cookie_string(self.session.cookies) - logger.success(f"虎牙短信登录成功: {self.phone}, Cookie长度: {len(cookie)}") + logger.success( + f"虎牙短信登录成功: {self.phone}, Cookie长度: {len(cookie)}" + ) return HuyaLoginResult( success=True, cookie=cookie, @@ -563,7 +592,9 @@ class HuyaSmsLogin: "middleUrl": self.middle_url, "createdAt": int(time.time()), } - raw = json.dumps(data, separators=(",", ":"), ensure_ascii=False).encode("utf-8") + raw = json.dumps(data, separators=(",", ":"), ensure_ascii=False).encode( + "utf-8" + ) return base64.urlsafe_b64encode(raw).decode("ascii") @classmethod @@ -588,7 +619,11 @@ class HuyaSmsLogin: login = cls( phone=expected_phone, - cookie=data.get("cookies") or {}, + cookie={ + key: value + for key, value in (data.get("cookies") or {}).items() + if value is not None + }, ua=str(data.get("ua") or DEFAULT_UA), proxies=proxies, timeout=timeout, @@ -600,10 +635,12 @@ class HuyaSmsLogin: login.device_id = str(data.get("deviceId") or login.device_id) login.middle_url = str(data.get("middleUrl") or login.middle_url) login.session_data = str(data.get("sessionData") or "") - login.session.headers.update({ - "Origin": "https://udblgn.huya.com", - "Referer": login.middle_url, - }) + login.session.headers.update( + { + "Origin": "https://udblgn.huya.com", + "Referer": login.middle_url, + } + ) return login diff --git a/core/huya/verification/ocr.py b/core/huya/verification/ocr.py index 22f7718..088e8f3 100644 --- a/core/huya/verification/ocr.py +++ b/core/huya/verification/ocr.py @@ -5,6 +5,7 @@ from __future__ import annotations from functools import lru_cache from io import BytesIO from pathlib import Path +from typing import cast import cv2 import numpy as np @@ -19,7 +20,12 @@ MODEL_DIR = Path(__file__).resolve().parent / "models" class YoloOnnx: """轻量 YOLO ONNX 推理封装。""" - def __init__(self, model_path: str | Path, classes: list[str], providers: list[str] | None = None): + def __init__( + self, + model_path: str | Path, + classes: list[str], + providers: list[str] | None = None, + ): providers = providers or ["CPUExecutionProvider"] ort.set_default_logger_severity(3) self.session = ort.InferenceSession(str(model_path), providers=providers) @@ -68,7 +74,9 @@ class YoloOnnx: return np.expand_dims(input_tensor, axis=0) @staticmethod - def nms(boxes: np.ndarray, scores: np.ndarray, iou_threshold: float = 0.3) -> list[int]: + def nms( + boxes: np.ndarray, scores: np.ndarray, iou_threshold: float = 0.3 + ) -> list[int]: order = scores.argsort()[::-1] keep = [] @@ -84,7 +92,9 @@ class YoloOnnx: width = np.maximum(0.0, xx2 - xx1) height = np.maximum(0.0, yy2 - yy1) intersection = width * height - area_i = (boxes[index, 2] - boxes[index, 0]) * (boxes[index, 3] - boxes[index, 1]) + area_i = (boxes[index, 2] - boxes[index, 0]) * ( + boxes[index, 3] - boxes[index, 1] + ) area_j = (boxes[order[1:], 2] - boxes[order[1:], 0]) * ( boxes[order[1:], 3] - boxes[order[1:], 1] ) @@ -134,20 +144,25 @@ class YoloOnnx: for index in keep_indices: class_id = int(class_ids[index]) x1, y1, x2, y2 = map(int, xyxy_boxes[index].tolist()) - results.append({ - "label_id": class_id, - "label_name": self.names[class_id], - "confidence": float(scores[index]), - "box_mid_xy": [(x1 + x2) // 2, (y1 + y2) // 2], - "xyxy": [x1, y1, x2, y2], - }) + results.append( + { + "label_id": class_id, + "label_name": self.names[class_id], + "confidence": float(scores[index]), + "box_mid_xy": [(x1 + x2) // 2, (y1 + y2) // 2], + "xyxy": [x1, y1, x2, y2], + } + ) results.sort(key=lambda item: item["confidence"], reverse=True) return results def detect(self, image: Image.Image) -> list[dict]: input_tensor = self.preprocess(image) - outputs = self.session.run([self.output_name], {self.input_name: input_tensor}) + outputs = cast( + list[np.ndarray], + self.session.run([self.output_name], {self.input_name: input_tensor}), + ) return self.postprocess(outputs) @@ -158,7 +173,7 @@ class SimilarityOnnx: providers = providers or ["CPUExecutionProvider"] ort.set_default_logger_severity(3) self.session = ort.InferenceSession(str(model_path), providers=providers) - self.input_shape = [64, 64] + self.input_shape: tuple[int, int] = (64, 64) @staticmethod def sigmoid(value: np.ndarray) -> np.ndarray: @@ -175,24 +190,38 @@ class SimilarityOnnx: return Image.open(value) def _tensor(self, value) -> np.ndarray: - image = self._to_image(value).convert("RGB").resize(tuple(reversed(self.input_shape)), 1) + image = ( + self._to_image(value) + .convert("RGB") + .resize((self.input_shape[1], self.input_shape[0]), 1) + ) array = np.array(image).astype(np.float32) / 255.0 return np.expand_dims(np.transpose(array, (2, 0, 1)), 0) def score(self, image_1, image_2) -> int: - out = self.session.run(None, {"x1": self._tensor(image_1), "x2": self._tensor(image_2)}) - similarity = self.sigmoid(out[0])[0][0] + out = self.session.run( + None, {"x1": self._tensor(image_1), "x2": self._tensor(image_2)} + ) + similarity = self.sigmoid(cast(np.ndarray, out[0]))[0][0] return int(round(similarity.item(), 2) * 100) class HuyaCaptchaOcr: """封装虎牙滑块和点选识别。""" - def __init__(self, model_dir: str | Path = MODEL_DIR, providers: list[str] | None = None): + def __init__( + self, model_dir: str | Path = MODEL_DIR, providers: list[str] | None = None + ): model_dir = Path(model_dir) - self.similarity = SimilarityOnnx(model_dir / "weights.onnx", providers=providers) - self.click_model = YoloOnnx(model_dir / "best.onnx", classes=["target", "char"], providers=providers) - self.slider_model = YoloOnnx(model_dir / "slider_2.onnx", classes=["slider"], providers=providers) + self.similarity = SimilarityOnnx( + model_dir / "weights.onnx", providers=providers + ) + self.click_model = YoloOnnx( + model_dir / "best.onnx", classes=["target", "char"], providers=providers + ) + self.slider_model = YoloOnnx( + model_dir / "slider_2.onnx", classes=["slider"], providers=providers + ) @staticmethod def _open_image(image_bytes: bytes) -> Image.Image: @@ -209,11 +238,16 @@ class HuyaCaptchaOcr: """识别点选坐标,返回按目标顺序排列的点击点。""" image = self._open_image(image_bytes) detections = self.click_model.detect(image) - results = [{**item, "cropped_image": image.crop(tuple(item["xyxy"]))} for item in detections] + results = [ + {**item, "cropped_image": image.crop(tuple(item["xyxy"]))} + for item in detections + ] char_list = [item for item in results if item.get("label_name") == "char"] target_list = [item for item in results if item.get("label_name") == "target"] - target_list = [item for item in target_list if item["xyxy"][2] - item["xyxy"][0] > 10] + target_list = [ + item for item in target_list if item["xyxy"][2] - item["xyxy"][0] > 10 + ] char_list.sort(key=lambda item: item["xyxy"][0]) target_list.sort(key=lambda item: item["xyxy"][0])