refactor(gateway/yuanbao-proto): hug signatures/dicts in proto+media, Counter-based sticker scoring

This commit is contained in:
Teknium
2026-09-02 20:20:41 -07:00
parent 24b573ce43
commit 127f1752d0
3 changed files with 34 additions and 80 deletions
+12 -38
View File
@@ -189,14 +189,8 @@ def _sorted_kv(d: dict[str, str]) -> list[tuple[str, str]]:
def _cos_sign(
method: str,
path: str,
params: dict[str, str],
headers: dict[str, str],
secret_id: str,
secret_key: str,
start_time: Optional[int] = None,
expire_seconds: int = 3600,
method: str, path: str, params: dict[str, str], headers: dict[str, str], secret_id: str, secret_key: str,
start_time: Optional[int] = None, expire_seconds: int = 3600,
) -> str:
"""COS Authorization 头(q-sign-algorithm=sha1;https://cloud.tencent.com/document/product/436/7778)。
@@ -208,12 +202,8 @@ def _cos_sign(
sign_key = _hmac_sha1_hex(secret_key, q_sign_time) # SignKey = HMAC-SHA1(SecretKey, q-sign-time)
sorted_params = _sorted_kv(params)
sorted_headers = _sorted_kv(headers)
http_string = "\n".join([
method.lower(), path,
"&".join(f"{k}={v}" for k, v in sorted_params),
"&".join(f"{k}={v}" for k, v in sorted_headers),
"",
])
kv = lambda pairs: "&".join(f"{k}={v}" for k, v in pairs) # noqa: E731
http_string = "\n".join([method.lower(), path, kv(sorted_params), kv(sorted_headers), ""])
string_to_sign = "\n".join(["sha1", q_sign_time, hashlib.sha1(http_string.encode("utf-8")).hexdigest(), ""])
return (
f"q-sign-algorithm=sha1&q-ak={secret_id}&q-sign-time={q_sign_time}&q-key-time={q_sign_time}"
@@ -226,13 +216,8 @@ def _cos_sign(
# ============ 主要公开 API ============
async def get_cos_credentials(
app_key: str,
api_domain: str,
token: str,
filename: str = "file",
file_id: Optional[str] = None,
bot_id: str = "",
route_env: str = "",
app_key: str, api_domain: str, token: str, filename: str = "file", file_id: Optional[str] = None,
bot_id: str = "", route_env: str = "",
) -> dict:
"""调用 genUploadInfo 获取 COS 临时密钥及上传配置。
@@ -260,22 +245,16 @@ async def get_cos_credentials(
async def upload_to_cos(
file_bytes: bytes,
filename: str,
content_type: str,
credentials: dict,
bucket: str,
region: str,
file_bytes: bytes, filename: str, content_type: str, credentials: dict, bucket: str, region: str,
) -> dict:
"""用临时凭证(get_cos_credentials() 返回的 dict)HMAC-SHA1 签名,httpx PUT 上传到 COS(走全球加速域名)。
Returns {url, uuid (内容 MD5), size, width?, height? (仅图片)}
Raises httpx.HTTPStatusError(COS 非 2xx)/ RuntimeError(credentials 字段缺失)
"""
secret_id: str = credentials.get("encryptTmpSecretId", "")
secret_key: str = credentials.get("encryptTmpSecretKey", "")
session_token: str = credentials.get("encryptToken", "")
cos_key: str = credentials.get("location", "")
secret_id, secret_key, session_token, cos_key = (
credentials.get(k, "") for k in ("encryptTmpSecretId", "encryptTmpSecretKey", "encryptToken", "location")
)
start_time: Optional[int] = credentials.get("startTime")
expired_time: Optional[int] = credentials.get("expiredTime")
if not secret_id or not secret_key or not cos_key:
@@ -311,13 +290,8 @@ async def upload_to_cos(
# ============ TIM 媒体消息构建(https://cloud.tencent.com/document/product/269/2720) ============
def build_image_msg_body(
url: str,
uuid: Optional[str] = None,
filename: Optional[str] = None,
size: int = 0,
width: int = 0,
height: int = 0,
mime_type: str = "",
url: str, uuid: Optional[str] = None, filename: Optional[str] = None, size: int = 0, width: int = 0,
height: int = 0, mime_type: str = "",
) -> list[dict]:
"""TIMImageElem 消息体(可直接放入 msg_body)。uuid 缺省依次退到 filename / URL basename / "image"。"""
return [{
+17 -24
View File
@@ -235,11 +235,10 @@ def _encode_head(
def _decode_head(data: bytes) -> dict:
fdict = _parse_dict(data)
fd = _parse_dict(data)
return {
"cmd_type": _get_varint(fdict, 1), "cmd": _get_string(fdict, 2), "seq_no": _get_varint(fdict, 3),
"msg_id": _get_string(fdict, 4), "module": _get_string(fdict, 5),
"need_ack": bool(_get_varint(fdict, 6)), "status": _get_varint(fdict, 10),
"cmd_type": _get_varint(fd, 1), "cmd": _get_string(fd, 2), "seq_no": _get_varint(fd, 3), "msg_id": _get_string(fd, 4),
"module": _get_string(fd, 5), "need_ack": bool(_get_varint(fd, 6)), "status": _get_varint(fd, 10),
}
@@ -322,11 +321,8 @@ def _decode_msg_content(data: bytes) -> dict:
fdict = _parse_dict(data)
content = _decode_spec(fdict, _MSG_CONTENT_SPEC)
imgs = [img for img in (_decode_spec(d, _IMAGE_INFO_SPEC) for d in _parse_repeated(fdict, 8)) if img]
if imgs:
content["image_info_array"] = imgs
ext_map = {_get_string(e, 1): _get_string(e, 2) for e in _parse_repeated(fdict, 999) if _get_string(e, 1)}
if ext_map:
content["ext_map"] = ext_map
content.update({k: v for k, v in (("image_info_array", imgs), ("ext_map", ext_map)) if v})
return content
@@ -395,21 +391,19 @@ def _decode_forward_msg_content(data: bytes) -> dict:
return content
def _decode_forward_msg(fdict: dict) -> dict:
return {
"sender": _get_string(fdict, 1), "time": _get_varint(fdict, 2), "plainText": _get_string(fdict, 3),
"msgContent": [_decode_forward_msg_content(b) for b in _get_repeated_bytes(fdict, 4)],
}
def _decode_forward_msg(fd: dict) -> dict:
return {"sender": _get_string(fd, 1), "time": _get_varint(fd, 2), "plainText": _get_string(fd, 3),
"msgContent": [_decode_forward_msg_content(b) for b in _get_repeated_bytes(fd, 4)]}
def decode_forward_msg_data(data: bytes) -> Optional[dict]:
"""Parse ForwardMsgData bytes (base64-decoded ext_map value) into the {sub_type, nick_name, msg, ...}
structure consumed by ForwardedRecordsParseMiddleware.build_forward_text; None on parse failure."""
try:
fdict = _parse_dict(data)
fd = _parse_dict(data)
return {
"sub_type": _get_varint(fdict, 1), "begin_time": _get_varint(fdict, 2), "end_time": _get_varint(fdict, 3),
"nick_name": _get_string(fdict, 4), "msg": [_decode_forward_msg(d) for d in _parse_repeated(fdict, 5)],
"sub_type": _get_varint(fd, 1), "begin_time": _get_varint(fd, 2), "end_time": _get_varint(fd, 3),
"nick_name": _get_string(fd, 4), "msg": [_decode_forward_msg(d) for d in _parse_repeated(fd, 5)],
}
except Exception:
return None
@@ -532,15 +526,14 @@ def decode_get_group_member_list_rsp(data: bytes) -> Optional[dict]:
member dict 过滤空值但保留 role;解析失败返回 None。"""
try:
fdict = _parse_dict(data)
members = []
for m in _parse_repeated(fdict, 3):
member = {
"user_id": _get_string(m, 1), "nickname": _get_string(m, 2), "role": _get_varint(m, 3),
"join_time": _get_varint(m, 4), "name_card": _get_string(m, 5),
}
members.append({k: v for k, v in member.items() if v or k == "role"})
members = [
{"user_id": _get_string(m, 1), "nickname": _get_string(m, 2), "role": _get_varint(m, 3),
"join_time": _get_varint(m, 4), "name_card": _get_string(m, 5)}
for m in _parse_repeated(fdict, 3)
]
return {
"code": _get_varint(fdict, 1), "message": _get_string(fdict, 2), "members": members,
"code": _get_varint(fdict, 1), "message": _get_string(fdict, 2),
"members": [{k: v for k, v in mem.items() if v or k == "role"} for mem in members],
"next_offset": _get_varint(fdict, 4), "is_complete": bool(_get_varint(fdict, 5)),
}
except Exception:
+5 -18
View File
@@ -12,6 +12,7 @@ import json
import random
import re
import unicodedata
from collections import Counter
from typing import Optional
# Sticker catalogue – ported from builtin-stickers.json. Every builtin sticker is in
@@ -131,27 +132,14 @@ def _compact_text(raw: str) -> str:
def _multiset_char_hit_ratio(needle: str, haystack: str) -> float:
if not needle:
return 0.0
bag: dict[str, int] = {}
for ch in haystack:
bag[ch] = bag.get(ch, 0) + 1
hits = 0
for ch in needle:
if bag.get(ch, 0) > 0:
hits += 1
bag[ch] -= 1
return hits / len(needle)
return sum((Counter(needle) & Counter(haystack)).values()) / len(needle) if needle else 0.0
def _bigram_jaccard(a: str, b: str) -> float:
if len(a) < 2 or len(b) < 2:
return 0.0
A = {a[i:i + 2] for i in range(len(a) - 1)}
B = {b[i:i + 2] for i in range(len(b) - 1)}
inter = len(A & B)
union = len(A) + len(B) - inter
return inter / union if union else 0.0
A, B = ({x[i:i + 2] for i in range(len(x) - 1)} for x in (a, b))
return len(A & B) / len(A | B)
def _longest_subsequence_ratio(needle: str, haystack: str) -> float:
@@ -161,8 +149,7 @@ def _longest_subsequence_ratio(needle: str, haystack: str) -> float:
for ch in haystack:
if j >= len(needle):
break
if ch == needle[j]:
j += 1
j += ch == needle[j]
return j / len(needle)