File size: 6,109 Bytes
e0ea7df
 
 
 
ae8b600
e5e756a
ae8b600
e0ea7df
e5e756a
e0ea7df
 
ae8b600
 
 
 
e0ea7df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cc826a1
ae8b600
e0ea7df
ae8b600
e0ea7df
 
e5e756a
ae8b600
e0ea7df
ae8b600
 
e0ea7df
 
ae8b600
e0ea7df
ae8b600
e0ea7df
 
 
 
 
 
ae8b600
 
 
 
 
 
 
 
 
 
cc826a1
ae8b600
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
"""OCR 插件核心。

对 PaddleOCR 3.x predict 接口做一层薄封装:解析识别结果、按需写 run 事件。
"""
import logging
import time
from pathlib import Path
from typing import Any, Optional

from app.plugins.run_log import get_run_log_service
from app.utils.ocr_engine import get_ocr_engine, get_ocr_model_info

logger = logging.getLogger(__name__)


def parse_result(result) -> list[str]:
    """从 PP-OCRv6 predict 返回中抽取文本行。"""
    if not result:
        return []
    lines: list[str] = []
    for page in result:
        if not isinstance(page, dict):
            continue
        # 仅保留非空字符串
        lines.extend(t for t in (page.get("rec_texts") or []) if t)
    return lines


def _normalize_poly(poly) -> list | None:
    """把 PP-OCRv6 各种检测框格式统一为 4 点 polygon [[x,y], ...]。

    rec_boxes 给的是 [xmin,ymin,xmax,ymax],polys/dt_polys 已经是 [[x,y], ...]。
    """
    if poly is None:
        return None
    arr = poly.tolist() if hasattr(poly, "tolist") else poly
    if not isinstance(arr, (list, tuple)) or len(arr) == 0:
        return None
    # 扁平 [xmin,ymin,xmax,ymax] 转四点
    if len(arr) == 4 and all(isinstance(v, (int, float)) for v in arr):
        x1, y1, x2, y2 = arr
        return [[x1, y1], [x2, y1], [x2, y2], [x1, y2]]
    return [list(p) for p in arr]


def extract_boxes(result) -> list[dict]:
    """抽取每行文本的检测框坐标与置信度,供前端叠加显示。"""
    boxes: list[dict] = []
    if not result:
        return boxes
    for page in result:
        if not isinstance(page, dict):
            continue
        texts = page.get("rec_texts") or []
        scores = page.get("rec_scores") or []
        # rec_boxes 是矩形,其他字段已是 polygon
        polys = page.get("rec_polys") or page.get("dt_polys") or []
        rects = page.get("rec_boxes") or []
        for i, text in enumerate(texts):
            if not text:
                continue
            # 优先用 polygon;没有则退回到 rec_boxes 矩形
            poly = _normalize_poly(polys[i] if i < len(polys) else None)
            if poly is None and i < len(rects):
                poly = _normalize_poly(rects[i])
            boxes.append({
                "text": str(text),
                "score": float(scores[i]) if i < len(scores) else None,
                "poly": poly,
            })
    return boxes


def extract_text(image_path: str, run_id: Optional[str] = None) -> dict:
    """从本地图片文件识别文字,可选地写 run 事件。

    返回字段: success / text / lines / boxes / raw_result / model_info / timings / error
    """
    run_service = get_run_log_service() if run_id else None

    def emit(stage: str, message: str, *, level: str = "info", detail: str | None = None) -> None:
        if run_service:
            run_service.add_event(run_id, stage, message, level=level, detail=detail)

    path = Path(image_path)
    if not path.exists():
        msg = f"图片文件不存在: {image_path}"
        emit("ocr", msg, level="error")
        return _fail(msg, run_service)

    start = time.perf_counter()
    try:
        emit("model_load", "OCR 引擎已加载")
        engine = get_ocr_engine()
        emit("recognize", "开始识别")
        result = engine.predict(str(path))
        lines = parse_result(result)
        boxes = extract_boxes(result)
        timings = {"total_ms": round((time.perf_counter() - start) * 1000, 2)}

        text = "\n".join(lines)
        emit("recognize", f"识别完成: {len(lines)} 行")
        return {
            "success": True,
            "text": text,
            "lines": lines,
            "boxes": boxes,
            "raw_result": result,
            "model_info": get_ocr_model_info(),
            "timings": timings,
            "error": None,
        }
    except Exception as e:
        logger.exception("OCR 识别失败: %s", image_path)
        emit("ocr", f"OCR 识别失败: {e}", level="error")
        return _fail(str(e), run_service)


def _fail(error: str, run_service) -> dict:
    """构造统一的失败返回;为减少模板,只暴露必要字段。"""
    return {
        "success": False,
        "text": "",
        "lines": [],
        "boxes": [],
        "raw_result": None,
        "model_info": get_ocr_model_info() if run_service else None,
        "timings": None,
        "error": error,
    }


def get_model_info_with_providers() -> dict:
    """聚合返回模型信息与 provider 可用性,前端一次拿全。"""
    from app.utils.ocr_engine import probe_ocr_providers
    return {
        "model_info": get_ocr_model_info(),
        "providers": {k: v.value for k, v in probe_ocr_providers().items()},
    }


# ---- 向后兼容层(被 content 调用)----
class OCRService:
    """薄壳:保持 `get_ocr_service().extract_text_from_file(path)` 旧接口可用。

    内部直接调用模块级 extract_text,不再额外缓存引擎实例。
    """

    def extract_text_from_file(self, image_path: str) -> dict:
        return extract_text(image_path)

    def extract_text_from_base64(self, image_base64: str, suffix: str = ".jpg") -> dict:
        import base64
        import tempfile
        try:
            content = base64.b64decode(image_base64)
        except Exception as e:
            return {"success": False, "text": "", "lines": [], "error": f"Base64 解码失败: {e}"}
        if suffix and not suffix.startswith("."):
            suffix = f".{suffix}"
        with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
            tmp.write(content)
            tmp_path = tmp.name
        try:
            return self.extract_text_from_file(tmp_path)
        finally:
            Path(tmp_path).unlink(missing_ok=True)


_ocr_service: OCRService | None = None


def get_ocr_service() -> OCRService:
    """返回 OCRService 单例,供 content 等历史调用方使用。"""
    global _ocr_service
    if _ocr_service is None:
        _ocr_service = OCRService()
    return _ocr_service