初始提交:识流 AI 助手项目

微信自动回复机器人,基于截图+OCR识别消息,支持关键词规则和 AI(OpenAI/DeepSeek/Dify)自动回复。
技术栈:PySide6 + Flask + Vue3 + RapidOCR + SQLite

注:OCR大模型文件(.onnx / .pdiparams)不纳入版本控制,需单独下载。

🤖 Generated with [Qoder][https://qoder.com]
This commit is contained in:
figmar
2026-05-30 14:57:45 +08:00
commit 81115dc23d
129 changed files with 56398 additions and 0 deletions
+81
View File
@@ -0,0 +1,81 @@
import requests
from app.infrastructure.service.backend.config import AI_PROVIDER, DEEPSEEK_API_BASE, DEEPSEEK_API_KEY, DEEPSEEK_MODEL, DIFY_API_BASE, DIFY_API_KEY, DIFY_USER, OPENAI_API_BASE, OPENAI_API_KEY, OPENAI_MODEL
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
def do_openai_like(url, headers, payload):
try:
resp = requests.post(url, headers=headers, json=payload, timeout=60)
data = resp.json() if resp.text else {}
if resp.status_code >= 400:
return f"抱歉,AI 服务请求失败({resp.status_code}"
content = (((data or {}).get("choices") or [{}])[0].get("message") or {}).get("content", "")
return content.strip() if content else "抱歉,AI 暂时没有合理的回复。"
except Exception:
return "抱歉,AI 服务暂时不可用,请稍后再试。"
def do_dify(url, headers, payload):
try:
resp = requests.post(url, headers=headers, json=payload, timeout=60)
text = resp.text or ""
if "data:" in text:
answer = ""
for line in text.splitlines():
line = line.strip()
if not line.startswith("data:"):
continue
chunk = line[5:].strip()
if not chunk or chunk == "[DONE]":
continue
try:
j = requests.models.complexjson.loads(chunk)
except Exception:
continue
if "answer" in j:
answer += j["answer"]
if answer.strip():
return answer.strip()
data = resp.json() if resp.text else {}
if resp.status_code >= 400:
return f"抱歉,Dify 服务请求失败({resp.status_code}"
return (data.get("answer") or "抱歉,Dify 暂时没有合理的回复。").strip()
except Exception:
return "抱歉,Dify 服务暂时不可用,请稍后再试。"
def call_ai(prompt, user_id=""):
trace_id = new_trace_id("ai")
provider = AI_PROVIDER
log_event("INFO", "ai", "ai.request", trace_id, "request", "ok", "发起AI请求", extra={"provider": provider, "user_id": user_id or "", "prompt_len": len(prompt or "")})
if provider == "mock":
result = "【自动回复】你刚才说了:" + (prompt or "")[:100]
log_event("INFO", "ai", "ai.response", trace_id, "response", "ok", "AI回复完成", extra={"provider": provider, "reply_len": len(result)})
return result
if provider == "openai":
result = do_openai_like(
OPENAI_API_BASE.rstrip("/") + "/chat/completions",
{"Content-Type": "application/json", "Authorization": f"Bearer {OPENAI_API_KEY}"},
{"model": OPENAI_MODEL, "messages": [{"role": "system", "content": "你是一个专业的微信私域运营助手,用简洁自然的中文回复用户。"}, {"role": "user", "content": prompt}], "temperature": 0.7, "user": user_id or None},
)
elif provider == "deepseek":
result = do_openai_like(
DEEPSEEK_API_BASE.rstrip("/") + "/chat/completions",
{"Content-Type": "application/json", "Authorization": f"Bearer {DEEPSEEK_API_KEY}"},
{"model": DEEPSEEK_MODEL, "messages": [{"role": "system", "content": "你是一个简洁高效的微信助手。回复要求:一句话,不超过50字。"}, {"role": "user", "content": prompt}], "temperature": 0.7, "max_tokens": 100, "user": user_id or None},
)
elif provider == "dify":
result = do_dify(
DIFY_API_BASE.rstrip("/") + "/chat-messages",
{"Content-Type": "application/json", "Authorization": f"Bearer {DIFY_API_KEY}"},
{"inputs": {}, "query": prompt, "response_mode": "streaming", "user": user_id or DIFY_USER, "conversation_id": ""},
)
else:
result = "AI_PROVIDER 未配置正确,请检查环境变量。"
if result.startswith("抱歉") or "未配置正确" in result:
log_event("WARNING", "ai", "ai.response", trace_id, "response", "failed", "AI回复异常或降级", reason="provider_error", extra={"provider": provider, "reply": result[:120]})
else:
log_event("INFO", "ai", "ai.response", trace_id, "response", "ok", "AI回复完成", extra={"provider": provider, "reply_len": len(result)})
return result
@@ -0,0 +1,45 @@
import os
from pathlib import Path
from flask import Flask, request
from app.configs.runtime_config import get_int, get_str
PROJECT_ROOT = Path(__file__).resolve().parents[4]
ASSETS_DIR = PROJECT_ROOT / "assets"
FRONTEND_DIST_DIR = PROJECT_ROOT / "frontend" / "dist"
_frontend_static_dir = get_str("FRONTEND_STATIC_DIR", "").strip()
FRONTEND_STATIC_DIR = Path(_frontend_static_dir).resolve() if _frontend_static_dir else None
STATIC_DIR = FRONTEND_STATIC_DIR or (FRONTEND_DIST_DIR if FRONTEND_DIST_DIR.exists() else ASSETS_DIR)
DB_HOST = get_str("DB_HOST", "127.0.0.1")
DB_PORT = get_int("DB_PORT", 3306)
DB_NAME = get_str("DB_NAME", "ai_shiliu")
DB_USER = get_str("DB_USER", "ai_shiliu")
DB_PASS = get_str("DB_PASS", "")
DB_CHARSET = get_str("DB_CHARSET", "utf8mb4")
AI_PROVIDER = get_str("AI_PROVIDER", "")
OPENAI_API_KEY = get_str("OPENAI_API_KEY", "")
OPENAI_API_BASE = get_str("OPENAI_API_BASE", "")
OPENAI_MODEL = get_str("OPENAI_MODEL", "")
DEEPSEEK_API_KEY = get_str("DEEPSEEK_API_KEY", "")
DEEPSEEK_API_BASE = get_str("DEEPSEEK_API_BASE", "")
DEEPSEEK_MODEL = get_str("DEEPSEEK_MODEL", "")
DIFY_API_KEY = get_str("DIFY_API_KEY", "")
DIFY_API_BASE = get_str("DIFY_API_BASE", "")
DIFY_USER = get_str("DIFY_USER", "")
app = Flask(__name__, static_folder=str(STATIC_DIR), static_url_path="")
@app.after_request
def add_cors_headers(response):
origin = (request.headers.get("Origin") or "").strip()
allow_origin = origin or "*"
response.headers["Access-Control-Allow-Origin"] = allow_origin
response.headers["Vary"] = "Origin"
response.headers["Access-Control-Allow-Headers"] = "Content-Type, Authorization"
response.headers["Access-Control-Allow-Methods"] = "GET, POST, OPTIONS"
return response
+218
View File
@@ -0,0 +1,218 @@
import json
import os
import re
import sqlite3
from pathlib import Path
from app.infrastructure.service.backend.config import PROJECT_ROOT
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
_SETTING_CACHE = {}
SETTINGS_FILE = PROJECT_ROOT / "logs" / "state" / "local_settings.json"
LOCAL_APPDATA_DIR = Path(os.environ.get("LOCALAPPDATA", str(Path.home() / "AppData" / "Local")))
SQLITE_DB_PATH = LOCAL_APPDATA_DIR / "com.shiliu.aiassistant" / "ai_shiliu.sqlite3"
class _SQLiteCursor:
def __init__(self, cursor):
self._cursor = cursor
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
self._cursor.close()
@staticmethod
def _adapt_sql(sql: str) -> str:
return sql.replace("%s", "?")
def execute(self, sql, params=None):
sql = self._adapt_sql(sql)
if params is None:
self._cursor.execute(sql)
else:
self._cursor.execute(sql, params)
return self
def fetchone(self):
row = self._cursor.fetchone()
return dict(row) if row is not None else None
def fetchall(self):
return [dict(row) for row in self._cursor.fetchall()]
@property
def lastrowid(self):
return self._cursor.lastrowid
class _SQLiteConn:
def __init__(self, conn):
self._conn = conn
def cursor(self):
return _SQLiteCursor(self._conn.cursor())
def commit(self):
self._conn.commit()
def rollback(self):
self._conn.rollback()
def close(self):
self._conn.close()
def _bootstrap_sqlite_file():
SQLITE_DB_PATH.parent.mkdir(parents=True, exist_ok=True)
def get_conn(db_name=None):
_bootstrap_sqlite_file()
conn = sqlite3.connect(str(SQLITE_DB_PATH))
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA foreign_keys = ON")
return _SQLiteConn(conn)
def init_db():
trace_id = new_trace_id("db")
log_event("INFO", "db", "db.init", trace_id, "start", "ok", "初始化数据库开始", extra={"path": str(SQLITE_DB_PATH)})
conn = get_conn()
try:
with conn.cursor() as cur:
cur.execute(
"""
CREATE TABLE IF NOT EXISTS messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
wx_user_id TEXT NOT NULL DEFAULT '',
wx_nickname TEXT NOT NULL DEFAULT '',
direction TEXT NOT NULL DEFAULT 'in',
content TEXT NOT NULL,
is_ai_reply INTEGER NOT NULL DEFAULT 0,
rule_id INTEGER NULL,
is_friend_request INTEGER NOT NULL DEFAULT 0,
reply_strategy TEXT NOT NULL DEFAULT '',
reply_reason TEXT NOT NULL DEFAULT '',
ocr_confidence TEXT NOT NULL DEFAULT '',
ocr_bubble_side TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime'))
)
"""
)
cur.execute("CREATE INDEX IF NOT EXISTS idx_messages_user_time ON messages(wx_user_id, created_at)")
cur.execute(
"""
CREATE TABLE IF NOT EXISTS auto_reply_rules (
id INTEGER PRIMARY KEY AUTOINCREMENT,
keyword TEXT NOT NULL,
match_type TEXT NOT NULL DEFAULT 'contain',
reply_text TEXT NOT NULL,
is_active INTEGER NOT NULL DEFAULT 1,
created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')),
updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime'))
)
"""
)
cur.execute("PRAGMA table_info(messages)")
cols = {str(x.get('name') or '') for x in (cur.fetchall() or [])}
if "reply_strategy" not in cols:
cur.execute("ALTER TABLE messages ADD COLUMN reply_strategy TEXT NOT NULL DEFAULT ''")
if "reply_reason" not in cols:
cur.execute("ALTER TABLE messages ADD COLUMN reply_reason TEXT NOT NULL DEFAULT ''")
if "ocr_confidence" not in cols:
cur.execute("ALTER TABLE messages ADD COLUMN ocr_confidence TEXT NOT NULL DEFAULT ''")
if "ocr_bubble_side" not in cols:
cur.execute("ALTER TABLE messages ADD COLUMN ocr_bubble_side TEXT NOT NULL DEFAULT ''")
conn.commit()
log_event("INFO", "db", "db.init", trace_id, "done", "ok", "初始化数据库完成")
except Exception as exc:
conn.rollback()
log_event("ERROR", "db", "db.init", trace_id, "done", "failed", "初始化数据库失败", reason="db_error", extra={"error": str(exc)})
raise
finally:
conn.close()
def _load_settings_file():
if _SETTING_CACHE:
return
try:
if SETTINGS_FILE.exists():
data = json.loads(SETTINGS_FILE.read_text(encoding="utf-8"))
if isinstance(data, dict):
for k, v in data.items():
_SETTING_CACHE[str(k)] = str(v)
except Exception:
pass
def _save_settings_file():
SETTINGS_FILE.parent.mkdir(parents=True, exist_ok=True)
SETTINGS_FILE.write_text(json.dumps(_SETTING_CACHE, ensure_ascii=False, indent=2), encoding="utf-8")
def get_setting(key, default=None):
_load_settings_file()
if key in _SETTING_CACHE:
return _SETTING_CACHE[key]
if default is None:
return None
val = str(default)
_SETTING_CACHE[key] = val
_save_settings_file()
return val
def set_setting(key, value):
_load_settings_file()
_SETTING_CACHE[str(key)] = str(value)
_save_settings_file()
def normalize_text(text):
t = (text or "").strip().lower()
t = re.sub(r"\s+", "", t)
t = t.replace("", ":")
return t
def find_rule_reply(content):
trace_id = new_trace_id("db")
conn = get_conn()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM auto_reply_rules WHERE is_active = 1 ORDER BY id ASC")
rules = cur.fetchall()
except Exception as exc:
log_event("ERROR", "db", "db.rule.query", trace_id, "query", "failed", "查询规则失败", reason="db_error", extra={"error": str(exc)})
raise
finally:
conn.close()
raw_content = (content or "").strip()
content_lower = raw_content.lower()
content_norm = normalize_text(raw_content)
for rule in rules:
kw = (rule.get("keyword") or "").strip()
if not kw:
continue
kw_lower = kw.lower()
kw_norm = normalize_text(kw)
match_type = rule.get("match_type")
if match_type == "equal":
if content_lower == kw_lower or content_norm == kw_norm:
log_event("INFO", "db", "db.rule.match", trace_id, "match", "ok", "命中规则", reason="rule_hit", extra={"rule_id": rule.get("id"), "match_type": match_type})
return rule
else:
if kw_lower in content_lower or kw_norm in content_norm:
log_event("INFO", "db", "db.rule.match", trace_id, "match", "ok", "命中规则", reason="rule_hit", extra={"rule_id": rule.get("id"), "match_type": match_type or "contain"})
return rule
log_event("INFO", "db", "db.rule.match", trace_id, "match", "ok", "未命中规则", reason="rule_miss", extra={"rule_count": len(rules)})
return None
@@ -0,0 +1,141 @@
import hashlib
import json
from datetime import datetime
from pathlib import Path
from app.infrastructure.service.wechat.config import LOG_ROOT_DIR
ALLOWED_MODULES = {"api", "bot", "ocr", "ai", "db", "capture", "audit", "error"}
DOMAIN_MODULES = {"api", "bot", "ocr", "ai", "db", "capture"}
def _infer_domain(row: dict) -> str:
module = str(row.get("module") or "").strip().lower()
if module in DOMAIN_MODULES:
return module
if module == "error":
event = str(row.get("event") or "")
prefix = event.split(".", 1)[0].strip().lower()
if prefix in DOMAIN_MODULES:
return prefix
return "api"
def _event_id(row: dict) -> str:
raw = json.dumps(row or {}, ensure_ascii=False, sort_keys=True)
return hashlib.md5(raw.encode("utf-8")).hexdigest()
def _enrich_row(row: dict) -> dict:
x = dict(row or {})
x["domain"] = _infer_domain(x)
x["event_id"] = _event_id(x)
return x
def _read_jsonl(module: str):
p = Path(LOG_ROOT_DIR) / f"{module}.jsonl"
if not p.exists():
return []
rows = []
with p.open("r", encoding="utf-8", errors="replace") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
rows.append(json.loads(line))
except Exception:
continue
return rows
def query_events(module=None, level=None, event=None, trace_id=None, start_ts=None, end_ts=None, keyword=None, page=1, size=50):
modules = [module] if module in ALLOWED_MODULES else sorted(ALLOWED_MODULES)
all_rows = []
for m in modules:
all_rows.extend(_read_jsonl(m))
def ok(row):
if level and str(row.get("level", "")).upper() != str(level).upper():
return False
if event and str(row.get("event", "")) != str(event):
return False
if trace_id and str(row.get("trace_id", "")) != str(trace_id):
return False
if keyword and keyword not in json.dumps(row, ensure_ascii=False):
return False
ts = str(row.get("ts") or "")
if start_ts and ts < start_ts:
return False
if end_ts and ts > end_ts:
return False
return True
rows = [_enrich_row(r) for r in all_rows if ok(r)]
rows.sort(key=lambda x: str(x.get("ts") or ""), reverse=True)
page = max(1, int(page or 1))
size = max(1, min(200, int(size or 50)))
start = (page - 1) * size
end = start + size
return {"total": len(rows), "page": page, "size": size, "items": rows[start:end]}
def query_trace(trace_id: str):
if not trace_id:
return []
result = query_events(trace_id=trace_id, size=500)
items = result.get("items") or []
items.sort(key=lambda x: str(x.get("ts") or ""))
return items
def query_event_json(event_id: str):
event_id = str(event_id or "").strip().lower()
if not event_id:
return None
for m in sorted(ALLOWED_MODULES):
rows = _read_jsonl(m)
for row in rows:
x = _enrich_row(row)
if str(x.get("event_id") or "") == event_id:
return x
return None
def clear_logs(module=None):
modules = [module] if module in ALLOWED_MODULES else sorted(ALLOWED_MODULES)
root = Path(LOG_ROOT_DIR)
deleted = []
for m in modules:
for p in root.glob(f"{m}.log*"):
p.write_text("", encoding="utf-8")
deleted.append(str(p.name))
for p in root.glob(f"{m}.jsonl*"):
p.write_text("", encoding="utf-8")
deleted.append(str(p.name))
return {"modules": modules, "files": deleted}
def query_summary(limit=300):
events = query_events(size=limit).get("items") or []
error_count = sum(1 for e in events if str(e.get("level", "")).upper() in {"ERROR"})
fallback_count = sum(1 for e in events if str(e.get("event", "")).endswith("fallback"))
reasons = {}
domain_counts = {"api": 0, "bot": 0, "ocr": 0, "ai": 0, "db": 0, "capture": 0}
for e in events:
reason = str(e.get("reason") or "").strip()
if reason:
reasons[reason] = reasons.get(reason, 0) + 1
domain = str(e.get("domain") or _infer_domain(e))
if domain in domain_counts:
domain_counts[domain] += 1
reason_top = sorted(reasons.items(), key=lambda x: x[1], reverse=True)[:20]
domain_top = sorted(domain_counts.items(), key=lambda x: x[1], reverse=True)
return {
"window": len(events),
"error_count": error_count,
"fallback_count": fallback_count,
"top_reasons": [{"reason": k, "count": v} for k, v in reason_top],
"domain_counts": [{"domain": k, "count": v} for k, v in domain_top],
"generated_at": datetime.now().astimezone().isoformat(timespec="seconds"),
}
@@ -0,0 +1,86 @@
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
MODULES = {"api", "bot", "ocr", "ai", "db", "capture", "audit", "error"}
DOMAIN_MODULES = {"api", "bot", "ocr", "ai", "db", "capture"}
ALLOWED_EVENTS = {
"api.bot.status",
"api.bot.start",
"api.bot.stop",
"api.messages.receive",
"audit.decision",
"bot.chat_analyze",
"bot.chat_snapshot",
"bot.loop",
"bot.session_scan",
"bot.session_service.init",
"bot.session_title",
"bot.submit",
"bot.unread.detect",
"bot.unread.scan",
"ocr.baidu.token",
"ocr.baidu.recognize",
"ocr.rapid.init",
"ocr.rapid.recognize",
"ocr.fallback",
"ocr.session_name",
"ocr.session_title",
"ocr.generic",
"capture.contact_list",
"capture.session_title",
"capture.chat_area",
"ai.request",
"ai.response",
"db.init",
"db.error",
"db.rule.query",
"db.rule.match",
}
def normalize_event(module: str, event: str) -> str:
e = str(event or "").strip().lower()
m = str(module or "").strip().lower()
if e in ALLOWED_EVENTS:
return e
if "." in e:
prefix = e.split(".", 1)[0]
if prefix in DOMAIN_MODULES:
return e
if m in DOMAIN_MODULES:
return f"{m}.unknown"
return "api.unknown"
@dataclass
class LogEvent:
level: str
module: str
event: str
trace_id: str
stage: str
status: str
message: str
reason: str = ""
extra: dict[str, Any] = field(default_factory=dict)
ts: str = ""
def to_dict(self) -> dict[str, Any]:
ts = self.ts or datetime.now().astimezone().isoformat(timespec="milliseconds")
module = self.module if self.module in MODULES else "api"
event = normalize_event(module, self.event)
return {
"ts": ts,
"level": (self.level or "INFO").upper(),
"module": module,
"event": event,
"trace_id": self.trace_id or "-",
"stage": self.stage or "-",
"status": self.status or "ok",
"reason": self.reason or "",
"message": self.message or "",
"extra": self.extra or {},
}
@@ -0,0 +1,102 @@
import json
import logging
import os
from logging.handlers import RotatingFileHandler
from pathlib import Path
from app.configs.runtime_config import get_bool, get_int, get_str
from app.infrastructure.service.logging.log_schema import ALLOWED_EVENTS, LogEvent, MODULES, normalize_event
from app.infrastructure.service.wechat.config import LOG_ROOT_DIR
_INITIALIZED = False
_LOGGERS: dict[str, logging.Logger] = {}
_JSON_LOGGERS: dict[str, logging.Logger] = {}
class _JsonFormatter(logging.Formatter):
def format(self, record):
if isinstance(record.msg, dict):
return json.dumps(record.msg, ensure_ascii=False)
return json.dumps({"message": str(record.msg)}, ensure_ascii=False)
class _TextFormatter(logging.Formatter):
def format(self, record):
if isinstance(record.msg, dict):
data = record.msg
extra = data.get("extra") or {}
kv = " ".join([f"{k}={v}" for k, v in extra.items()])
suffix = f" | {kv}" if kv else ""
return f"[{data.get('ts')}][{data.get('level')}][{data.get('module')}][{data.get('event')}][{data.get('trace_id')}] {data.get('message')}{suffix}"
return str(record.msg)
def _build_logger(name: str, file_path: Path, formatter: logging.Formatter, level: int, rotate_mb: int, backup_count: int):
logger = logging.getLogger(name)
logger.handlers.clear()
logger.propagate = False
logger.setLevel(level)
file_path.parent.mkdir(parents=True, exist_ok=True)
handler = RotatingFileHandler(str(file_path), maxBytes=rotate_mb * 1024 * 1024, backupCount=backup_count, encoding="utf-8")
handler.setLevel(level)
handler.setFormatter(formatter)
logger.addHandler(handler)
return logger
def init_logging():
global _INITIALIZED
if _INITIALIZED:
return
enabled = get_bool("LOG_ENABLED", True)
if not enabled:
logging.disable(logging.CRITICAL)
_INITIALIZED = True
return
level_name = (get_str("LOG_LEVEL", "INFO") or "INFO").upper()
level = getattr(logging, level_name, logging.INFO)
rotate_mb = max(1, get_int("LOG_ROTATE_MB", 5))
backup_count = max(1, get_int("LOG_BACKUP_COUNT", 7))
root = Path(LOG_ROOT_DIR)
for module in MODULES:
text_logger = _build_logger(f"solo.{module}.text", root / f"{module}.log", _TextFormatter(), level, rotate_mb, backup_count)
json_logger = _build_logger(f"solo.{module}.json", root / f"{module}.jsonl", _JsonFormatter(), level, rotate_mb, backup_count)
_LOGGERS[module] = text_logger
_JSON_LOGGERS[module] = json_logger
_INITIALIZED = True
def log_event(level: str, module: str, event: str, trace_id: str, stage: str, status: str, message: str, reason: str = "", extra: dict | None = None):
if not _INITIALIZED:
init_logging()
if logging.root.manager.disable >= logging.CRITICAL:
return
normalized_event = normalize_event(module, event)
payload_extra = dict(extra or {})
if normalized_event != str(event or "").strip().lower() and normalized_event not in ALLOWED_EVENTS:
payload_extra["event_raw"] = event
payload_extra["event_normalized"] = normalized_event
payload = LogEvent(
level=level,
module=module,
event=normalized_event,
trace_id=trace_id,
stage=stage,
status=status,
reason=reason,
message=message,
extra=payload_extra,
).to_dict()
module_name = payload["module"]
lvl = getattr(logging, payload["level"], logging.INFO)
_LOGGERS[module_name].log(lvl, payload)
_JSON_LOGGERS[module_name].log(lvl, payload)
if lvl >= logging.ERROR and module_name != "error":
_LOGGERS["error"].log(lvl, payload)
_JSON_LOGGERS["error"].log(lvl, payload)
def new_trace_id(prefix: str = "trace") -> str:
import uuid
return f"{prefix}_{uuid.uuid4().hex[:12]}"
File diff suppressed because it is too large Load Diff
+124
View File
@@ -0,0 +1,124 @@
import os
from pathlib import Path
from app.configs.runtime_config import get_bool, get_float, get_int, get_str
PROJECT_ROOT = Path(__file__).resolve().parents[4]
def _local_appdata_root() -> Path:
local_appdata = (os.getenv("LOCALAPPDATA") or "").strip()
if local_appdata:
return Path(local_appdata)
return Path.home() / "AppData" / "Local"
def _resolve_path(path_value: str, base_dir: Path) -> str:
path = Path(path_value)
if not path.is_absolute():
path = base_dir / path
return str(path.resolve())
APP_NAME = (get_str("APP_NAME", "AiShiliu") or "AiShiliu").strip() or "AiShiliu"
APP_DATA_DIR = _resolve_path(
get_str("APP_DATA_DIR", "") or get_str("OPENCLAW_APP_DATA_DIR", "") or str(_local_appdata_root() / APP_NAME),
PROJECT_ROOT,
)
LOG_ROOT_DIR = _resolve_path(get_str("LOG_ROOT_DIR", os.path.join(APP_DATA_DIR, "logs")), PROJECT_ROOT)
BACKEND_LOG_DIR = _resolve_path(get_str("BACKEND_LOG_DIR", os.path.join(LOG_ROOT_DIR, "backend")), PROJECT_ROOT)
FRONTEND_LOG_DIR = _resolve_path(get_str("FRONTEND_LOG_DIR", os.path.join(LOG_ROOT_DIR, "frontend")), PROJECT_ROOT)
BACKEND_PYTHON_LOG_FILE = _resolve_path(
get_str("BACKEND_PYTHON_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "python", "backend.log")),
PROJECT_ROOT,
)
BOT_LOG_FILE = _resolve_path(
get_str("BOT_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "bot", "bot.log")),
PROJECT_ROOT,
)
BOT_SESSION_LIST_LOG_FILE = _resolve_path(
get_str("BOT_SESSION_LIST_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "bot", "session_list.log")),
PROJECT_ROOT,
)
BOT_SESSION_DETAIL_LOG_FILE = _resolve_path(
get_str("BOT_SESSION_DETAIL_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "bot", "session_detail.log")),
PROJECT_ROOT,
)
OCR_LOG_FILE = _resolve_path(
get_str("OCR_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "ocr", "ocr.log")),
PROJECT_ROOT,
)
OCR_BAIDU_LOG_FILE = _resolve_path(
get_str("OCR_BAIDU_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "ocr", "baidu.log")),
PROJECT_ROOT,
)
OCR_RAPID_LOG_FILE = _resolve_path(
get_str("OCR_RAPID_LOG_FILE", os.path.join(BACKEND_LOG_DIR, "ocr", "rapid.log")),
PROJECT_ROOT,
)
TAURI_LOG_FILE = _resolve_path(
get_str("TAURI_LOG_FILE", os.path.join(FRONTEND_LOG_DIR, "tauri", "tauri.log")),
PROJECT_ROOT,
)
VUE_LOG_FILE = _resolve_path(
get_str("VUE_LOG_FILE", os.path.join(FRONTEND_LOG_DIR, "vue", "vue.log")),
PROJECT_ROOT,
)
OCR_SAVE_DIR = _resolve_path(get_str("OCR_SAVE_DIR", os.path.join(BACKEND_LOG_DIR, "ocr_debug_images")), PROJECT_ROOT)
BLOCKED_ROW_CACHE_FILE = _resolve_path(
get_str("BLOCKED_ROW_CACHE_FILE", os.path.join(BACKEND_LOG_DIR, "state", "blocked_rows.json")),
PROJECT_ROOT,
)
BAIDU_API_KEY = get_str("BAIDU_API_KEY", "")
BAIDU_SECRET_KEY = get_str("BAIDU_SECRET_KEY", "")
OCR_PROVIDER = get_str("OCR_PROVIDER", "baidu").strip().lower()
RAPID_OCR_DET_MODEL_PATH = get_str("RAPID_OCR_DET_MODEL_PATH", "app/resources/ocr_models/ch_PP-OCRv4_det.onnx").strip()
RAPID_OCR_REC_MODEL_PATH = get_str("RAPID_OCR_REC_MODEL_PATH", "app/resources/ocr_models/ch_PP-OCRv4_rec.onnx").strip()
RAPID_OCR_CLS_MODEL_PATH = get_str("RAPID_OCR_CLS_MODEL_PATH", "app/resources/ocr_models/ch_ppocr_mobile_v2.0_cls.onnx").strip()
BACKEND_URL = get_str("BACKEND_URL", "http://127.0.0.1:5000/api/messages/receive")
LOOP_INTERVAL = get_int("BOT_LOOP_INTERVAL", 3)
CLICK_AFTER_DELAY = get_float("BOT_CLICK_AFTER_DELAY", 1.2)
TITLE_AFTER_DELAY = get_float("BOT_TITLE_AFTER_DELAY", 1.0)
CONTACT_SWITCH_DELAY = get_float("BOT_CONTACT_SWITCH_DELAY", 1.0)
LOOP_ERROR_DELAY = get_float("BOT_LOOP_ERROR_DELAY", 3)
WECHAT_WINDOW_TARGET_WIDTH = get_int("WECHAT_WINDOW_TARGET_WIDTH", 1080)
WECHAT_WINDOW_TARGET_HEIGHT = get_int("WECHAT_WINDOW_TARGET_HEIGHT", 820)
WECHAT_WINDOW_TARGET_LEFT = get_int("WECHAT_WINDOW_TARGET_LEFT", 120)
WECHAT_WINDOW_TARGET_TOP = get_int("WECHAT_WINDOW_TARGET_TOP", 80)
OCR_SAVE_IMAGES = get_bool("OCR_SAVE_IMAGES", True)
CONTACT_ROW_HEIGHT = get_int("CONTACT_ROW_HEIGHT", 64)
CONTACT_ROW_WIDTH = get_int("CONTACT_ROW_WIDTH", 240)
CONTACT_LIST_LEFT_OFFSET = get_int("CONTACT_LIST_LEFT_OFFSET", 68)
CONTACT_LIST_TOP_OFFSET = get_int("CONTACT_LIST_TOP_OFFSET", 82)
CONTACT_LIST_BOTTOM_OFFSET = get_int("CONTACT_LIST_BOTTOM_OFFSET", 0)
SESSION_NAME_LEFT_OFFSET = get_int("SESSION_NAME_LEFT_OFFSET", 56)
SESSION_NAME_TOP_OFFSET = get_int("SESSION_NAME_TOP_OFFSET", 8)
SESSION_NAME_WIDTH = get_int("SESSION_NAME_WIDTH", 134)
SESSION_NAME_HEIGHT = get_int("SESSION_NAME_HEIGHT", 24)
SESSION_NAME_OCR_SCALE = get_int("SESSION_NAME_OCR_SCALE", 4)
SESSION_NAME_OCR_EXTRA_SCALE = get_int("SESSION_NAME_OCR_EXTRA_SCALE", 6)
CHAT_CAPTURE_LEFT_OFFSET = get_int("CHAT_CAPTURE_LEFT_OFFSET", 310)
CHAT_CAPTURE_TOP_OFFSET = get_int("CHAT_CAPTURE_TOP_OFFSET", 70)
CHAT_CAPTURE_WIDTH = get_int("CHAT_CAPTURE_WIDTH", 750)
CHAT_CAPTURE_HEIGHT = get_int("CHAT_CAPTURE_HEIGHT", 550)
OCR_TOP_PENALTY_RATIO = get_float("OCR_TOP_PENALTY_RATIO", 0.18)
OCR_TOP_PENALTY_BIN_FACTOR = get_float("OCR_TOP_PENALTY_BIN_FACTOR", 2.0)
OCR_TOP_PENALTY_COLOR_FACTOR = get_float("OCR_TOP_PENALTY_COLOR_FACTOR", 2.2)
TITLE_OCR_AREA_LEFT_OFFSET = get_int("TITLE_OCR_AREA_LEFT_OFFSET", 240)
TITLE_OCR_AREA_TOP_OFFSET = get_int("TITLE_OCR_AREA_TOP_OFFSET", 4)
TITLE_OCR_AREA_WIDTH = get_int("TITLE_OCR_AREA_WIDTH", 600)
TITLE_OCR_AREA_HEIGHT = get_int("TITLE_OCR_AREA_HEIGHT", 64)
NO_REPLY_KEYWORDS = [
"谢谢", "好的", "", "", "ok", "收到",
"[图片]", "[语音]", "[视频]", "[文件]"
]
BLOCKED_SESSION_KEYWORDS = [
"服务号", "公众号", "微信公众平台", "文件传输助手"
]
UI_NOISE_KEYWORDS = [
"微信", "Weixin", "WeChat", "聊天信息", "搜索", "更多", "表情", "发送", "Message", "Messages"
]
+346
View File
@@ -0,0 +1,346 @@
import base64
from io import BytesIO
from pathlib import Path
import cv2
import numpy as np
import requests
from PIL import Image
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
from app.infrastructure.service.wechat.config import (
BAIDU_API_KEY,
BAIDU_SECRET_KEY,
OCR_PROVIDER,
RAPID_OCR_CLS_MODEL_PATH,
RAPID_OCR_DET_MODEL_PATH,
RAPID_OCR_REC_MODEL_PATH,
SESSION_NAME_OCR_EXTRA_SCALE,
SESSION_NAME_OCR_SCALE,
)
BAIDU_FALLBACK_ERROR_CODES = {17, 18, 110, 111}
def _runtime_roots() -> list[Path]:
roots: list[Path] = []
meipass = getattr(__import__("sys"), "_MEIPASS", None)
if meipass:
roots.append(Path(meipass))
file_root = Path(__file__).resolve().parents[4]
roots.append(file_root)
cwd = Path.cwd().resolve()
roots.append(cwd)
roots.append(cwd / "resources")
roots.append(cwd / "app")
try:
exe_parent = Path(__import__("sys").executable).resolve().parent
roots.append(exe_parent)
roots.append(exe_parent / "resources")
roots.append(exe_parent / "app")
except Exception:
pass
unique_roots: list[Path] = []
seen = set()
for root in roots:
key = str(root)
if key in seen:
continue
seen.add(key)
unique_roots.append(root)
return unique_roots
def _resolve_project_path(path_str: str) -> str:
path = Path(path_str)
if path.is_absolute():
return str(path)
candidates = [(root / path).resolve() for root in _runtime_roots()]
for candidate in candidates:
if candidate.exists():
return str(candidate)
return str(candidates[0])
class OCRBase:
provider_name = "base"
def recognize(self, image_data, scene="generic", mode="generic"):
raise NotImplementedError
class BaiduOCR(OCRBase):
provider_name = "baidu"
def __init__(self, api_key, secret_key):
self.api_key = api_key
self.secret_key = secret_key
self.access_token = None
self.last_error_code = None
self.last_error_msg = ""
self.get_access_token()
def get_access_token(self):
trace_id = new_trace_id("ocr")
if not self.api_key or not self.secret_key:
log_event("WARNING", "ocr", "ocr.baidu.token", trace_id, "token", "failed", "百度OCR凭据缺失", reason="credential_missing")
return
url = "https://aip.baidubce.com/oauth/2.0/token"
params = {"grant_type": "client_credentials", "client_id": self.api_key, "client_secret": self.secret_key}
try:
response = requests.post(url, params=params, timeout=10)
if response.status_code == 200:
self.access_token = response.json().get("access_token")
if self.access_token:
log_event("INFO", "ocr", "ocr.baidu.token", trace_id, "token", "ok", "百度OCR token获取成功")
else:
log_event("WARNING", "ocr", "ocr.baidu.token", trace_id, "token", "failed", "百度OCR token为空", reason="token_empty")
else:
log_event("WARNING", "ocr", "ocr.baidu.token", trace_id, "token", "failed", "百度OCR token获取失败", reason="http_error", extra={"status_code": response.status_code})
except Exception as e:
log_event("ERROR", "ocr", "ocr.baidu.token", trace_id, "token", "failed", "百度OCR token请求异常", reason="request_error", extra={"error": str(e)})
def _reset_last_error(self):
self.last_error_code = None
self.last_error_msg = ""
def should_fallback_to_rapid(self):
return self.last_error_code in BAIDU_FALLBACK_ERROR_CODES
def recognize(self, image_data, scene="generic", mode="generic"):
trace_id = new_trace_id("ocr")
self._reset_last_error()
if not self.access_token:
self.last_error_msg = "no_access_token"
log_event("WARNING", "ocr", "ocr.baidu.recognize", trace_id, "recognize", "failed", "百度OCR无可用token", reason="no_access_token", extra={"scene": scene})
return []
url = f"https://aip.baidubce.com/rest/2.0/ocr/v1/general_basic?access_token={self.access_token}"
payload = {"image": base64.b64encode(image_data).decode(), "language_type": "CHN_ENG", "detect_direction": "true", "probability": "true"}
try:
response = requests.post(url, data=payload, timeout=10)
if response.status_code == 200:
result = response.json()
if "error_code" in result:
self.last_error_code = result.get("error_code")
self.last_error_msg = result.get("error_msg") or ""
log_event("WARNING", "ocr", "ocr.baidu.recognize", trace_id, "recognize", "failed", "百度OCR返回错误码", reason="baidu_error", extra={"scene": scene, "error_code": self.last_error_code, "error_msg": self.last_error_msg})
return []
if "words_result" in result:
lines = []
for item in result["words_result"]:
text = item.get("words", "")
prob = item.get("probability", {}).get("average", 0.9)
if text and prob > 0.6:
lines.append(text)
log_event("INFO", "ocr", "ocr.baidu.recognize", trace_id, "recognize", "ok", "百度OCR识别完成", extra={"scene": scene, "line_count": len(lines)})
return lines
else:
self.last_error_msg = f"http_{response.status_code}"
log_event("WARNING", "ocr", "ocr.baidu.recognize", trace_id, "recognize", "failed", "百度OCR请求失败", reason="http_error", extra={"scene": scene, "status_code": response.status_code})
except Exception as e:
self.last_error_msg = str(e)
log_event("ERROR", "ocr", "ocr.baidu.recognize", trace_id, "recognize", "failed", "百度OCR请求异常", reason="request_error", extra={"scene": scene, "error": str(e)})
return []
class RapidLocalOCR(OCRBase):
provider_name = "rapid"
def __init__(self):
self.ready = False
self.engine = None
self._init_engine()
def ensure_ready(self):
return self.ready and self.engine is not None
def _init_engine(self):
trace_id = new_trace_id("ocr")
try:
from rapidocr_onnxruntime import RapidOCR
model_paths = {
"det_model_path": _resolve_project_path(RAPID_OCR_DET_MODEL_PATH),
"rec_model_path": _resolve_project_path(RAPID_OCR_REC_MODEL_PATH),
"cls_model_path": _resolve_project_path(RAPID_OCR_CLS_MODEL_PATH),
}
existing_model_paths = {key: value for key, value in model_paths.items() if Path(value).exists()}
if len(existing_model_paths) == len(model_paths):
self.engine = RapidOCR(**existing_model_paths)
log_extra = existing_model_paths
else:
self.engine = RapidOCR()
log_extra = {**model_paths, "missing_models": [value for value in model_paths.values() if not Path(value).exists()]}
self.ready = True
log_event("INFO", "ocr", "ocr.rapid.init", trace_id, "init", "ok", "RapidOCR初始化成功", extra=log_extra)
except Exception as e:
self.ready = False
log_event("WARNING", "ocr", "ocr.rapid.init", trace_id, "init", "failed", "RapidOCR初始化失败", reason="init_error", extra={"error": str(e)})
def recognize(self, image_data, scene="generic", mode="generic"):
trace_id = new_trace_id("ocr")
if not self.ready or self.engine is None:
log_event("WARNING", "ocr", "ocr.rapid.recognize", trace_id, "recognize", "failed", "RapidOCR未就绪", reason="not_ready", extra={"scene": scene})
return []
try:
img_np = np.frombuffer(image_data, dtype=np.uint8)
img = cv2.imdecode(img_np, cv2.IMREAD_COLOR)
if img is None:
log_event("WARNING", "ocr", "ocr.rapid.recognize", trace_id, "recognize", "failed", "RapidOCR图像解码失败", reason="decode_failed", extra={"scene": scene})
return []
result = self.engine(img)
if not result or len(result) < 1:
log_event("INFO", "ocr", "ocr.rapid.recognize", trace_id, "recognize", "ok", "RapidOCR识别结果为空", reason="empty_result", extra={"scene": scene})
return []
rec_res = result[0] or []
lines = []
for item in rec_res:
if not item or len(item) < 2:
continue
text = str(item[1]).strip()
if text:
lines.append(text)
log_event("INFO", "ocr", "ocr.rapid.recognize", trace_id, "recognize", "ok", "RapidOCR识别完成", extra={"scene": scene, "line_count": len(lines)})
return lines
except Exception as e:
log_event("ERROR", "ocr", "ocr.rapid.recognize", trace_id, "recognize", "failed", "RapidOCR识别异常", reason="recognize_error", extra={"scene": scene, "error": str(e)})
return []
class OCRService(OCRBase):
provider_name = "service"
def __init__(self, provider=None):
self.provider_requested = (provider or OCR_PROVIDER or "baidu").strip().lower()
self.baidu_provider = BaiduOCR(BAIDU_API_KEY, BAIDU_SECRET_KEY)
self.rapid_provider = RapidLocalOCR()
self.provider = self._build_provider(self.provider_requested)
def _build_provider(self, provider_name: str):
if provider_name in {"rapid", "rapidocr"}:
return self.rapid_provider
if provider_name in {"baidu", "baiduocr"}:
return self.baidu_provider
if provider_name == "auto":
if self.baidu_provider.access_token:
return self.baidu_provider
if self.rapid_provider.ensure_ready():
return self.rapid_provider
return self.baidu_provider
return self.baidu_provider
def _provider_recognize(self, image_data, scene):
trace_id = new_trace_id("ocr")
lines = self.provider.recognize(image_data, scene=scene)
if self.provider.provider_name != "baidu":
return lines
if lines:
return lines
no_token_fallback = self.baidu_provider.last_error_msg == "no_access_token"
should_fallback = self.baidu_provider.should_fallback_to_rapid() or no_token_fallback
if not should_fallback:
return lines
if not self.rapid_provider.ensure_ready():
log_event("WARNING", "ocr", "ocr.fallback", trace_id, "fallback", "failed", "触发Rapid回退但引擎未就绪", reason="rapid_not_ready", extra={"scene": scene, "baidu_error": self.baidu_provider.last_error_msg or ""})
return lines
rapid_lines = self.rapid_provider.recognize(image_data, scene=f"{scene}_rapid_fallback")
if rapid_lines:
log_event("INFO", "ocr", "ocr.fallback", trace_id, "fallback", "ok", "百度OCR回退Rapid成功", reason="fallback_success", extra={"scene": scene, "line_count": len(rapid_lines)})
else:
log_event("WARNING", "ocr", "ocr.fallback", trace_id, "fallback", "failed", "百度OCR回退Rapid失败", reason="fallback_empty", extra={"scene": scene})
return rapid_lines
def _encode_image(self, image_obj):
buf = BytesIO()
image_obj.save(buf, format="PNG")
return buf.getvalue()
def _normalize_lines(self, lines, min_len=1, exclude=None):
exclude = set(exclude or [])
normalized = []
for line in lines or []:
text = str(line).strip()
if not text:
continue
if len(text) < min_len:
continue
if text in exclude:
continue
normalized.append(text)
return normalized
def _build_session_name_variants(self, image_data):
image = Image.open(BytesIO(image_data)).convert("RGB")
gray = image.convert("L")
base_scale = max(2, int(SESSION_NAME_OCR_SCALE))
extra_scale = max(base_scale, int(SESSION_NAME_OCR_EXTRA_SCALE))
enlarged = gray.resize(
(gray.width * base_scale, gray.height * base_scale),
resample=Image.Resampling.LANCZOS,
)
contrast = cv2.equalizeHist(np.array(enlarged))
binary = cv2.threshold(contrast, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)[1]
binary_inv = cv2.threshold(contrast, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)[1]
extra_enlarged = gray.resize(
(gray.width * extra_scale, gray.height * extra_scale),
resample=Image.Resampling.LANCZOS,
)
extra_contrast = cv2.equalizeHist(np.array(extra_enlarged))
return [
("name_crop", image_data),
(f"name_orig_{base_scale}x", self._encode_image(enlarged)),
(f"name_eq_{base_scale}x", self._encode_image(Image.fromarray(contrast))),
(f"name_bin_{base_scale}x", self._encode_image(Image.fromarray(binary))),
(f"name_bin_inv_{base_scale}x", self._encode_image(Image.fromarray(binary_inv))),
(f"name_orig_{extra_scale}x", self._encode_image(extra_enlarged)),
(f"name_eq_{extra_scale}x", self._encode_image(Image.fromarray(extra_contrast))),
]
def _recognize_session_name(self, image_data, scene):
trace_id = new_trace_id("ocr")
for variant_name, variant_bytes in self._build_session_name_variants(image_data):
lines = self._normalize_lines(
self._provider_recognize(variant_bytes, scene=f"{scene}_{variant_name}"),
min_len=1,
)
if lines:
log_event("INFO", "ocr", "ocr.session_name", trace_id, "recognize", "ok", "会话名识别成功", extra={"scene": scene, "variant": variant_name, "line_count": len(lines)})
return lines
log_event("INFO", "ocr", "ocr.session_name", trace_id, "recognize", "failed", "会话名识别为空", reason="empty_result", extra={"scene": scene})
return []
def _recognize_session_title(self, image_data, scene):
trace_id = new_trace_id("ocr")
lines = self._provider_recognize(image_data, scene=scene)
normalized = self._normalize_lines(lines, min_len=1)
if normalized:
log_event("INFO", "ocr", "ocr.session_title", trace_id, "recognize", "ok", "会话标题识别成功", extra={"scene": scene, "line_count": len(normalized)})
else:
log_event("INFO", "ocr", "ocr.session_title", trace_id, "recognize", "failed", "会话标题识别为空", reason="empty_result", extra={"scene": scene})
return normalized
def recognize_session_name(self, image_data, scene="session_name"):
return self._recognize_session_name(image_data, scene=scene)
def recognize_session_title(self, image_data, scene="session_title"):
return self._recognize_session_title(image_data, scene=scene)
def recognize(self, image_data, scene="generic", mode="generic"):
trace_id = new_trace_id("ocr")
if mode == "session_name":
return self.recognize_session_name(image_data, scene=scene)
if mode == "session_title":
return self.recognize_session_title(image_data, scene=scene)
lines = self._normalize_lines(self._provider_recognize(image_data, scene=scene), min_len=1)
log_event("INFO", "ocr", "ocr.generic", trace_id, "recognize", "ok", "通用OCR识别完成", extra={"scene": scene, "line_count": len(lines)})
return lines
@@ -0,0 +1,314 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict
import cv2
import numpy as np
from PIL import ImageGrab
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
from app.infrastructure.service.wechat.config import (
CHAT_CAPTURE_HEIGHT,
CHAT_CAPTURE_LEFT_OFFSET,
CHAT_CAPTURE_TOP_OFFSET,
CHAT_CAPTURE_WIDTH,
CONTACT_LIST_BOTTOM_OFFSET,
CONTACT_LIST_LEFT_OFFSET,
CONTACT_LIST_TOP_OFFSET,
CONTACT_ROW_WIDTH,
SESSION_NAME_HEIGHT,
SESSION_NAME_LEFT_OFFSET,
SESSION_NAME_TOP_OFFSET,
SESSION_NAME_WIDTH,
TITLE_OCR_AREA_HEIGHT,
TITLE_OCR_AREA_LEFT_OFFSET,
TITLE_OCR_AREA_TOP_OFFSET,
TITLE_OCR_AREA_WIDTH,
)
@dataclass
class CaptureBox:
left: int
top: int
right: int
bottom: int
@property
def width(self) -> int:
return self.right - self.left
@property
def height(self) -> int:
return self.bottom - self.top
def as_tuple(self):
return (self.left, self.top, self.right, self.bottom)
def as_dict(self) -> Dict[str, int]:
return {
"left": self.left,
"top": self.top,
"right": self.right,
"bottom": self.bottom,
"width": self.width,
"height": self.height,
}
class ScreenshotService:
def build_box(self, left: int, top: int, width: int, height: int) -> CaptureBox:
return CaptureBox(
left=int(left),
top=int(top),
right=int(left + width),
bottom=int(top + height),
)
def build_box_from_window(self, window_rect: dict, left_offset: int, top_offset: int, width: int, height: int) -> CaptureBox:
return self.build_box(
left=window_rect["left"] + int(left_offset),
top=window_rect["top"] + int(top_offset),
width=int(width),
height=int(height),
)
def build_contact_list_box(self, window_rect: dict, left_offset: int, top_offset: int, width: int, bottom_offset: int) -> CaptureBox:
left = window_rect["left"] + int(left_offset)
top = window_rect["top"] + int(top_offset)
right = left + int(width)
bottom = window_rect["bottom"] - int(bottom_offset)
return CaptureBox(left=left, top=top, right=right, bottom=bottom)
def is_valid_box(self, box: CaptureBox) -> bool:
return box.right > box.left and box.bottom > box.top
def capture_box(self, left: int, top: int, width: int, height: int):
box = self.build_box(left, top, width, height)
if not self.is_valid_box(box):
raise ValueError(f"invalid capture box: {box.as_dict()}")
return ImageGrab.grab(bbox=box.as_tuple())
def capture_from_window(self, window_rect: dict, left_offset: int, top_offset: int, width: int, height: int):
box = self.build_box_from_window(window_rect, left_offset, top_offset, width, height)
if not self.is_valid_box(box):
raise ValueError(f"invalid window capture box: {box.as_dict()}")
return ImageGrab.grab(bbox=box.as_tuple())
def capture_contact_list(self, window_rect: dict, left_offset: int, top_offset: int, width: int, bottom_offset: int):
box = self.build_contact_list_box(window_rect, left_offset, top_offset, width, bottom_offset)
if not self.is_valid_box(box):
raise ValueError(f"invalid contact list box: {box.as_dict()}")
return ImageGrab.grab(bbox=box.as_tuple())
def get_contact_list_box(self, window_rect: dict) -> CaptureBox:
return self.build_contact_list_box(
window_rect,
left_offset=CONTACT_LIST_LEFT_OFFSET,
top_offset=CONTACT_LIST_TOP_OFFSET,
width=CONTACT_ROW_WIDTH,
bottom_offset=CONTACT_LIST_BOTTOM_OFFSET,
)
def capture_contact_list_default(self, window_rect: dict):
trace_id = new_trace_id("capture")
box = self.get_contact_list_box(window_rect)
log_event("INFO", "capture", "capture.contact_list", trace_id, "capture", "ok", "截图会话列表区域", extra=box.as_dict())
return self.capture_contact_list(
window_rect,
left_offset=CONTACT_LIST_LEFT_OFFSET,
top_offset=CONTACT_LIST_TOP_OFFSET,
width=CONTACT_ROW_WIDTH,
bottom_offset=CONTACT_LIST_BOTTOM_OFFSET,
)
def get_session_title_box(self, window_rect: dict) -> CaptureBox:
return self.build_box_from_window(
window_rect,
left_offset=TITLE_OCR_AREA_LEFT_OFFSET,
top_offset=TITLE_OCR_AREA_TOP_OFFSET,
width=TITLE_OCR_AREA_WIDTH,
height=TITLE_OCR_AREA_HEIGHT,
)
def capture_session_title(self, window_rect: dict):
trace_id = new_trace_id("capture")
box = self.get_session_title_box(window_rect)
log_event("INFO", "capture", "capture.session_title", trace_id, "capture", "ok", "截图会话标题区域", extra=box.as_dict())
return self.capture_area_from_box(box)
def get_chat_capture_box(self, window_rect: dict) -> CaptureBox:
base_height = max(120, CHAT_CAPTURE_HEIGHT)
max_height = max(base_height, window_rect["height"] - CHAT_CAPTURE_TOP_OFFSET)
return self.build_box_from_window(
window_rect,
left_offset=CHAT_CAPTURE_LEFT_OFFSET,
top_offset=CHAT_CAPTURE_TOP_OFFSET,
width=CHAT_CAPTURE_WIDTH,
height=max_height,
)
def capture_chat_area(self, window_rect: dict):
trace_id = new_trace_id("capture")
box = self.get_chat_capture_box(window_rect)
image = self.capture_area_from_box(box)
chat_bottom = self._detect_chat_bottom_by_binary_merge(image)
if chat_bottom is not None:
image = image.crop((0, 0, image.size[0], chat_bottom))
extra = box.as_dict()
extra["dynamic_bottom"] = chat_bottom or ""
extra["final_width"] = image.size[0]
extra["final_height"] = image.size[1]
log_event("INFO", "capture", "capture.chat_area", trace_id, "capture", "ok", "截图聊天区域", extra=extra)
return image
def crop_session_name(self, row_img):
return self.crop_from_image(
row_img,
left=SESSION_NAME_LEFT_OFFSET,
top=SESSION_NAME_TOP_OFFSET,
width=SESSION_NAME_WIDTH,
height=SESSION_NAME_HEIGHT,
)
def capture_area_from_box(self, box: CaptureBox):
if not self.is_valid_box(box):
raise ValueError(f"invalid capture box: {box.as_dict()}")
return ImageGrab.grab(bbox=box.as_tuple())
def _build_merged_binary_array(self, image_obj):
arr = np.array(image_obj.convert("RGB"))
gray = cv2.cvtColor(arr, cv2.COLOR_RGB2GRAY)
blurred = cv2.GaussianBlur(gray, (5, 5), 0)
_, binary_inv = cv2.threshold(blurred, 248, 255, cv2.THRESH_BINARY_INV)
adaptive_inv = cv2.adaptiveThreshold(
blurred,
255,
cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
cv2.THRESH_BINARY_INV,
13,
1,
)
merged = cv2.bitwise_or(binary_inv, adaptive_inv)
kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (5, 1))
merged = cv2.morphologyEx(merged, cv2.MORPH_CLOSE, kernel, iterations=1)
merged = cv2.morphologyEx(merged, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (2, 2)))
return merged
def _detect_chat_bottom_by_binary_merge(self, image_obj) -> int | None:
if image_obj is None:
return None
merged = self._build_merged_binary_array(image_obj)
img_h, img_w = merged.shape[:2]
start_y = max(0, int(img_h * 0.55))
roi = merged[start_y:, :]
if roi.size == 0:
return None
row_density = (roi > 0).mean(axis=1)
min_run = max(28, int(img_h * 0.045))
dense_limit = 0.018
run_start = None
candidates = []
for idx, density in enumerate(row_density.tolist() + [1.0]):
is_blank = density <= dense_limit
if is_blank and run_start is None:
run_start = idx
continue
if is_blank:
continue
if run_start is not None:
run_end = idx
if run_end - run_start >= min_run:
top = start_y + run_start
bottom = start_y + run_end
if top >= img_h * 0.58 and bottom <= img_h - 8:
candidates.append((top, bottom))
run_start = None
if not candidates:
return self._detect_chat_bottom(image_obj)
top, _ = candidates[-1]
bottom = max(120, int(top - 4))
if bottom >= img_h - 20:
return None
return bottom
def crop_from_image(self, image_obj, left: int, top: int, width: int, height: int):
if image_obj is None:
return None
img_w, img_h = image_obj.size
crop_left = min(max(0, int(left)), img_w)
crop_top = min(max(0, int(top)), img_h)
crop_right = min(img_w, crop_left + max(1, int(width)))
crop_bottom = min(img_h, crop_top + max(1, int(height)))
if crop_right <= crop_left or crop_bottom <= crop_top:
return None
return image_obj.crop((crop_left, crop_top, crop_right, crop_bottom))
def _detect_chat_bottom(self, image_obj) -> int | None:
if image_obj is None:
return None
img_rgb = np.array(image_obj.convert("RGB"))
if img_rgb.size == 0:
return None
gray = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2GRAY)
img_h, img_w = gray.shape[:2]
start_y = max(0, int(img_h * 0.55))
focus = gray[start_y:, :]
if focus.size == 0:
return None
row_mean = focus.mean(axis=1)
row_std = focus.std(axis=1)
bright_mask = (row_mean >= 242) & (row_std <= 18)
run = self._find_last_run(bright_mask, min_len=max(18, int(img_h * 0.035)))
candidate_y = None
if run is not None:
run_top, run_bottom = run
candidate_y = start_y + run_top
edge_img = cv2.Canny(focus, 40, 120)
edge_strength = edge_img.mean(axis=1)
if len(row_mean) >= 2:
transition = np.abs(np.diff(row_mean, prepend=row_mean[0]))
else:
transition = np.zeros_like(row_mean)
score = edge_strength * 1.8 + transition * 2.4
score[: max(8, int(len(score) * 0.15))] = 0
if candidate_y is not None:
local_limit = max(0, candidate_y - start_y + 4)
score[local_limit:] = 0
best_idx = int(np.argmax(score)) if score.size else -1
best_score = float(score[best_idx]) if best_idx >= 0 else 0.0
edge_candidate = None
if best_idx >= 0 and best_score >= 12.0:
edge_candidate = start_y + best_idx
final_y = None
if candidate_y is not None and edge_candidate is not None:
if abs(candidate_y - edge_candidate) <= 28:
final_y = min(candidate_y, edge_candidate)
else:
final_y = candidate_y
else:
final_y = candidate_y if candidate_y is not None else edge_candidate
if final_y is None:
return None
final_y = max(120, min(img_h, int(final_y - 6)))
if final_y >= img_h - 20:
return None
return final_y
def _find_last_run(self, mask: np.ndarray, min_len: int) -> tuple[int, int] | None:
run_start = None
best = None
for idx, flag in enumerate(mask.tolist() + [False]):
if flag and run_start is None:
run_start = idx
continue
if flag:
continue
if run_start is None:
continue
run_len = idx - run_start
if run_len >= min_len:
best = (run_start, idx)
run_start = None
return best
@@ -0,0 +1,288 @@
from __future__ import annotations
from dataclasses import dataclass
from io import BytesIO
import json
import os
from typing import Callable
import cv2
import numpy as np
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
from app.infrastructure.service.wechat.chat_snapshot_analyzer import analyze_pil_image
from app.infrastructure.service.wechat.unread_session_analyzer import UnreadSessionAnalyzer
from app.infrastructure.service.wechat.config import (
BLOCKED_SESSION_KEYWORDS,
CONTACT_ROW_HEIGHT,
OCR_SAVE_DIR,
OCR_SAVE_IMAGES,
SESSION_NAME_HEIGHT,
SESSION_NAME_LEFT_OFFSET,
SESSION_NAME_TOP_OFFSET,
SESSION_NAME_WIDTH,
UI_NOISE_KEYWORDS,
)
# 会话扫描结果数据类,包含所有会话列表和未读会话列表
@dataclass
class SessionScanResult:
sessions: list[dict]
unread_sessions: list[dict]
# 聊天快照分析结果数据类
@dataclass
class ChatAnalyzeResult:
ok: bool
file_name: str
latest_text: str
confidence: str | float
bubble_side: str
screenshot_path: str
# 微信会话服务类,处理会话列表扫描、红点检测和聊天截图分析
class WechatSessionService:
def __init__(self, screenshot_service, ocr_service, save_debug_image: Callable | None = None):
self.screenshot = screenshot_service
self.ocr = ocr_service
self.save_debug_image = save_debug_image
self._session_title_cache = {"value": "", "ts": 0.0}
self.unread_analyzer = UnreadSessionAnalyzer()
log_event("INFO", "bot", "bot.session_service.init", new_trace_id("bot"), "init", "ok", "会话服务初始化完成")
# 根据窗口矩形计算会话列表区域的位置
def get_contact_list_rect(self, window_rect):
box = self.screenshot.get_contact_list_box(window_rect)
return {
'left': box.left,
'top': box.top,
'right': box.right,
'bottom': box.bottom,
}
# 从会话行图片中裁剪出会话名称区域
def extract_session_name_image(self, row_img):
return self.screenshot.crop_session_name(row_img)
# 检测会话列表中的所有红点位置(红色圆点表示未读消息)
def detect_red_dots(self, window_rect):
contact_rect = self.get_contact_list_rect(window_rect)
screenshot = self.screenshot.capture_contact_list_default(window_rect)
return self.unread_analyzer.detect_red_dots(contact_rect, screenshot)
# 检测单行会话图片中是否有未读红点标记(严格模式)
def row_has_red_dot(self, row_img, relaxed=False):
return self.unread_analyzer.row_has_red_dot(row_img, relaxed=relaxed)
# 检测单行会话图片中是否有未读红点标记(宽松模式)
def row_has_red_dot_weak(self, row_img):
return self.unread_analyzer.row_has_red_dot_weak(row_img)
# 扫描所有会话行,识别哪些有未读消息标记
def get_all_sessions_with_unread(self, window_rect, round_count):
trace_id = new_trace_id("bot")
contact_rect = self.get_contact_list_rect(window_rect)
screenshot = self.screenshot.capture_contact_list_default(window_rect)
sessions, unread_sessions = self.unread_analyzer.get_all_sessions_with_unread(
contact_rect=contact_rect,
screenshot=screenshot,
round_count=round_count,
save_debug_image=lambda image_obj, filename: self._save_debug_image(image_obj, filename),
)
self._save_session_scan_debug(round_count=round_count, sessions=sessions, unread_sessions=unread_sessions, contact_rect=contact_rect)
log_event("INFO", "bot", "bot.session_scan", trace_id, "scan", "ok", "会话扫描完成", extra={"round": int(round_count), "total": len(sessions), "unread": len(unread_sessions)})
return SessionScanResult(sessions=sessions, unread_sessions=unread_sessions)
# 标准化文本用于匹配:去除空格并转为小写
def normalize_match_text(self, text):
if not text:
return ""
text = str(text).strip().lower()
return "".join(ch for ch in text if not ch.isspace())
# 生成会话屏蔽关键字的唯一标识key,用于缓存比对
def make_block_key(self, text):
normalized = self.normalize_match_text(text)
if not normalized:
return ""
return f"title:{normalized}"
# 重置当前会话标题缓存
def reset_session_title_cache(self):
self._session_title_cache = {"value": "", "ts": 0.0}
# 通过OCR识别当前会话窗口的标题文字
def get_session_title_by_ocr(self, window_rect):
trace_id = new_trace_id("bot")
try:
if not window_rect:
return ""
area_name = "main"
screenshot = self.screenshot.capture_session_title(window_rect)
img_bytes = BytesIO()
screenshot.save(img_bytes, format='PNG')
valid = self.ocr.recognize_session_title(img_bytes.getvalue(), scene=f"session_title_{area_name}")
if valid:
title = valid[0]
log_event("INFO", "bot", "bot.session_title", trace_id, "ocr", "ok", "会话标题识别成功", extra={"title": title})
return title
log_event("INFO", "bot", "bot.session_title", trace_id, "ocr", "failed", "会话标题识别为空", reason="empty_result")
return ""
except Exception as e:
log_event("ERROR", "bot", "bot.session_title", trace_id, "ocr", "failed", "会话标题识别异常", reason="ocr_error", extra={"error": str(e)})
return ""
# 获取当前会话标题,优先使用缓存避免频繁OCR调用
def get_current_session_title(self, window_rect):
try:
import time
now_ts = time.time()
cached_title = (self._session_title_cache.get("value") or "").strip()
cached_ts = float(self._session_title_cache.get("ts") or 0.0)
if cached_title and now_ts - cached_ts <= 1.2:
return cached_title
title = (self.get_session_title_by_ocr(window_rect) or "").strip()
if title and title not in UI_NOISE_KEYWORDS:
self._session_title_cache = {"value": title, "ts": now_ts}
return title
except Exception as e:
return ""
# 判断当前选中的会话是否应被跳过(点击后标题检查阶段)
def should_skip_current_session(self, window_rect, session, blocked_row_cache, save_blocked_row_cache: Callable):
title = self.get_current_session_title(window_rect)
block_key = self.make_block_key(title)
if block_key and block_key in blocked_row_cache:
return True
normalized_title = self.normalize_match_text(title)
for keyword in BLOCKED_SESSION_KEYWORDS:
if self.normalize_match_text(keyword) in normalized_title:
if block_key:
blocked_row_cache[block_key] = title or keyword
save_blocked_row_cache()
return True
return False
# 比较两个会话名称是否匹配(考虑模糊匹配和大小写)
def is_same_session(self, expected_session, current_session):
expected = self.normalize_match_text(expected_session)
current = self.normalize_match_text(current_session)
if not expected or not current:
return False
return expected in current or current in expected
# 根据OCR识别结果判断会话列表中的会话是否应被跳过
def should_skip_session_by_ocr(self, session, blocked_row_cache, save_blocked_row_cache: Callable):
image_obj = session.get('row_img')
if image_obj is None:
return False
try:
name_img = self.extract_session_name_image(image_obj)
if name_img is None:
return False
crop_box = {
'left': SESSION_NAME_LEFT_OFFSET,
'top': SESSION_NAME_TOP_OFFSET,
'width': SESSION_NAME_WIDTH,
'height': SESSION_NAME_HEIGHT,
'row_w': image_obj.size[0],
'row_h': image_obj.size[1],
'crop_w': name_img.size[0],
'crop_h': name_img.size[1],
}
if OCR_SAVE_IMAGES:
file_name = f"row_{int(session.get('row_idx', 0)):03d}_name_raw.png"
self._save_debug_image(name_img, os.path.join('sessions', 'name_ocr', file_name))
img_bytes = BytesIO()
name_img.save(img_bytes, format='PNG')
lines = self.ocr.recognize_session_name(img_bytes.getvalue(), scene=f"session_row_{session.get('row_idx')}")
line_text = ' '.join(lines)
session['list_ocr_title'] = line_text
normalized_text = self.normalize_match_text(line_text)
block_key = self.make_block_key(line_text)
if block_key and block_key in blocked_row_cache:
return True
for keyword in BLOCKED_SESSION_KEYWORDS:
if self.normalize_match_text(keyword) in normalized_text:
if block_key:
blocked_row_cache[block_key] = line_text or keyword
save_blocked_row_cache()
return True
return False
except Exception as e:
return False
# 分析点击后的聊天区域截图,提取最新消息文本并返回分析结果
def analyze_clicked_session(self, window_rect, round_count, row_idx):
trace_id = new_trace_id("bot")
chat_box = self.screenshot.get_chat_capture_box(window_rect)
if not self.screenshot.is_valid_box(chat_box):
log_event("WARNING", "bot", "bot.chat_analyze", trace_id, "capture", "failed", "聊天区截图区域无效", reason="invalid_box")
return ChatAnalyzeResult(ok=False, file_name='', latest_text='', confidence='', bubble_side='', screenshot_path='')
screenshot = self.screenshot.capture_chat_area(window_rect)
file_name = f"round_{round_count:04d}_row_{row_idx:03d}_chat.png"
rel_path = os.path.join('sessions', 'clicked', file_name)
self._save_debug_image(screenshot, rel_path)
result = analyze_pil_image(screenshot, stem=os.path.splitext(file_name)[0], file_name=file_name)
latest_text = (getattr(result, 'latest_text', None) or '').strip()
confidence = getattr(result, 'confidence', '')
bubble_side = getattr(result, 'bubble_side', '')
log_event("INFO", "bot", "bot.chat_analyze", trace_id, "analyze", "ok", "聊天截图分析完成", extra={"round": int(round_count), "row_idx": int(row_idx), "has_text": bool(latest_text), "bubble_side": bubble_side or "", "confidence": confidence})
return ChatAnalyzeResult(
ok=bool(latest_text),
file_name=file_name,
latest_text=latest_text,
confidence=confidence,
bubble_side=bubble_side,
screenshot_path=rel_path,
)
# 保存会话扫描调试数据(类似聊天分析输出 result.json
def _save_session_scan_debug(self, round_count: int, sessions: list[dict], unread_sessions: list[dict], contact_rect: dict):
if not OCR_SAVE_IMAGES:
return
try:
debug_dir = os.path.join(OCR_SAVE_DIR, 'sessions', 'scan_debug')
os.makedirs(debug_dir, exist_ok=True)
file_name = f"round_{round_count:04d}_scan.json"
file_path = os.path.join(debug_dir, file_name)
rows = []
for session in sessions:
rows.append({
'row_idx': session.get('row_idx'),
'has_red_dot': bool(session.get('has_red_dot')),
'has_red_by_global': bool(session.get('has_red_by_global')),
'has_red_by_row': bool(session.get('has_red_by_row')),
'has_red_by_row_weak': bool(session.get('has_red_by_row_weak')),
'click_x': session.get('click_x'),
'click_y': session.get('click_y'),
'list_ocr_title': session.get('list_ocr_title', ''),
})
payload = {
'round': int(round_count),
'contact_rect': contact_rect,
'total_sessions': len(sessions),
'unread_count': len(unread_sessions),
'unread_rows': [s.get('row_idx') for s in unread_sessions],
'rows': rows,
}
with open(file_path, 'w', encoding='utf-8') as f:
json.dump(payload, f, ensure_ascii=False, indent=2)
except Exception as e:
pass
# 保存调试图片的内部方法
def _save_debug_image(self, image_obj, filename):
if not self.save_debug_image:
return
self.save_debug_image(image_obj, filename)
@@ -0,0 +1,258 @@
from __future__ import annotations
import cv2
import numpy as np
from app.infrastructure.service.logging.log_service import log_event, new_trace_id
from app.infrastructure.service.wechat.config import CONTACT_ROW_HEIGHT
class UnreadSessionAnalyzer:
def __init__(self):
pass
def detect_red_dots(self, contact_rect: dict, screenshot) -> list[dict]:
trace_id = new_trace_id("bot")
try:
img_np = np.array(screenshot)
hsv = cv2.cvtColor(cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR), cv2.COLOR_BGR2HSV)
mask = cv2.inRange(hsv, np.array([0, 80, 80]), np.array([12, 255, 255])) + cv2.inRange(hsv, np.array([168, 80, 80]), np.array([180, 255, 255]))
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
contact_width = contact_rect['right'] - contact_rect['left']
red_dots_raw = []
for contour in contours:
area = cv2.contourArea(contour)
if 12 < area < 220:
perimeter = cv2.arcLength(contour, True)
if perimeter <= 0:
continue
circularity = 4 * np.pi * area / (perimeter * perimeter)
if circularity <= 0.45:
continue
moments = cv2.moments(contour)
if moments['m00'] == 0:
continue
cx = int(moments['m10'] / moments['m00'])
cy = int(moments['m01'] / moments['m00'])
if cx > contact_width * 0.1:
red_dots_raw.append({'x': contact_rect['left'] + cx, 'y': contact_rect['top'] + cy, 'rel_y': cy})
snapped_map = {}
for dot in red_dots_raw:
row_idx = int(round((dot['y'] - contact_rect['top']) / max(1, CONTACT_ROW_HEIGHT)))
snapped_y = int(contact_rect['top'] + row_idx * CONTACT_ROW_HEIGHT + CONTACT_ROW_HEIGHT // 2)
if row_idx not in snapped_map:
snapped_map[row_idx] = {'x': dot['x'], 'y': snapped_y, 'row_idx': row_idx}
red_dots_final = sorted(snapped_map.values(), key=lambda d: d['y'])
log_event("INFO", "bot", "bot.unread.detect", trace_id, "detect", "ok", "红点检测完成", extra={"dot_count": len(red_dots_final)})
return red_dots_final
except Exception as e:
log_event("ERROR", "bot", "bot.unread.detect", trace_id, "detect", "failed", "红点检测异常", reason="detect_error", extra={"error": str(e)})
return []
def row_has_red_dot(self, row_img, relaxed: bool = False) -> bool:
try:
row_np = np.array(row_img)
h, w = row_np.shape[:2]
if h < 30 or w < 100:
return False
margin_left = max(6, int(w * 0.012))
avatar_size = int(h * 0.72)
avatar_y = (h - avatar_size) // 2
avatar_x = margin_left
avatar_cx = avatar_x + avatar_size / 2.0
avatar_cy = avatar_y + avatar_size / 2.0
avatar_r = avatar_size * 0.50
probe_x1 = avatar_x + int(avatar_size * 0.42)
probe_y1 = max(0, avatar_y - int(avatar_size * 0.10))
probe_x2 = min(w, avatar_x + int(avatar_size * 1.00))
probe_y2 = min(h, avatar_y + int(avatar_size * 0.36))
if probe_x2 <= probe_x1 or probe_y2 <= probe_y1:
return False
probe = row_np[probe_y1:probe_y2, probe_x1:probe_x2]
if probe.size == 0:
return False
probe_hsv = cv2.cvtColor(probe, cv2.COLOR_RGB2HSV)
mask1 = cv2.inRange(probe_hsv, np.array([0, 115, 125]), np.array([12, 255, 255]))
mask2 = cv2.inRange(probe_hsv, np.array([168, 115, 125]), np.array([180, 255, 255]))
mask = cv2.bitwise_or(mask1, mask2)
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3))
mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)
mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel)
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
self.debug_log(f"row_red w={w} h={h} candidates=0")
return False
candidates = []
pw = probe_x2 - probe_x1
ph = probe_y2 - probe_y1
for cnt in contours:
area = cv2.contourArea(cnt)
if not (8 <= area <= 260):
continue
x, y, cw, ch = cv2.boundingRect(cnt)
peri = cv2.arcLength(cnt, True)
if peri <= 0:
continue
circ = 4 * np.pi * area / (peri * peri)
ar = max(cw, ch) / max(1, min(cw, ch))
cx = x + cw / 2.0
cy = y + ch / 2.0
gx = probe_x1 + cx
gy = probe_y1 + cy
in_upper_right = cx > pw * 0.30 and cx < pw * 0.90 and cy < ph * 0.68
near_avatar_corner = (
gx >= avatar_x + avatar_size * 0.66 and
gx <= avatar_x + avatar_size * 1.00 and
gy >= avatar_y - avatar_size * 0.06 and
gy <= avatar_y + avatar_size * 0.24
)
if not (in_upper_right and near_avatar_corner):
continue
comp_mask = np.zeros(mask.shape, dtype=np.uint8)
cv2.drawContours(comp_mask, [cnt], -1, 255, thickness=-1)
ys, xs = np.where(comp_mask > 0)
if len(xs) == 0:
continue
global_xs = xs + probe_x1
global_ys = ys + probe_y1
d2 = (global_xs - avatar_cx) ** 2 + (global_ys - avatar_cy) ** 2
outside_ratio = float(np.count_nonzero(d2 > (avatar_r * 0.92) ** 2)) / len(d2)
min_area = 10 if relaxed else 14
min_small_outside = 0.18 if relaxed else 0.25
min_small_circ = 0.72 if relaxed else 0.82
min_match_score = 7 if relaxed else 8
min_match_outside = 0.10 if relaxed else 0.15
if area < min_area:
continue
if area < 20 and outside_ratio < min_small_outside:
continue
if gy > avatar_y + avatar_size * 0.24:
continue
if area < 20 and circ < min_small_circ:
continue
if area < 20 and ar > 1.20:
continue
if area >= 120:
shape_ok = circ > 0.26 and ar < 2.6
elif area >= 28:
shape_ok = circ > 0.45 and ar < 1.9
else:
shape_ok = circ > 0.82 and ar <= 1.20 and outside_ratio >= 0.25
if not shape_ok:
continue
white_ratio = 0.0
if cw >= 7 and ch >= 7:
inner = probe[max(0, y):min(probe.shape[0], y + ch), max(0, x):min(probe.shape[1], x + cw)]
if inner.size > 0:
gray = cv2.cvtColor(inner, cv2.COLOR_RGB2GRAY)
white_ratio = np.count_nonzero(gray > 190) / gray.size
score = 3
if area >= 14:
score += 2
if circ > 0.85:
score += 2
elif circ > 0.70:
score += 1
if ar <= 1.15:
score += 2
elif ar <= 1.35:
score += 1
if outside_ratio >= 0.35:
score += 4
elif outside_ratio >= 0.25:
score += 3
elif outside_ratio >= 0.15:
score += 1
if 0.05 <= white_ratio <= 0.60:
score += 1
candidates.append({
'score': score,
'area': area,
'circ': circ,
'ar': ar,
'outside_ratio': outside_ratio,
'white_ratio': white_ratio,
'center': (gx, gy),
'bbox': (probe_x1 + x, probe_y1 + y, cw, ch),
'min_match_score': min_match_score,
'min_match_outside': min_match_outside,
})
if not candidates:
return False
best = max(candidates, key=lambda x: x['score'])
matched = best['score'] >= best['min_match_score'] and best['outside_ratio'] >= best['min_match_outside']
return matched
except Exception as e:
return False
def row_has_red_dot_weak(self, row_img) -> bool:
return self.row_has_red_dot(row_img, relaxed=True)
def get_all_sessions_with_unread(self, contact_rect: dict, screenshot, round_count: int, save_debug_image: Callable | None = None) -> tuple[list[dict], list[dict]]:
trace_id = new_trace_id("bot")
red_dots = self.detect_red_dots(contact_rect, screenshot)
red_y_list = [dot['y'] for dot in red_dots]
row_count = max(1, int((contact_rect['bottom'] - contact_rect['top']) / max(1, CONTACT_ROW_HEIGHT)))
sessions = []
for row_idx in range(row_count):
top = int(row_idx * CONTACT_ROW_HEIGHT)
bottom = int(min((row_idx + 1) * CONTACT_ROW_HEIGHT, screenshot.height))
if bottom <= top:
continue
row_img = screenshot.crop((0, top, screenshot.width, bottom))
center_y = int(contact_rect['top'] + row_idx * CONTACT_ROW_HEIGHT + CONTACT_ROW_HEIGHT // 2)
has_red_by_global = any(abs(center_y - y) <= max(7, CONTACT_ROW_HEIGHT // 4) for y in red_y_list)
has_red_by_row = self.row_has_red_dot(row_img)
has_red_by_row_weak = self.row_has_red_dot_weak(row_img) if has_red_by_global and not has_red_by_row else has_red_by_row
has_red = has_red_by_row or (has_red_by_global and has_red_by_row_weak)
row_name = f"round_{round_count:04d}_row_{row_idx:03d}.png"
if save_debug_image:
save_debug_image(row_img, f"sessions/all/{row_name}")
if has_red:
save_debug_image(row_img, f"sessions/unread/{row_name}")
sessions.append({
'row_idx': row_idx,
'has_red_dot': has_red,
'has_red_by_global': has_red_by_global,
'has_red_by_row': has_red_by_row,
'has_red_by_row_weak': has_red_by_row_weak,
'click_x': int((contact_rect['left'] + contact_rect['right']) // 2),
'click_y': center_y,
'row_img': row_img.copy(),
})
unread_sessions = [s for s in sessions if s['has_red_dot']]
global_hits = sum(1 for s in sessions if s['has_red_by_global'])
row_hits = sum(1 for s in sessions if s['has_red_by_row'])
row_weak_hits = sum(1 for s in sessions if s['has_red_by_row_weak'])
log_event("INFO", "bot", "bot.unread.scan", trace_id, "scan", "ok", "未读会话扫描完成", extra={"round": int(round_count), "rows": len(sessions), "unread": len(unread_sessions), "global_hits": global_hits, "row_hits": row_hits, "row_weak_hits": row_weak_hits})
return sessions, unread_sessions