# -*- coding: utf-8 -*-
"""全量 OCR 脚本：图片 → 涨停复盘 CSV（RapidOCR 本地识别）
用法: python ocr_zt.py <图片路径> <输出CSV路径> [段数] [统计行数]
特性:
  1. 重叠切段（每段上下重叠 60px）放大 2 倍识别，消除切段边界截断
  2. 跨段去重（同文本 + y差<30 + x差<80）
  3. 水印/噪声过滤（韭研水印、页脚 URL、关键词区碎片），保留板数列与数值列小框
  4. 板块分组标题独立提取（形如 "AI硬件*11"），按 y 归属每行
  5. 行聚类 + 列边界分配（板块/板数/代码/个股/涨停时间/流通市值/成交额/涨停关键词）
  6. 同列归一化去重（忽略 +、空格、括号差异，防重叠残留整串重复）
  7. 全表质量校验（代码/时间/数值/板数/关键词格式）+ 行数核对
  8. 中间结果 boxes.json 落盘，供定点复核脚本复查
"""
import os, sys, time, json, re, csv
from PIL import Image
from collections import Counter

SRC = sys.argv[1]
OUT = sys.argv[2]
N = int(sys.argv[3]) if len(sys.argv) > 3 else 8
EXPECT = int(sys.argv[4]) if len(sys.argv) > 4 else 0
OVERLAP = 60
JSON_PATH = os.path.join(os.path.dirname(OUT), "_boxes.json")

t_start = time.time()
im = Image.open(SRC)
W, H = im.size
print(f"[1/6] 图片 {W}x{H}，重叠切段 {N} 段 x{OVERLAP}px")

from rapidocr_onnxruntime import RapidOCR
engine = RapidOCR()
t0 = time.time()
boxes = []
for i in range(N):
    y0 = max(0, i * H // N - OVERLAP)
    y1 = min(H, (i + 1) * H // N + OVERLAP)
    crop = im.crop((0, y0, W, y1))
    crop = crop.resize((crop.width * 2, crop.height * 2), Image.LANCZOS)
    tmp = os.path.join(os.path.dirname(SRC), "_o.png"); crop.save(tmp)
    res, _ = engine(tmp)
    for box, text, score in res:
        xs = [p[0] for p in box]; ys = [p[1] for p in box]
        boxes.append({"y0": y0 + min(ys) / 2, "x0": min(xs) / 2,
                      "y1": y0 + max(ys) / 2, "x1": max(xs) / 2,
                      "text": text.strip(), "score": round(float(score), 3)})
    os.remove(tmp)
print(f"[2/6] OCR 完成 {time.time()-t0:.1f}s，原始框 {len(boxes)}")

# 去重（重叠区同一行被两段识别；必须比较 x，否则同行的板数"1"与成交额"1"互相误杀）
dedup = []
for b in sorted(boxes, key=lambda b: -b["score"]):
    if any(k["text"] == b["text"] and abs(k["y0"] - b["y0"]) < 30 and abs(k["x0"] - b["x0"]) < 80 for k in dedup):
        continue
    dedup.append(b)
print(f"      去重后 {len(dedup)}")

# 噪声过滤 + 板块标题
NOISE = re.compile(r"(韭研|jiuyangongshe|jiuyangongsh|jiuy|ngongs|gongshe|开公社|公社|www\.|\.com|Zngcngs|Zngongs|com|非研|N$|om$|\.co|查看详细版|统计数据|不含)")
def is_noise(d):
    if d["x0"] > 2000 or d["x1"] < 30:
        return True
    if d["x0"] < 200:  # 板数列小框保留
        return False
    if d["x0"] > 1135 and d["x1"] - d["x0"] < 45 and d["y1"] - d["y0"] < 30:  # 关键词区右侧碎片
        return True
    if NOISE.search(d["text"]):
        return True
    return False

sector_re = re.compile(r"^([\u4e00-\u9fa5A-Za-z0-9 ]+)\*(\d+)$")
sectors = []
clean = []
for d in dedup:
    m = sector_re.match(d["text"].strip())
    if m and 1000 < (d["x0"] + d["x1"]) / 2 < 1350 and d["y0"] > 230:
        sectors.append({"y": d["y0"], "name": d["text"]})
        continue
    if is_noise(d):
        continue
    clean.append(d)
sectors.sort(key=lambda s: s["y"])

def sector_for(y):
    cur = ""
    for s in sectors:
        if s["y"] < y:
            cur = s["name"]
        else:
            break
    return cur

# 行聚类
clean.sort(key=lambda d: d["y0"])
rows = []
cur = []
last = None
for d in clean:
    if last is None or d["y0"] - last <= 25:
        cur.append(d)
    else:
        rows.append(cur); cur = [d]
    last = d["y0"]
if cur:
    rows.append(cur)

def col_of(cx):
    if cx < 195: return "板数"
    if cx < 425: return "代码"
    if cx < 620: return "个股"
    if cx < 815: return "涨停时间"
    if cx < 1000: return "流通市值"
    if cx < 1135: return "成交额"
    return "涨停关键词"

# 归一化（去空格、+、全角/半角括号差异），用于同列去重
def norm(t):
    return re.sub(r"[\s+＋（）()·\-．\u3000]", "", t)

records = []
for row in rows:
    row.sort(key=lambda d: d["x0"])
    joined = "".join(d["text"] for d in row)
    if row[0]["y0"] < 230:
        continue
    if "不含" in joined or "统计数据" in joined:
        continue
    items = {}
    for d in row:
        c = col_of((d["x0"] + d["x1"]) / 2)
        items.setdefault(c, []).append(d)
    records.append((sector_for(min(d["y0"] for d in row)), items))

def get_col(items, col):
    lst = items.get(col)
    if not lst:
        return ""
    lst = sorted(lst, key=lambda d: d["x0"])
    # 1) y 差 < 30 视为同一识别块（同一行文字被识别多次）；块内保留最长文本，
    #    丢弃归一化后是其子串的框（如 "完整关键词" 与 "完整关键词"+"（华为）" 拆分残留）
    groups = []
    for d in lst:
        placed = False
        for g in groups:
            if abs(g[0]["y0"] - d["y0"]) < 30:
                g.append(d); placed = True; break
        if not placed:
            groups.append([d])
    picked = []
    for g in groups:
        g.sort(key=lambda d: len(norm(d["text"])), reverse=True)
        kept = []
        for d in g:
            n = norm(d["text"])
            if any(n in norm(k["text"]) for k in kept):
                continue
            kept.append(d)
        picked.extend(kept)
    picked.sort(key=lambda d: d["x0"])
    # 2) 归一化完全相同的兜底去重
    out = []
    seen = set()
    for d in picked:
        n = norm(d["text"])
        if n and n in seen:
            continue
        if n:
            seen.add(n)
        out.append(d)
    return "".join(d["text"] for d in out)

def row_items(items):
    out = []
    for v in items.values():
        out.extend(v)
    return out

out_rows = []
issues = []
for sector, items in records:
    code = get_col(items, "代码").strip()
    name = get_col(items, "个股").strip()
    if not code and not name:
        continue
    row = [sector, get_col(items, "板数").strip(), code, name,
           get_col(items, "涨停时间").strip(), get_col(items, "流通市值").strip(),
           get_col(items, "成交额").strip(), get_col(items, "涨停关键词").strip()]
    out_rows.append(row)
    # 行内 y 跨度诊断：>100px 提示可能有噪声串行
    ys = [d["y0"] for d in row_items(items)]
    if ys and (max(ys) - min(ys)) > 100:
        issues.append(("跨度>100px", name, round(max(ys) - min(ys))))

print(f"[3/6] 重建 {len(out_rows)} 行")
c = Counter(r[0] for r in out_rows)
for k, v in sorted(c.items()):
    print(f"      {k}: {v}")
if issues:
    print("  [跨度诊断]")
    for it in issues:
        print("   !", it)

# 全表校验
print("[4/6] 质量校验")
bad = []
for r in out_rows:
    reasons = []
    if not re.match(r"^\d{6}\.(SZ|SH|BJ)$", r[2]):
        reasons.append("代码:" + r[2])
    if not re.match(r"^\d{1,2}:\d{2}:\d{2}$", r[4]):
        reasons.append("时间:" + r[4])
    if not re.match(r"^\d+(\.\d+)?$", r[5]) or not re.match(r"^\d+(\.\d+)?$", r[6]):
        reasons.append("数值:" + r[5] + "/" + r[6])
    if r[1] and not re.match(r"^(\d+天\d+板|1)$", r[1]):
        reasons.append("板数:" + r[1])
    if not r[7]:
        reasons.append("空关键词")
    if reasons:
        bad.append((r[3], "; ".join(reasons)))
for name, why in bad:
    print("  !", name, "->", why)

# 中间结果落盘（供定点复核脚本使用）
with open(JSON_PATH, "w", encoding="utf-8") as f:
    json.dump({"boxes": clean, "sectors": sectors, "rows": out_rows}, f, ensure_ascii=False)

total = sum(v for k, v in c.items() if k)
print(f"[5/6] 行数 {len(out_rows)} vs 期望 {EXPECT}，总耗时 {time.time()-t_start:.1f}s")
if EXPECT and len(out_rows) != EXPECT:
    print(f"  [WARN] 行数不一致！差异 {len(out_rows)-EXPECT}")

with open(OUT, "w", encoding="utf-8-sig", newline="") as f:
    w = csv.writer(f)
    w.writerow(["板块", "板数", "代码", "个股", "涨停时间", "流通市值(亿元)", "成交额(亿元)", "涨停关键词"])
    w.writerows(out_rows)
print(f"[6/6] 已写入: {OUT}")
